Advancing Generalizable Tumor Segmentation with Anomaly-Aware Open-Vocabulary Attention Maps and Frozen Foundation Diffusion Models (CVPR 2025)

Treeboy·2025년 8월 24일

CVPR준비

목록 보기
6/14

2025 CVPR 을 보다보니 foundation model 의 zero-shot application 에 참 관심이 많다는게 느껴집니다. 이번 논문은 의료 영상의 diffusion model 인 MAISI (3D) 를 활용해서 본 적 없는 도메인에서 (!!) zero-shot tumor segmentation 을 수행하였습니다. 뭐, 예를 들어 CT brain tumor segmentation 을 시작하고 싶은데, annotation 이 된 데이터가 없으니까 일단 MRI brain tumor segmentation 을 한다음에 그걸 좀 수정해서 쓴달까요?

Methods

DiffuGTS

논문에서 제시한 DiffuGTS 는 간단하게 네 부분으로 나눠볼 수 있겠습니다:

  1. Diffusion model (MAISI)
  2. Adaptor
  3. AOVA map generation
  4. Mask refinement

MAISI

MAISI 는 frozen diffusion model 이며, 3D CT & MRI (무려 512x512x768 !!) 을 생성하기 위해 만들어졌습니다. 또한, CT 의 127 가지 anatomical structure 의 segmentation 을 활용한 ControlNet 컴포넌트도 가지고 있습니다.

훈련 stages 를 가져왔는데, 사실 엄청 단순합니다.

  • VAE 훈련하기 (KL loss 는 standard deviation 을 0.9~1.1 로 constrain, VQVAE 는 아닌듯?)
  • Diffusion 훈련하기 (latent diffusion!)
  • Segmentation mask 를 condition 으로 ControlNet 훈련

그러니까 얘는 CT 랑 MR 만들기 위한 Frozen foundation model 이라고 생각하면 됩니다. 훈련에 활용된 데이터는 아래 표에서 참고하시길 바랍니다.

Adapter

위의 MAISI는 generation 에 특화되어있기 때문에, segmentation 같은 discriminative 한 task 에 적용하기 위해서는 representation 을 조금 손봐야합니다. 이를 위해서 저자들은 VAE 의 각 layer 마다 learnable feature adapater (linear transformation 두번과 residual connection) 을 활용해 원본 feature 을 약간 수정합니다.

Fl=αTl(Fl)+(1α)FlF_l =\alpha T_l(F_l)+(1-\alpha)F_l

(Default α=0.1\alpha=0.1)

AOVA

AOVA 는 Anomaly-aware open-vocabulary attention 을 줄인 말입니다. 하나씩 해석을 해보자면,

anomaly-aware - 영상에서의 이상치 (tumor) 을 인식할 수 있음
open-vocabulary - 학습하지 않았던 데이터도 포함
attention - attention map 을 활용

그럼 구조를 자세히 봐봅시다.

Structure

Figure 2 의 왼쪽 및 부분

Text prompt 와 image feature Fl,l{1,2,3}F_l, l\in\{{1,2,3}\} 사이의 cross attention

Text prompt 는 두가지 경우로 나뉩니다:

  • Normal: "A normal CT scan/MRI of {organ name}
  • Abnormal: "An abnormal CT scan/MRI of {disease name}

이 prompt 는 Clinical-bert 에 입력되어 text embedding eRN×de \in \R^{N\times d} 으로 변환되고, MLP 를 활용해 각 image feature level 의 dimension 에 맞게 차원축소를 진행합니다. 마지막으로, key projection matrix 인 WKl\mathcal{W}^l_K 에 곱하여 attribution keys KlRN×ClK_l \in\R^{N\times C_l} 을 얻습니다.

Key 를 얻었으니 query 만 있으면 attention map 을 구할 수 있겠지요? Query 도 비슷하게 query projection matrix 인 WQl\mathcal{W}^l_Q 에 곱하여 pixel queries QlRHl×Wl×Dl×ClQ_l \in\R^{H_l\times W_l \times D_l\times C_l} 을 구합니다.

그러면, attention map A(Ql,Kl)A(Q_l, K_l)

A(Ql,Kl)=Softmax(QlKlTCl)RHl×Wl×Dl×NA(Q_l, K_l)=\text{Softmax}(\frac{Q_l K_l^T}{\sqrt{C_l}}) \in\R^{H_l\times W_l \times D_l\times N}

이 되겠습니다. 이 때, training categories NN 은 13개 (Tumor 7종류, 정상 장기 6종류) 이며, 각 카테고리에 대해 attention map (즉 segmentation mask!) 을 만든다고 생각하시면 됩니다. 이 attention map 을 layer 마다 구한 뒤 아래와 같이 aggregate 한 것이 AOVA maps 입니다.

MI,e=l=13RI(A(Ql,Kl))RH×W×D×NM_{\mathcal{I},e}=\sum_{l=1}^{3}{\textbf{RI}}(A(Q_l,K_l))\in\R^{H\times W\times D\times N}

RI()\textbf{RI}(\cdot) 는 bilinear interpolation 을 활용한 reshape operation 이고요.

Training Objectives

  • Anomaly Classification Loss (Lano\mathcal{L}_{ano})

이미지 전체가 '정상' 인지 '비정상' 인지에 대한 binary classification 을 수행

각 AOVA map (총 N개) 를 MLP 에 넣어 class embedding giRd,i[1,N]g_i\in\R^d,i\in[1,N] 을 얻고, 이걸 그대로 maxpool 해서 anomaly score ascore=Sigmoid(MaxPool(g)):Rd[0,1]ascore=\text{Sigmoid}(MaxPool(g)):\R^d\rarr[0,1] 를 얻습니다.

  • CLIP-Style Contrastive Learning Loss (Lsim\mathcal{L}_{sim})

Class embedding (gig_i) 과 text embedding (eje_j) 의 semantic alignment 강화

Lsim=1Ni=1Nlogexps(gi,ej)j=1Nexps(gi,ej)L_{sim} = - \frac{1}{N} \sum_{i=1}^N \log \frac{\exp s(g_i, e_j)}{\sum_{j=1}^N \exp s(g_i, e_j)}

여기서, s(gi,ej)s(g_i, e_j) 는 클래스 임베딩과 텍스트 임베딩 사이의 유사도 점수입니다. 이 유사도 점수로 softmax 를 하게 되는데, AOVA map 을 활용한 예측과 정답 텍스트 임베딩이 높은 유사도를 가지도록 (즉 비정상 AOVA 와 비정상 textual prompt 가 비슷하도록) contrastive loss 를 줍니다.

  • Dice Loss (Ldice\mathcal{L}_{dice})

Segmentation loss, 근데 이거 zero-shot 이라매? ㅋㅋ

Partially labeled segmentation annotations 와 AOVA map 과의 Dice Loss 를 구합니다.

Mask Refinement

자. 이제 위 사진의 (b) 까지 왔습니다. 그런데, tumor attention map 이 조금 coarse 하다고 합니다. 이걸 MAISI 를 활용해 pseudo-healthy organ 을 impainting 하여 더 정확한 anomaly segmentation map 을 만들겠다고 합니다.

원래 MAISI 는 ControlNet 을 활용해서 그 장기를 생성할 수 있긴 합니다. 하지만, tumor 이 없는 부위는 복원하면 안되기 때문에, tumor 이 있을 것 같은 mask (아마 AOVA 그대로 썼을겁니다) 안에서 latent impainting 을 수행합니다.

먼저, 이미지를 인코더에 넣어 latent representation zz 를 구합니다. 여기서 reverse process 를 수행하는데, 아까 활용한 mask 를 downsampling 하여 masked regeneration 을 수행합니다.

zt1=(1D(ML))zt1other+D(ML)zt1tumorz_{t-1} = (1 - D(M_L)) \otimes z^{other}_{t-1} + D(M_L) \otimes z^{tumor}_{t-1}

DD가 downsampling, MLM_L 이 마스크입니다. 그렇게 해서 얻은 Pseudo-healthy equivalent 의 latent embeddings z0z_0 을 다시 디코더에 넣어 pseudo-healthy image H=VD(z0)\mathcal{H}=V_D(z_0) 을 생성하고, 아래의 두 방법을 활용하여 최종 segmentation 을 수행합니다.

  • Pixel-level residual learning PrP_r

원본 이미지와 pseudo-healthy 이미지의 차이

  • Feature-level residual learning FrF_r

Latent embeddings 의 차이 fr=zz0f_r=z-z_0 인데, 이 때 frRh×w×d×cf_r \in \R^{h\times w\times d\times c} 여서 channel 의 정보를 aggregate 해야 합니다. 저자들은 단순히 channel-wise averaging 하지 않고, text embedding eje_jcc channel 로 linear projection 시킨 뒤 frf_r 과 내적을 해버리고, 원본 이미지 크기로 upsampling 합니다.

또한, tumor segmentation map 을 활용해서 dice loss 로 추가적인 지도학습을 수행합니다. (아니 이러면 zero-shot 맞냐니까?)

최종 anomaly segmentation map 은 PrP_rFrF_r 을 평균내서 만듭니다.

Main

1. Generalization to unseen tumors

KiTS23 데이터셋과 MSD에서 가져온 5가지 종양(간, 대장, 췌장, 폐, 간혈관 종양) segmentation 데이터셋, 총 6가지 종양 카테고리에서 leave-one-out 으로 실험했습니다. 즉 6가지 카테고리중 5가지 카테고리로 훈련한 뒤, 나머지 한 카테고리에서 테스트 한거죠.

성능은 그럭저럭 괜찮았습니다. 모든 카테고리에서 SOTA 였고, runner-up 알고리즘인 Malenia 보다 DSC가 평균 4 point 올랐습니다. 다만 그렇게 impressive 하진 않습니다.

2. Generalization to unseen modality

이번엔 CT 에서 훈련한다음에 MRI 로 가져갔습니다!! (이번엔 좀 흥미롭네요).

이번엔 성능 격차가 좀 눈에 들어옵니다. 최소 DSC가 27 point 증가하는 모습을 보였으며, 50정도면 annotation 의 시작점으로는 활용해볼법한 궤도에 들어갔다고 생각합니다.

다만, 저자들이 pseudo-healthy 가 퀄리티가 좋다고 하는 점은 동의가 어렵네요.. 별로 뇌 같이 생기진 않았고 그냥 안에 회색칠만 한 기분입니다.

3. Computational overhead

이건 사실 단점이라고도 봅니다. MAISI 를 frozen 상태로 두었기 때문에 trainable params 는 훨씬 줄었지만, 그래도 foundation model 을 활용해야 하기 때문에 FLOPs 는 훨씬 높습니다. Knowledge distillation 을 해볼 수 있다고는 하지만, 이 연구의 범위를 벗어났다고 선을 긋습니다.

이제 ablation study 를 봅시다.

Ablation studies

Ablation studies 는 CT의 5개 category (KiTS23, MSD) 에서 훈련한 뒤 MSD Liver (CT), MSD Brain (MRI) 에서 테스트를 진행했습니다.

1. Adapters

Adapter 을 증명하는 실험입니다. 3개의 대조군과 비교를 했는데,

  • nnUNet: train segmentation from scratch
  • w/o Adapter: MAISI 의 latent representation 을 그대로 활용
  • Fine-tuning: MAISI 의 latent representation 을 fine-tuning

이 때 Adapter 을 활용한 것이 MSD Liver (CT) 에서 제일 좋은 성능을 보여준 것 뿐만이 아니라, MRI modality 인 MSD Brain 데이터셋에서도 높은 성능을 보여주었습니다. 특히 MRI 에서 성능 방어가 잘 된 것이 주목할만한데, nnUNet 과 fine-tuning 은 foundation model 의 rich representation 이 소실되며, adapter 을 사용하지 않는 것은 generative task 에만 특화되어서 그렇다고 설명합니다.

2. AOVA

  • Mask2Former + MR: Mask2Former 도 query-based open-vocabulary zero shot segmentation 에 활용되는 모델입니다
  • DiffuGTS (AOVA only): Mask refinement (MR) 없이 AOVA 를 그대로 사용합니다.

먼저, AOVA 혼자서는 충분히 좋은 성능을 뽑지 못하고, MR 이 필수적이라는 결과입니다. 이는 AOVA 가 latent space 에서 upsample 된 coarse 한 attention map 이기 때문인데, 이걸 pixel space 에서 한번 더 다듬어서 supervised learning 하는 것이 중요하다고 볼 수 있습니다.

3. Mask refinement loss

Pixel-level 과 feature-level 을 따로따로 했을 때의 실험 결과입니다.

Figure 5 를 보면, Pixel-level map 만을 활용했을 때는 중요한 semantic 을 놓치는 경향이 있습니다 (다 채우지 못한다던가, tumor 을 아예 놓친다던가). 반면에, feature level 만을 활용했을 때는 정교한 segmentation map 을 만들지 못합니다. 이 두가지를 적절히 활용해서 ground-truth 에 가장 가깝게 뽑을 수 있었다고 합니다.

Summary

이 논문은 medical foundation model 을 이용해서 보지 못한 modality 에서의 zero-shot segmentation 을 수행하였습니다. 다만, pseudo-healthy image 를 보았다시피 unseen modality 에 대한 영상의 이해가 부족한 듯 싶습니다. MAISI 에 brain MRI 이 들어갔을텐데, 왜 이럴까요?

profile
지식이 모자라서 논문리뷰를...

0개의 댓글