PANDA / scripts /common /zero_shot.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
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()