CoolFace
Modelpublic

taejoon89/openpath

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
1likes17downloads
st_bench.py177 linesDownload Raw Back to eval
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