Download scripts/common/run_cv.py from bryan7264/PANDA: direct link, hf CLI and curl.
- Browser
- Download file 8.95 kB
-
https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/run_cv.py
- Command line
-
hf download hf://bryan7264/PANDA/scripts/common/run_cv.py
-
curl -L -o run_cv.py https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/run_cv.py
8.95 kB
| """5-fold stratified CV, 3 systems x 2 variants. 5 epochs/fold (shorter than train_panda).""" | |
| from __future__ import annotations | |
| import argparse, sys, json, pickle, warnings, numpy as np, pandas as pd, torch, torch.nn.functional as F | |
| from pathlib import Path | |
| import anndata as ad, scanpy as sc, scipy.sparse as sp, yaml | |
| warnings.filterwarnings("ignore"); sc.settings.verbosity = 0 | |
| from sklearn.model_selection import StratifiedKFold | |
| from sklearn.metrics import accuracy_score, f1_score, roc_auc_score, classification_report | |
| import os as _os | |
| from pathlib import Path as _Path | |
| PANDA_ROOT = _Path(_os.environ.get("PANDA_ROOT", str(_Path(__file__).resolve().parents[2]))) | |
| sys.path.insert(0, str(PANDA_ROOT)) | |
| from panda import (PANDAEncoder, supcon_loss, vicreg_loss, hsic_biased, | |
| subcenter_angular_infonce, prototype_repulsion) | |
| ROOT = Path(str(PANDA_ROOT)) | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| def load_and_prepare(system, variant): | |
| """canonical corpus + features (PCA, optional marker channel).""" | |
| a = ad.read_h5ad(ROOT / f"data/corpus/{system}/harmonized/corpus.h5ad") | |
| stats = np.load(ROOT / f"data/corpus/{system}/harmonized/corpus_stats.npz", allow_pickle=True) | |
| pca = pickle.load(open(ROOT / f"data/corpus/{system}/harmonized/pca_basis.pkl", "rb")) | |
| hvgs = [str(g) for g in stats["shared_hvgs"]] | |
| hvg2i = {g: i for i, g in enumerate(hvgs)} | |
| common = [g for g in a.var_names.astype(str) if g in hvg2i] | |
| a_c = a[:, common].copy() | |
| sc.pp.normalize_total(a_c, target_sum=1e4); sc.pp.log1p(a_c) | |
| X = a_c.X.toarray().astype(np.float32) if sp.issparse(a_c.X) else a_c.X.astype(np.float32) | |
| Xf = np.zeros((a.n_obs, len(hvgs)), dtype=np.float32) | |
| Xf[:, np.array([hvg2i[g] for g in common])] = X | |
| Xz = np.clip((Xf - stats["mean"].astype(np.float32)) / stats["std"].astype(np.float32), -10, 10) | |
| Xpca = pca.transform(Xz).astype(np.float32) | |
| Xmark = None; marker_genes = [] | |
| if variant == "marker": | |
| marker_genes = yaml.safe_load(open(ROOT / "panda/markers.yaml"))[system] | |
| mv = np.zeros((a.n_obs, len(marker_genes)), dtype=np.float32) | |
| for j, g in enumerate(marker_genes): | |
| if g in a.var_names: | |
| col = a[:, g].X | |
| if sp.issparse(col): col = col.toarray() | |
| mv[:, j] = col.flatten().astype(np.float32) | |
| mmu = mv.mean(axis=0, keepdims=True); msig = mv.std(axis=0, keepdims=True) + 1e-6 | |
| Xmark = np.clip((mv - mmu) / msig, -5, 5).astype(np.float32) | |
| labels = a.obs["canonical_label"].astype(str).values | |
| classes = sorted(set(labels)) | |
| y = np.array([classes.index(l) for l in labels], dtype=np.int64) | |
| datasets = sorted(set(a.obs["dataset"].astype(str).values)) | |
| y_dset = np.array([datasets.index(d) for d in a.obs["dataset"].astype(str).values], dtype=np.int64) | |
| return Xpca, Xmark, y, classes, y_dset, datasets, marker_genes | |
| def train_fold(Xpca, Xmark, y, y_dset, classes, variant, tr_ix, epochs=5, batch=256, lr=1e-3, seed=0): | |
| n_classes = len(classes) | |
| n_datasets = int(max(y_dset[tr_ix].max() + 1, 1)) | |
| n_markers = Xmark.shape[1] if Xmark is not None else 0 | |
| torch.manual_seed(seed); np.random.seed(seed) | |
| model = PANDAEncoder(variant=variant, n_pca=50, n_markers=n_markers, | |
| n_classes=n_classes, n_sub=3, n_datasets=n_datasets, dropout=0.2).to(DEVICE) | |
| opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4) | |
| Xt = Xpca[tr_ix]; yt = y[tr_ix]; ydt = y_dset[tr_ix] | |
| Xmt = Xmark[tr_ix] if Xmark is not None else None | |
| rng = np.random.default_rng(seed) | |
| n = len(tr_ix) | |
| for epoch in range(epochs): | |
| stage = 0 if epoch < 1 else 1 if epoch < 3 else 2 | |
| perm = rng.permutation(n) | |
| for bstart in range(0, n, batch): | |
| idx = perm[bstart:bstart+batch] | |
| x = torch.from_numpy(Xt[idx]).to(DEVICE) | |
| xm = torch.from_numpy(Xmt[idx]).to(DEVICE) if Xmt is not None else None | |
| yy = torch.from_numpy(yt[idx]).to(DEVICE) | |
| yd = torch.from_numpy(ydt[idx]).to(DEVICE) | |
| aux = torch.zeros(len(idx), 2, device=DEVICE) | |
| lam = 0.1 if stage >= 2 else 0.0 | |
| out = model(x, aux, x_markers=xm, lam_dann=lam) | |
| z = out["z"] | |
| L = supcon_loss(z, yy, 0.1) + 1.0 * vicreg_loss(z) + 0.4 * F.cross_entropy(out["logits"], yy) | |
| if stage >= 1: | |
| L = L + 0.6 * subcenter_angular_infonce(z, yy, model.prototypes.detach().clone(), | |
| margin=0.15, temperature=0.07) | |
| if stage >= 2: | |
| L = L + F.cross_entropy(out["dom"], yd) | |
| opt.zero_grad(); L.backward(); opt.step() | |
| if stage >= 1: | |
| with torch.no_grad(): model.update_prototypes(z.detach(), yy) | |
| return model | |
| def evaluate(model, Xpca, Xmark, y, val_ix, classes): | |
| model.eval() | |
| preds, probs = [], [] | |
| Xv = Xpca[val_ix]; Xmv = Xmark[val_ix] if Xmark is not None else None | |
| with torch.no_grad(): | |
| for i in range(0, len(val_ix), 2048): | |
| xb = torch.from_numpy(Xv[i:i+2048]).to(DEVICE) | |
| xmb = torch.from_numpy(Xmv[i:i+2048]).to(DEVICE) if Xmv is not None else None | |
| aux = torch.zeros(len(xb), 2, device=DEVICE) | |
| out = model(xb, aux, x_markers=xmb, lam_dann=0.0) | |
| z = out["z"] | |
| mc = model.max_sub_cos(z) | |
| preds.append(mc.argmax(dim=1).cpu().numpy()) | |
| probs.append(F.softmax(mc / 0.07, dim=1).cpu().numpy()) | |
| preds = np.concatenate(preds); probs = np.concatenate(probs) | |
| yv = y[val_ix] | |
| acc = accuracy_score(yv, preds) | |
| f1 = f1_score(yv, preds, average="macro", zero_division=0) | |
| try: | |
| auc = roc_auc_score(np.eye(len(classes))[yv], probs, average="macro", multi_class="ovr") | |
| except Exception: | |
| auc = float("nan") | |
| rep = classification_report(yv, preds, labels=list(range(len(classes))), | |
| target_names=classes, output_dict=True, zero_division=0) | |
| return acc, f1, auc, rep | |
| def cv(system, variant, folds=5, epochs=5, seed=0): | |
| print(f"\n=== CV {system}/{variant} ({folds}-fold, {epochs} epochs) ===", flush=True) | |
| Xpca, Xmark, y, classes, y_dset, datasets, _ = load_and_prepare(system, variant) | |
| print(f"[cv] n={len(y):,} K={len(classes)}", flush=True) | |
| skf = StratifiedKFold(n_splits=folds, shuffle=True, random_state=seed) | |
| accs, f1s, aucs = [], [], [] | |
| last_rep = None | |
| for fold, (tr, va) in enumerate(skf.split(np.zeros(len(y)), y)): | |
| model = train_fold(Xpca, Xmark, y, y_dset, classes, variant, tr, | |
| epochs=epochs, seed=seed * 100 + fold) | |
| acc, f1, auc, rep = evaluate(model, Xpca, Xmark, y, va, classes) | |
| accs.append(acc); f1s.append(f1); aucs.append(auc) | |
| last_rep = rep | |
| print(f"[fold {fold+1}] acc={acc:.4f} F1={f1:.4f} AUC={auc:.4f}", flush=True) | |
| result = { | |
| "system": system, "variant": variant, "folds": folds, "epochs": epochs, "seed": seed, | |
| "n_cells": int(len(y)), "n_classes": len(classes), | |
| "per_class_report_note": "per_class_report is from the LAST fold only, not aggregated", | |
| "per_fold_acc": accs, "per_fold_f1": f1s, "per_fold_auc": aucs, | |
| "mean_acc": float(np.mean(accs)), "std_acc": float(np.std(accs)), | |
| "mean_f1": float(np.mean(f1s)), "std_f1": float(np.std(f1s)), | |
| "mean_auc": float(np.nanmean(aucs)), "std_auc": float(np.nanstd(aucs)), | |
| "per_class_report": last_rep, | |
| } | |
| out_dir = ROOT / f"discovery/{system}/{variant}" | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| # seed 0 is the canonical file; other seeds get a suffix so true seed-replicates | |
| # (same corpus, same script, same curriculum) can be compared side by side | |
| fname = "cv_5fold.json" if seed == 0 else f"cv_5fold_seed{seed}.json" | |
| (out_dir / fname).write_text(json.dumps(result, indent=2, default=str)) | |
| print(f"[cv] mean acc={result['mean_acc']:.4f}±{result['std_acc']:.4f} " | |
| f"F1={result['mean_f1']:.4f} AUC={result['mean_auc']:.4f}", flush=True) | |
| if __name__ == "__main__": | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--systems", nargs="*", default=["pan_skin", "hematopoiesis", "pancreas"]) | |
| ap.add_argument("--variants", nargs="*", default=["pca", "marker"]) | |
| ap.add_argument("--folds", type=int, default=5) | |
| ap.add_argument("--epochs", type=int, default=5) | |
| ap.add_argument("--seed", type=int, default=0, | |
| help="fold-assignment + model-init seed; non-zero seeds write cv_5fold_seed{N}.json") | |
| args = ap.parse_args() | |
| for s in args.systems: | |
| for v in args.variants: | |
| cv(s, v, args.folds, args.epochs, seed=args.seed) | |