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