File size: 6,880 Bytes
141bacd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 | """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()
|