Download scripts/common/run_all_zero_shot.py from bryan7264/PANDA: direct link, hf CLI and curl.
- Browser
- Download file 10.4 kB
-
https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/run_all_zero_shot.py
- Command line
-
hf download hf://bryan7264/PANDA/scripts/common/run_all_zero_shot.py
-
curl -L -o run_all_zero_shot.py https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/run_all_zero_shot.py
10.4 kB
| """zero-shot inference over every held-out target x (pca, marker) checkpoint.""" | |
| from pathlib import Path | |
| import warnings, json, sys, pickle, argparse, numpy as np, pandas as pd, anndata as ad, scanpy as sc, scipy.sparse as sp, torch, torch.nn.functional as F | |
| warnings.filterwarnings("ignore"); sc.settings.verbosity = 0 | |
| 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 | |
| from sklearn.metrics import accuracy_score, f1_score, classification_report | |
| ROOT = Path(str(PANDA_ROOT)) | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| def infer(a, system, variant): | |
| ck = torch.load(ROOT / f"checkpoints/{system}/{variant}/panda_final.pt", | |
| map_location=DEVICE, weights_only=False) | |
| classes = ck["classes"]; marker_genes = ck.get("marker_genes", []) | |
| 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)} | |
| # human->mouse symbol case-fold (same heuristic as zero_shot.py): corpora + markers.yaml | |
| # use mouse Title-case symbols; human targets (e.g. veres) ship ALL-CAPS HGNC symbols. | |
| # Without this the HVG intersection collapses to ~0 and predictions are meaningless. | |
| vn = a.var_names.astype(str) | |
| n_upper = sum(1 for g in vn[:1000] if g.isupper() and len(g) > 1) | |
| if n_upper > 500: | |
| a = a.copy() | |
| a.var_names = [g.capitalize() for g in vn] | |
| a.var_names_make_unique() | |
| print(f"[infer] case-folded {n_upper}/1000 uppercase symbols human->mouse", flush=True) | |
| common = [g for g in a.var_names.astype(str) if g in hvg2i] | |
| if len(common) < 0.2 * len(hvgs): | |
| print(f"[infer] WARNING: only {len(common)}/{len(hvgs)} corpus HVGs present in target; " | |
| f"predictions will be unreliable", flush=True) | |
| 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 | |
| if variant == "marker": | |
| 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) | |
| # prefer the training-corpus marker stats stored in the checkpoint; refitting on the | |
| # target puts the marker channel on a target-dependent scale the model never saw | |
| if ck.get("marker_mu") is not None and ck.get("marker_sig") is not None: | |
| mmu = np.asarray(ck["marker_mu"], dtype=np.float32) | |
| msig = np.asarray(ck["marker_sig"], dtype=np.float32) | |
| else: | |
| print("[infer] WARNING: checkpoint lacks marker_mu/sig; z-scoring markers on the " | |
| "target itself (legacy behaviour, target-dependent scale)", flush=True) | |
| 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) | |
| model = PANDAEncoder(variant=variant, n_pca=50, | |
| n_markers=len(marker_genes) if variant == "marker" else 0, | |
| n_classes=len(classes), n_sub=3, | |
| n_datasets=len(ck["datasets"])).to(DEVICE).eval() | |
| model.load_state_dict(ck["model"]) | |
| preds, probs, coss = [], [], [] | |
| with torch.no_grad(): | |
| for i in range(0, a.n_obs, 4096): | |
| xb = torch.from_numpy(Xpca[i:i+4096]).to(DEVICE) | |
| xmb = torch.from_numpy(Xmark[i:i+4096]).to(DEVICE) if Xmark is not None else None | |
| aux = torch.zeros(len(xb), 2, device=DEVICE) | |
| out = model(xb, aux, x_markers=xmb, lam_dann=0.0) | |
| mc = model.max_sub_cos(out["z"]) | |
| preds.append(mc.argmax(dim=1).cpu().numpy()) | |
| coss.append(mc.max(dim=1).values.cpu().numpy()) | |
| probs.append(F.softmax(mc / 0.07, dim=1).cpu().numpy()) | |
| return (np.array([classes[i] for i in np.concatenate(preds)]), | |
| np.concatenate(probs), np.concatenate(coss), classes) | |
| TARGETS = { | |
| "pan_skin": [ | |
| ("dingwall", ROOT / "data/raw/GSE220977_combined.h5ad", None), | |
| # WARNING: all 4,683 sulic cells (incl. this 4,183-cell "test" slice) are inside | |
| # data/corpus/pan_skin/harmonized/corpus.h5ad (verified by barcode overlap 2026-08-19). | |
| # Scoring the standard corpus checkpoint here is a TRAIN-SET evaluation, not held-out. | |
| # Use scripts/pan_skin/92_retrain_with_sulic_anchor.py (500-cell anchor, rest held out) | |
| # for an honest Sulic number. | |
| ("sulic", ROOT / "data/corpus/pan_skin/held_out_labeled/sulic_GSE212673_test.h5ad", "canonical_label"), | |
| ("belote", ROOT / "data/corpus/pan_skin/held_out_labeled/belote_GSE151091_test.h5ad", "canonical_label"), | |
| ], | |
| "hematopoiesis": [ | |
| ("nestorowa", ROOT / "data/corpus/hematopoiesis/held_out_labeled/nestorowa_GSE81682_test.h5ad", "cell_type"), | |
| ("dahlin", None, None), # loaded per-file via loader (61k cells across 8 samples) | |
| ], | |
| "pancreas": [ | |
| ("baron", ROOT / "data/corpus/pancreas/held_out_labeled/baron_GSE84133_mouse_test.h5ad", "canonical_label"), | |
| ("veres", ROOT / "data/corpus/pancreas/held_out_labeled/veres_GSE114412_test.h5ad", "canonical_label"), | |
| ], | |
| } | |
| def load_dahlin(): | |
| """dahlin 61k held-out unlabeled hsc target, 8 sample files.""" | |
| D = ROOT / "data/corpus/hematopoiesis/held_out_unlabeled/dahlin_extract" | |
| GT = {"SIGAB1":"WT","SIGAC1":"WT","SIGAD1":"WT","SIGAF1":"WT","SIGAG1":"WT", | |
| "SIGAH1":"WT","SIGAG8":"Kit_W41","SIGAH8":"Kit_W41"} | |
| parts = [] | |
| for f in sorted(D.glob("*.txt.gz")): | |
| sample = f.name.split("_")[1].split(".")[0] | |
| df = pd.read_csv(f, sep="\t", compression="gzip", index_col=0) | |
| X = sp.csr_matrix(df.values.T.astype(np.float32)) | |
| obs = pd.DataFrame(index=[f"{sample}_{bc}" for bc in df.columns.astype(str)]) | |
| obs["sample"] = sample; obs["genotype"] = GT.get(sample, "unknown") | |
| var = pd.DataFrame(index=df.index.astype(str)) | |
| parts.append(ad.AnnData(X=X, obs=obs, var=var)) | |
| a = ad.concat(parts, join="outer") | |
| import mygene | |
| mg = mygene.MyGeneInfo() | |
| res = mg.querymany(a.var_names.astype(str).tolist(), scopes="ensembl.gene", | |
| fields="symbol", species="mouse", verbose=False) | |
| id2sym = {r["query"]: r["symbol"] for r in res if "symbol" in r} | |
| syms = pd.Series(a.var_names.astype(str)).map(id2sym).values | |
| keep = pd.notna(syms) | |
| a = a[:, keep].copy(); a.var_names = syms[keep]; a.var_names_make_unique() | |
| return a | |
| def process(system, variant): | |
| print(f"\n===== {system} / {variant} =====", flush=True) | |
| for tgt_name, tgt_path, tgt_label in TARGETS[system]: | |
| print(f"\n[{tgt_name}] loading", flush=True) | |
| if tgt_name == "dahlin": | |
| a = load_dahlin() | |
| else: | |
| a = ad.read_h5ad(tgt_path) | |
| print(f"[{tgt_name}] {a.shape}", flush=True) | |
| pred, probs, max_cos, classes = infer(a, system, variant) | |
| out_dir = ROOT / f"discovery/{system}/{variant}" | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| # max_cos is the genuine prototype cosine; max_prob is softmax(max_cos/0.07). | |
| # (earlier revisions wrote the softmax value under the name max_cos) | |
| pd.DataFrame({ | |
| "cell_id": a.obs_names, | |
| "pred_label": pred, | |
| "max_cos": max_cos, | |
| "max_prob": probs.max(axis=1), | |
| }).to_csv(out_dir / f"{tgt_name}_predictions.csv", index=False) | |
| summary = { | |
| "system": system, "variant": variant, "target": tgt_name, | |
| "n_cells": int(a.n_obs), "n_classes_model": len(classes), | |
| "predicted_class_dist": pd.Series(pred).value_counts().head(30).to_dict(), | |
| "max_cos_p50": float(np.median(max_cos)), | |
| "max_cos_p05": float(np.quantile(max_cos, 0.05)), | |
| "max_prob_p50": float(np.median(probs.max(axis=1))), | |
| "max_prob_p05": float(np.quantile(probs.max(axis=1), 0.05)), | |
| } | |
| if tgt_label and tgt_label in a.obs.columns: | |
| y_true = a.obs[tgt_label].astype(str).values | |
| mask = np.isin(y_true, classes) | |
| if mask.sum() > 0: | |
| acc = accuracy_score(y_true[mask], pred[mask]) | |
| f1 = f1_score(y_true[mask], pred[mask], average="macro", zero_division=0) | |
| rep = classification_report(y_true[mask], pred[mask], | |
| zero_division=0, output_dict=True) | |
| summary["labeled_eval"] = { | |
| "n_eval": int(mask.sum()), "acc": float(acc), | |
| "n_excluded_off_vocab": int((~mask).sum()), | |
| "excluded_label_dist": pd.Series(y_true[~mask]).value_counts().head(20).to_dict(), | |
| "macro_f1": float(f1), "per_class_report": rep, | |
| } | |
| print(f"[{tgt_name}] acc={acc:.4f} F1={f1:.4f} on {mask.sum()} labeled cells", flush=True) | |
| (out_dir / f"{tgt_name}_summary.json").write_text(json.dumps(summary, indent=2, default=str)) | |
| print(f"[{tgt_name}] wrote {out_dir}/{tgt_name}_predictions.csv + summary.json", flush=True) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--systems", nargs="*", default=["pan_skin", "hematopoiesis", "pancreas"]) | |
| ap.add_argument("--variants", nargs="*", default=["pca", "marker"]) | |
| args = ap.parse_args() | |
| for sys_ in args.systems: | |
| for var in args.variants: | |
| process(sys_, var) | |
| if __name__ == "__main__": | |
| main() | |