PANDA / scripts /pancreas /06_zero_shot_baron.py
bryan7264's picture
Correction pass: gate-matched Dahlin, retracted unsupported claims, complete HF-placode DEG set, restyled figures
141bacd verified
Raw
History Blame Contribute Delete
4.96 kB
"""zero-shot PANDA on held-out Baron 2016 mouse pancreas."""
from __future__ import annotations
from pathlib import Path
import warnings, json, sys, pickle
warnings.filterwarnings("ignore")
import numpy as np, pandas as pd, anndata as ad, scanpy as sc, scipy.sparse as sp
import torch
from pathlib import Path as _P_root
ROOT = _P_root(__file__).resolve().parents[2]
ROOT_STR = str(ROOT)
sys.path.insert(0, ROOT_STR)
from panda import PANDAEncoder
from panda.data.pancreas_loaders import load_baron
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CKPT = Path(f"{ROOT_STR}/checkpoints/pancreas")
HARM = Path(f"{ROOT_STR}/data/corpus/pancreas/harmonized")
OUT = Path(f"{ROOT_STR}/discovery/pancreas/marker")
OUT.mkdir(parents=True, exist_ok=True)
def project(a, shared_hvgs, mu, sig):
G = len(shared_hvgs); n = a.n_obs
hvg2i = {g: i for i, g in enumerate(shared_hvgs)}
common = [g for g in a.var_names.astype(str) if g in hvg2i]
present_frac = len(common) / G
print(f"[proj] {len(common)}/{G} HVGs present ({present_frac:.1%})", 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((n, G), dtype=np.float32)
cols = [hvg2i[g] for g in common]
Xf[:, cols] = X
Xz = np.clip((Xf - mu.astype(np.float32)) / sig.astype(np.float32), -10, 10)
return Xz, present_frac
def main():
ck = torch.load(CKPT / "panda_final.pt", map_location=DEVICE, weights_only=False)
classes = ck["classes"]; datasets = ck["datasets"]
model = PANDAEncoder(n_pca=50, n_classes=len(classes),
n_datasets=len(datasets)).to(DEVICE).eval()
model.load_state_dict(ck["model"])
print(f"[model] classes: {classes}", flush=True)
stats = np.load(HARM / "corpus_stats.npz", allow_pickle=True)
shared_hvgs = [str(g) for g in stats["shared_hvgs"]]
mu, sig = stats["mean"], stats["std"]
with open(HARM / "pca_basis.pkl", "rb") as f: pca = pickle.load(f)
a = load_baron()
print(f"[target] Baron shape: {a.shape}", flush=True)
print(f"[target] labels (assigned_cluster head): "
f"{a.obs.get('assigned_cluster', pd.Series(dtype=str)).value_counts().head(12).to_dict()}",
flush=True)
Xz, present_frac = project(a, shared_hvgs, mu, sig)
counts = np.asarray(a.X.sum(axis=1)).ravel()
logc = np.log10(counts + 1); logc = (logc - logc.mean()) / (logc.std() + 1e-6)
Xpca = pca.transform(Xz).astype(np.float32)
mhf = np.full(a.n_obs, 1.0 - present_frac, dtype=np.float32)
all_z, batch = [], 4096
with torch.no_grad():
for i in range(0, a.n_obs, batch):
xb = torch.from_numpy(Xpca[i:i+batch]).to(DEVICE)
aux = torch.from_numpy(np.stack([mhf[i:i+batch], logc[i:i+batch]], 1).astype(np.float32)).to(DEVICE)
out = model(xb, aux, lam_dann=0.0)
all_z.append(out["z"].cpu().numpy())
Z = np.concatenate(all_z, axis=0)
protos = ck["prototypes"]
protos = protos / (np.linalg.norm(protos, axis=1, keepdims=True) + 1e-8)
cos = Z @ protos.T
pred_ix = cos.argmax(axis=1)
pred_label = np.array([classes[i] for i in pred_ix], dtype=object)
conf = cos.max(axis=1)
a.obs["pred_label"] = pred_label
a.obs["pred_conf"] = conf.astype(np.float32)
a.obs.to_csv(OUT / "baron_predictions.csv")
print(f"\n[summary] Baron pred_label breakdown:")
print(pd.Series(pred_label).value_counts())
if "assigned_cluster" in a.obs.columns:
norm_map = {
"alpha": "alpha", "beta": "beta", "delta": "delta", "gamma": "gamma",
"epsilon": "epsilon", "ductal": "ductal", "acinar": "acinar",
"endothelial": "endothelial", "activated_stellate": "other",
"quiescent_stellate": "other", "schwann": "other",
"mast": "immune", "macrophage": "immune", "t_cell": "immune", "T_cell": "immune",
}
true_lbl = a.obs["assigned_cluster"].astype(str).str.lower().map(norm_map).fillna("other")
mask = true_lbl.isin(classes)
if mask.sum() > 0:
from sklearn.metrics import classification_report, accuracy_score
true = true_lbl[mask].values
pred = pred_label[mask.values]
print(f"\n[eval] shared-class accuracy: {accuracy_score(true, pred):.4f} on {mask.sum()} cells")
print(classification_report(true, pred, digits=3, zero_division=0))
with open(OUT / "baron_accuracy.json", "w") as f:
json.dump({"acc": accuracy_score(true, pred),
"n_evaluated": int(mask.sum())}, f, indent=2)
else:
print("[eval] no shared classes; skipping")
if __name__ == "__main__":
main()