taejoon89/openpath
117
1#!/usr/bin/env python2"""AMC-HCC-ST 벤치마크 — Asan Medical Center HCC Visium spatial-transcriptomics 코호트(비공개).3각 spot에서 WSI 패치 추출 → FM 임베딩(CLS) → PCA → Ridge 회귀로 상위 HVG 발현 예측4→ Pearson 상관(유전자평균). Leave-one-patient-out CV.5 6사용:7 PYTHONPATH=OpenPath:eval venv_eva/bin/python eval/st_bench.py \8 --backbone openmidnight # 참조: OpenMidnight9 ... --backbone phikon10 ... --backbone openpath --weights data/runs/openpath_run/eval/training_316250/teacher_checkpoint.pth11"""12import os, sys, json, glob, argparse, re13import numpy as np, pandas as pd14import torch, openslide15from PIL import Image16from sklearn.linear_model import Ridge17from sklearn.decomposition import PCA18from sklearn.preprocessing import StandardScaler19from scipy.stats import pearsonr20import torchvision.transforms as T21 22ROOT = os.environ.get("OPENPATH_ROOT", os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) # repo 루트(eval/의 상위)23DATA = os.environ.get("ST_ROOT", f"{ROOT}/data/st_bench") # AMC-HCC-ST 코호트(비공개; 코드만 공개)24IMAGENET_MEAN = (0.485, 0.456, 0.406); IMAGENET_STD = (0.229, 0.224, 0.225)25 26 27def build_backbone(name, weights):28 sys.path.insert(0, f"{ROOT}/OpenPath"); sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))29 import openpath_eva_backbone as B30 ref = {"uni": (B.build_uni, 1024), "uni2": (B.build_uni2, 1536),31 "gigapath": (B.build_gigapath, 1536),32 "virchow2": (B.build_virchow2, 1280),33 "phikon": (B.build_phikon, 1024),34 "openmidnight": (B.build_openmidnight, 1536)}35 if name in ref:36 fn, dim = ref[name]; return fn(), dim37 # openpath: 우리 teacher_checkpoint(dinov2 teacher 포맷) 로더38 return B.build_openpath(weights), 153639 40 41def slide_dirs():42 return sorted([d for d in glob.glob(f"{DATA}/*") if os.path.isdir(d)])43 44 45def patient_of(slide_dir):46 b = os.path.basename(slide_dir) # 예: pt<N>-<M> 또는 pt<N>47 m = re.match(r"(pt\d+)", b)48 return m.group(1) if m else b # pt<N>-<M>, pt<N>-<K> → 동일 환자 pt<N>49 50 51def load_one(slide_dir):52 exp_f = glob.glob(f"{slide_dir}/*.spatial.data.exp.csv")53 pos_f = glob.glob(f"{slide_dir}/*.tissue_positions_fullres.csv")54 sf_f = glob.glob(f"{slide_dir}/*.scalefactors_json.json")55 wsi_f = glob.glob(f"{slide_dir}/*_p0.tif")56 if not (exp_f and pos_f and sf_f and wsi_f):57 return None58 exp = pd.read_csv(exp_f[0], index_col=0) # spots × genes (log-norm)59 pos = pd.read_csv(pos_f[0], index_col=0) # barcode → pxl coords60 sf = json.load(open(sf_f[0]))61 diam = int(round(sf["spot_diameter_fullres"])) # ~147px62 common = exp.index.intersection(pos.index)63 exp = exp.loc[common]; pos = pos.loc[common]64 return dict(dir=slide_dir, exp=exp, pos=pos, diam=diam, wsi=wsi_f[0])65 66 67def _load_patches(sl):68 """224 uint8 패치 (N,224,224,3). 슬라이드별 캐시 재사용(체크포인트마다 동일 패치)."""69 import numpy as np70 cache = f"{DATA}/_pcache/{os.path.basename(sl['dir'])}.pt"71 if os.path.exists(cache):72 try: return torch.load(cache)73 except Exception: pass74 sldx = openslide.OpenSlide(sl["wsi"]); d = sl["diam"]75 rs = T.Resize((224, 224))76 xs = sl["pos"]["pxl_x_in_fullres"].values.astype(int)77 ys = sl["pos"]["pxl_y_in_fullres"].values.astype(int)78 plist = []79 for x, y in zip(xs, ys):80 patch = sldx.read_region((int(x - d // 2), int(y - d // 2)), 0, (d, d)).convert("RGB")81 plist.append(torch.from_numpy(np.asarray(rs(patch)))) # (224,224,3) uint882 sldx.close()83 patches = torch.stack(plist)84 os.makedirs(os.path.dirname(cache), exist_ok=True)85 tmp = cache + f".tmp{os.getpid()}"86 torch.save(patches, tmp); os.replace(tmp, cache) # atomic87 return patches88 89 90@torch.no_grad()91def embed_slide(sl, model, device, bs=256):92 patches = _load_patches(sl) # (N,224,224,3) uint893 mean = torch.tensor(IMAGENET_MEAN).view(1, 3, 1, 1)94 std = torch.tensor(IMAGENET_STD).view(1, 3, 1, 1)95 embs = []96 for i in range(0, len(patches), bs):97 b = patches[i:i + bs].permute(0, 3, 1, 2).float().div_(255.0) # (B,3,224,224)98 b = ((b - mean) / std).to(device)99 embs.append(model(b).float().cpu().numpy())100 return np.concatenate(embs, 0) # (n_spots, dim)101 102 103def top_hvg(exp_all, k=50):104 # 학습셋 전체 log-norm 발현서 분산 상위 k 유전자105 v = exp_all.var(axis=0)106 return v.sort_values(ascending=False).index[:k].tolist()107 108 109def main():110 ap = argparse.ArgumentParser()111 ap.add_argument("--backbone", required=True, choices=["openpath","openmidnight","phikon","uni","uni2","gigapath","virchow2"])112 ap.add_argument("--weights", default=None)113 ap.add_argument("--k-genes", type=int, default=50)114 ap.add_argument("--pca", type=int, default=256)115 ap.add_argument("--alpha", type=float, default=100.0)116 ap.add_argument("--tag", default=None)117 args = ap.parse_args()118 device = "cuda"119 120 model, dim = build_backbone(args.backbone, args.weights)121 model = model.to(device).eval()122 123 dirs = slide_dirs()124 print(f"[st]slides={len(dirs)} backbone={args.backbone} dim={dim}", flush=True)125 slides = []126 for d in dirs:127 sl = load_one(d)128 if sl is None: print(f" skip {os.path.basename(d)} (파일 부족)"); continue129 sl["emb"] = embed_slide(sl, model, device)130 sl["pat"] = patient_of(d)131 slides.append(sl)132 print(f" {os.path.basename(d)}: spots={len(sl['pos'])} emb={sl['emb'].shape} pat={sl['pat']}", flush=True)133 134 # 공통 유전자135 genes = slides[0]["exp"].columns136 for s in slides[1:]: genes = genes.intersection(s["exp"].columns)137 genes = list(genes)138 print(f"[st]공통 유전자 {len(genes)}", flush=True)139 140 pats = sorted(set(s["pat"] for s in slides))141 # leave-one-patient-out142 per_gene_corr = []143 for held in pats:144 tr = [s for s in slides if s["pat"] != held]145 te = [s for s in slides if s["pat"] == held]146 Xtr = np.concatenate([s["emb"] for s in tr], 0)147 Ytr = np.concatenate([s["exp"][genes].values for s in tr], 0)148 Xte = np.concatenate([s["emb"] for s in te], 0)149 Yte = np.concatenate([s["exp"][genes].values for s in te], 0)150 # HVG는 학습셋서 선택151 hvg_idx = np.argsort(-Ytr.var(0))[:args.k_genes]152 Ytr_h, Yte_h = Ytr[:, hvg_idx], Yte[:, hvg_idx]153 # 표준화 + PCA + Ridge154 sc = StandardScaler().fit(Xtr)155 Xtr2, Xte2 = sc.transform(Xtr), sc.transform(Xte)156 p = PCA(n_components=min(args.pca, Xtr2.shape[1])).fit(Xtr2)157 Xtr3, Xte3 = p.transform(Xtr2), p.transform(Xte2)158 reg = Ridge(alpha=args.alpha).fit(Xtr3, Ytr_h)159 pred = reg.predict(Xte3)160 cors = []161 for g in range(Yte_h.shape[1]):162 if Yte_h[:, g].std() < 1e-8 or pred[:, g].std() < 1e-8:163 cors.append(0.0)164 else:165 cors.append(pearsonr(Yte_h[:, g], pred[:, g])[0])166 m = float(np.nanmean(cors))167 per_gene_corr.append(m)168 print(f" [fold {held}] test_spots={Yte.shape[0]} meanPearson={m:.4f}", flush=True)169 170 overall = float(np.mean(per_gene_corr))171 tag = args.tag or args.backbone172 print(f"[AMC-HCC-ST RESULT] backbone={tag} | LOPO mean Pearson = {overall:.4f} (folds={len(pats)}, HVG={args.k_genes})", flush=True)173 174 175if __name__ == "__main__":176 main()177 