Download scripts/common/zero_shot.py from bryan7264/PANDA: direct link, hf CLI and curl.
- Browser
- Download file 6.88 kB
-
https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/zero_shot.py
- Command line
-
hf download hf://bryan7264/PANDA/scripts/common/zero_shot.py
-
curl -L -o zero_shot.py https://huggingface.co/bryan7264/PANDA/resolve/main/scripts/common/zero_shot.py
6.88 kB
| """zero-shot inference on held-out discovery targets (dingwall / dahlin / veres).""" | |
| from __future__ import annotations | |
| from pathlib import Path | |
| import sys, warnings, pickle, json, argparse, numpy as np, pandas as pd, anndata as ad, scanpy as sc | |
| import scipy.sparse as sp, torch, yaml | |
| warnings.filterwarnings("ignore") | |
| 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 | |
| ROOT = Path(str(PANDA_ROOT)) | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| def prep_input(adata, system, variant, hvgs, mu, sig, pca, marker_genes): | |
| """log-normalise, PCA-50, optional marker channel.""" | |
| hvg2i = {g: i for i, g in enumerate(hvgs)} | |
| common = [g for g in adata.var_names.astype(str) if g in hvg2i] | |
| a_c = adata[:, 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((adata.n_obs, len(hvgs)), dtype=np.float32) | |
| cols = np.array([hvg2i[g] for g in common]) | |
| Xf[:, cols] = X | |
| Xz = np.clip((Xf - mu.astype(np.float32)) / sig.astype(np.float32), -10, 10) | |
| Xpca = pca.transform(Xz).astype(np.float32) | |
| Xmark = None | |
| if variant == "marker": | |
| mvals = np.zeros((adata.n_obs, len(marker_genes)), dtype=np.float32) | |
| for j, g in enumerate(marker_genes): | |
| if g in adata.var_names: | |
| col = adata[:, g].X | |
| if sp.issparse(col): col = col.toarray() | |
| mvals[:, j] = col.flatten().astype(np.float32) | |
| mmu = mvals.mean(axis=0, keepdims=True); msig = mvals.std(axis=0, keepdims=True) + 1e-6 | |
| Xmark = np.clip((mvals - mmu) / msig, -5, 5).astype(np.float32) | |
| return Xpca, Xmark | |
| def infer(system, variant, target_anndata, target_name): | |
| """load checkpoint, run inference, return per-cell (pred, max_cos) + summary.""" | |
| ckpt = torch.load(ROOT / f"checkpoints/{system}/{variant}/panda_final.pt", | |
| map_location=DEVICE, weights_only=False) | |
| classes = ckpt["classes"] | |
| marker_genes = ckpt.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"]] | |
| # case-fold human symbols → mouse-style when hvgs are mouse (e.g. veres cross-species) | |
| a = target_anndata.copy() | |
| n_upper = sum(1 for g in a.var_names[:1000].astype(str) if g.isupper()) | |
| if n_upper > 500: | |
| new = [g[0].upper() + g[1:].lower() if len(g) > 1 else g for g in a.var_names.astype(str)] | |
| a.var_names = new; a.var_names_make_unique() | |
| Xpca, Xmark = prep_input(a, system, variant, hvgs, stats["mean"], stats["std"], pca, marker_genes) | |
| 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(ckpt["datasets"]), | |
| ).to(DEVICE).eval() | |
| model.load_state_dict(ckpt["model"]) | |
| protos = model.prototypes # (K, n_sub, D) | |
| preds, max_cos_list = [], [] | |
| 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) | |
| z = out["z"] | |
| mc = model.max_sub_cos(z) # (B, K) | |
| preds.append(mc.argmax(dim=1).cpu().numpy()) | |
| max_cos_list.append(mc.max(dim=1).values.cpu().numpy()) | |
| preds = np.concatenate(preds); max_cos = np.concatenate(max_cos_list) | |
| pred_labels = np.array([classes[i] for i in preds]) | |
| out_dir = ROOT / f"discovery/{system}/{variant}" | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| df = pd.DataFrame({ | |
| "cell_id": a.obs_names, | |
| "pred_label": pred_labels, | |
| "max_cos": max_cos, | |
| }) | |
| df.to_csv(out_dir / f"{target_name}_predictions.csv", index=False) | |
| dist = pd.Series(pred_labels).value_counts() | |
| summary = { | |
| "system": system, "variant": variant, "target": target_name, | |
| "n_cells": int(a.n_obs), | |
| "n_classes": len(classes), | |
| "predicted_class_dist": dist.to_dict(), | |
| "max_cos_p50": float(np.median(max_cos)), | |
| "max_cos_p05": float(np.quantile(max_cos, 0.05)), | |
| "abstain_frac_cos_lt_0.5": float((max_cos < 0.5).mean()), | |
| } | |
| (out_dir / f"{target_name}_summary.json").write_text(json.dumps(summary, indent=2, default=str)) | |
| print(f"[{system}/{variant}/{target_name}] {a.n_obs} cells, top preds: {dist.head(5).to_dict()}", flush=True) | |
| return summary | |
| def load_target(name): | |
| if name == "dingwall": | |
| return ad.read_h5ad(ROOT / "data/raw/GSE220977_combined.h5ad") | |
| if name == "veres": | |
| SHARON_DIR = ROOT / "data/corpus/pancreas/held_out_unlabeled/sharon_extract" | |
| parts = [] | |
| for meta_file in sorted(SHARON_DIR.glob("*.cell_metadata.tsv.gz")): | |
| counts_file = str(meta_file).replace("cell_metadata", "processed_counts") | |
| if not Path(counts_file).exists(): continue | |
| meta = pd.read_csv(meta_file, sep="\t", compression="gzip") | |
| counts = pd.read_csv(counts_file, sep="\t", compression="gzip", index_col=0) | |
| obs = meta.set_index("library.barcode") | |
| obs = obs.loc[obs.index.intersection(counts.index)] | |
| counts_al = counts.loc[obs.index] | |
| X = sp.csr_matrix(counts_al.values.astype(np.float32)) | |
| a = ad.AnnData(X=X, obs=obs, | |
| var=pd.DataFrame(index=counts_al.columns)) | |
| a.var_names_make_unique() | |
| parts.append(a) | |
| return ad.concat(parts, join="outer") | |
| if name == "dahlin": | |
| # skipped here — needs mygene ENSMUSG→symbol conversion (see run_all_zero_shot) | |
| return None | |
| raise ValueError(name) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("system", choices=["pan_skin", "hematopoiesis", "pancreas"]) | |
| ap.add_argument("--variant", choices=["pca", "marker"], required=True) | |
| ap.add_argument("--target", choices=["dingwall", "dahlin", "veres"], required=True) | |
| args = ap.parse_args() | |
| a = load_target(args.target) | |
| if a is None: | |
| print(f"[!] target {args.target} loader deferred", flush=True); return | |
| infer(args.system, args.variant, a, args.target) | |
| if __name__ == "__main__": | |
| main() | |