
"Unimodal보다 multimodal 설정에서 모델 성능이 우수하며, full-ST, 즉 좌표 정보를 보존한 데이터를 활용했을 때가 best이다"
를 보이기 위해 총 네 가지 모달리티 설정,
- Unimodal (ST-only)
- Unimodal (Image-only)
- Multimodal (Bulk)
- Multimodal (Full ST)
에서 모델 성능을 비교했다.
코드 참고: https://github.com/EWHA-CAPSTONE-VISION/team_project_repo
이미지 모달리티와 ST 모달리티의 기여도를 비교하기 위해 encoder on/off 방식의 ablation 실험 구조를 구현하였다.
Default 모델 코드 model.py는 항상 이미지 encoder와 ST encoder를 모두 사용한 뒤 fusion module을 거쳐 MIL classification을 수행하는 구조였으나, ablation 버전에서는 use_image, use_st 플래그를 추가하여 조건을 동일한 학습 및 평가 파이프라인에서 실행할 수 있도록 하였다.
1. 모델 코드
기본 모델인 models/models.py의 MultimodalMILModel은 ImageEncoder, SpatialEncoder, SpotFusionModule, MILAttentionPooling, classifier를 모두 생성하고,
forward에서 항상 이미지 feature와 ST feature를 계산한 뒤 fusion을 수행한다. 즉, 기본 구조는 다음과 같다.
image patch → ImageEncoder → img_head
gene expression + coords → SpatialSTEncoder
img_feat + st_feat → SpotFusionModule
spot embeddings → MILAttentionPooling → classifier
반면 ablation-support 모델인 models/model_ablation.py에서는 MultiModalMILModel 생성자에 use_image=True, use_st=True 플래그가 추가되었다. 이 두 값은 최소 하나 이상 켜져(True) 있어야 하며, 둘 다 꺼지는 잘못된 실험 설정은 assert use_image or use_st로 차단한다. 모델 내부에서는 이들 플래그에 따라 조건부로 인코더를 생성한다.
또한 fusion 모듈은 두 모달리티가 모두 활성화된 경우에만 생성된다. 이미지 단독(Img-only 모델) 또는 ST 단독(ST-only 모델) 실험에서는 fusion이 필요하지 않으므로, self.fusion=None으로 둔다. Forward 단계에서도 동일한 라우팅을 적용하며, 위 내용들을 정리하면 다음과 같다.
image + ST → fusion(img_feat, st_feat)
image only → spot_embeds = img_feat
ST only → spot_embeds = st_feat
즉, ablation 모델은 최종적으로 항상 동일한 형태의 spot_embeds를 MIL pooling에 전달하되, 그 spot_embeds가 어떤 인코더에서 왔는지만이 실험 조건에 따라 달라지도록 설계되었다. 덕분에 encoder 이후의 MIL pooling과 classifier는 모든 실험에서 동일하게 유지된다.
2. YAML 설정
Ablation 실험용 설정은 configs/train_ablation.yaml에 분리되어 있다. 기존 configs/train.yaml에는 멀티모달 학습 설정만 존재하지만, ablation을 위해 model 항목 아래에 다음 두 옵션이 추가되었다.
model:
fusion_option: concat
top_k_genes: 512
use_image: true
use_st: true
위 값을 바꿈으로써, 동일한 코드를 이용해 아래 세 가지 실험을 수행할 수 있다.
use_image: true, use_st: true → Image + ST 멀티모달
use_image: true, use_st: false → Image-only
use_image: false, use_st: true → ST-only
3. Train 코드
Ablation 학습은 train_ablation.py에서 수행된다. load_config() 함수는 configs/train_ablation.yaml을 읽고, use_iamge, use_st 값을 CONFIG dictionary에 저장한다.
"use_image": cfg["model"].get("use_image", True),
"use_st": cfg["model"].get("use_st", True),
이후 assert CONFIG[“use_image”] or CONFIG[“use_st”]로 최소 하나의 모달리티가 활성화되었는지 확인한다.
Train 코드의 핵심은 encode_spots_chunkwise() 함수이며, 이 함수는 하나의 WSI에 포함된 여러 spot을 batch 크기로 나누어 처리하면서, 설정된 모달리티에 필요한 입력만 사용한다.
use_image=True → batch["images"] 사용
use_st=True → batch["expr"], batch["coords"] 사용
즉, 각 chunk마다 아래와 같이 라우팅한다.
image + ST:
img_encoder → img_head
st_encoder
fusion(img_feat, st_feat)
image only:
img_encoder → img_head
ST only:
st_encoder
이렇게 얻은 spot_embeds는 이후 모든 실험에서 동일하게 model.mil_pooling()과 model.classifier()로 전달된다. 즉, 모달리티 비교 실험에서 encoder와 fusion 입력만 달라지고, MIL pooling 및 classifier는 동일하게 유지되도록 구현하였다.
4. Test 코드
평가 및 XAI 분석용 ablation 코드는 test_ablation.py에 구현되어 있다.
parser.add_argument("--use_image", action="store_true")
parser.add_argument("--use_st", action="store_true")
parser.add_argument("--use_image", action="store_true")
parser.add_argument("--use_st", action="store_true")
학습은 YAML 기반으로 모달리티를 설정하지만, 테스트 코드는 command-line argument 방식으로 encoder on/off를 지정하며, defaut는 멀티모달로 하여 아무 옵션도 지정되지 않을 경우 멀티모달 실험이 진행되도록 처리하였다.
images = batch["images"].to(args.device) if args.use_image else None
expr = batch["expr"].to(args.device) if args.use_st else None
coords = batch["coords"].to(args.device) if args.use_st else None
테스트 시에도 모델 생성자에 동일하게 use_image, use_st가 전달된다. 또한 batch에서 필요한 모달리티만 GPU로 올리도록 구현하였다.
더불어 ST 모달리티가 활성화되어 있을 때만 gene attention을 반환하도록 return_gene_attn=bool(args.use_st)를 사용한다. 따라서 image-only 실험에서는 gene-level XAI를 계산하지 않고, ST 또는 multimodal 실험에서만 gene attention 기반 top gene 분석을 수행한다. 반대로 top patch 저장은 image 모달리티가 켜진 경우에만 수행된다.
5. DataLoader와 Modality 정렬
DataLoader인 dataset/loader.py는 기본적으로 image patch, gene expression, spatial coordinate를 모두 반환한다. ST spot과 image patch는 barcode 기준으로 정렬되며, 반환되는 주요 항목은 다음과 같다.
images → image encoder 입력
expr → ST encoder의 gene expression 입력
coords → ST encoder의 spatial token 입력
label → WSI-level classification label
즉, dataloader는 모든 모달리티를 준비하지만, 실제 train 및 test 단계에서 use_image, use_st 플래그에 따라 필요한 입력만 선택적으로 사용한다. 이 방식은 데이터 split과 sample 구성을 동일하게 유지하면서 인코더의 on/off만을 수행하므로, 모달리티별 성능 비교의 공정성을 높인다.
6. 정리
결과적으로, 본 ablation 구현은 use_image와 use_st 두 개의 플래그를 중심으로 구성된다. YAML에서는 학습 실험 조건을 설정하고, train code는 해당 설정을 읽어 모델 생성 및 spot encoding 라우팅에 반영한다. test code에서는 command-line argument로 동일한 조건을 재현하며, 활성화된 모달리티에 맞춰 patch-level XAI 또는 gene-level XAI를 선택적으로 저장한다. 모델 내부에서는 encoder 생성, freeze 처리, fusion 생성, forward 라우팅이 모두 플래그 기반으로 동작하므로, image-only, ST-only, image+ST 실험을 동일한 MIL backbone 위에서 비교할 수 있도록 구현하였다.
'capstone_2025_fall' 카테고리의 다른 글
| [데이터셋] HEST-1k Overview (1) | 2025.12.01 |
|---|---|
| [논문 리뷰] Celcomen: spatial causal disentanglement for single-cell and tissue perturbation modeling (1) | 2025.11.27 |
| [데이터셋] HEST-1k 사용 방법 (0) | 2025.11.26 |
| 25-2 졸업 프로젝트 (0) | 2025.11.18 |