"""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()