"""true zero-shot on baron test-half (943 mouse islet cells held out of corpus) under pca+marker variants.""" from pathlib import Path import warnings, json, sys, pickle, numpy as np, pandas as pd, anndata as ad, scanpy as sc, scipy.sparse as sp, torch 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") BARON = ROOT / "data/corpus/pancreas/held_out_labeled/baron_GSE84133_mouse_test.h5ad" def infer(a, variant): ck = torch.load(ROOT / f"checkpoints/pancreas/{variant}/panda_final.pt", map_location=DEVICE, weights_only=False) classes = ck["classes"]; marker_genes = ck.get("marker_genes", []) stats = np.load(ROOT / "data/corpus/pancreas/harmonized/corpus_stats.npz", allow_pickle=True) pca = pickle.load(open(ROOT / "data/corpus/pancreas/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 if variant == "marker": mvals = 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() 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) 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 = [], [] 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()) probs.append(torch.softmax(mc / 0.07, dim=1).cpu().numpy()) return np.array([classes[i] for i in np.concatenate(preds)]), np.concatenate(probs), classes def main(): print(f"[baron] loading {BARON}", flush=True) a = ad.read_h5ad(BARON) y_true = a.obs["canonical_label"].astype(str).values print(f"[baron] {a.shape} true labels: {pd.Series(y_true).value_counts().to_dict()}", flush=True) for variant in ("pca", "marker"): print(f"\n=== {variant.upper()} ===", flush=True) pred, probs, classes = infer(a, variant) # eval only on cells whose true label is in our class vocabulary mask = np.isin(y_true, classes) 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) print(f"[eval-{variant}] n={mask.sum()} acc={acc:.4f} macro-f1={f1:.4f}", flush=True) out = ROOT / f"discovery/pancreas/{variant}" out.mkdir(parents=True, exist_ok=True) (out / "93_baron_zero_shot.json").write_text(json.dumps({ "variant": variant, "n_cells": int(mask.sum()), "n_classes_eval": int(len(set(y_true[mask]))), "acc": float(acc), "macro_f1": float(f1), "per_class": rep, "predicted_dist": pd.Series(pred).value_counts().to_dict(), "true_dist": pd.Series(y_true).value_counts().to_dict(), }, indent=2, default=str)) pd.DataFrame({"cell_id": a.obs_names, "true_label": y_true, "pred_label": pred, "max_cos": probs.max(axis=1)}).to_csv( out / "93_baron_predictions.csv", index=False) if __name__ == "__main__": main()