File size: 5,512 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
"""true zero-shot on Nestorowa GSE81682 (1920 smart-seq2 FACS-labeled 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")
NEST = ROOT / "data/raw/nestorowa_combined.h5ad"

# panda fine class -> nestorowa coarse FACS gate (LT-HSC vs HSPC)
COARSE = {
    "LT-HSC":       "LT-HSC",
    "MPP":          "HSPC",
    "GMP":          "HSPC",
    "myeloid":      "HSPC",
    "erythroid":    "HSPC",
    "megakaryocyte":"HSPC",
    "basophil-mast":"HSPC",
    "lymphoid":     "HSPC",
    "unassigned":   "HSPC",
    "UNK":          "HSPC",
}


def infer(a, variant):
    ck = torch.load(ROOT / f"checkpoints/hematopoiesis/{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/hematopoiesis/harmonized/corpus_stats.npz", allow_pickle=True)
    pca = pickle.load(open(ROOT / "data/corpus/hematopoiesis/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():
    a = ad.read_h5ad(NEST)
    print(f"[nest] {a.shape} facs gates: {a.obs['cell_type'].value_counts().to_dict()}", flush=True)

    for variant in ("pca", "marker"):
        print(f"\n=== {variant.upper()} ===", flush=True)
        pred, probs, classes = infer(a, variant)
        pred_coarse = np.array([COARSE.get(p, "HSPC") for p in pred])
        y_true = a.obs["cell_type"].astype(str).values
        mask = y_true != "unknown"
        acc = accuracy_score(y_true[mask], pred_coarse[mask])
        f1 = f1_score(y_true[mask], pred_coarse[mask], average="macro", zero_division=0)
        rep = classification_report(y_true[mask], pred_coarse[mask], zero_division=0, output_dict=True)
        print(f"[eval-{variant}] n_labeled={mask.sum()} coarse-acc={acc:.4f} macro-f1={f1:.4f}", flush=True)
        fine_by_gate = pd.crosstab(a.obs["cell_type"].astype(str), pd.Series(pred))
        print(fine_by_gate.to_string(), flush=True)

        out = ROOT / f"discovery/hematopoiesis/{variant}"
        out.mkdir(parents=True, exist_ok=True)
        (out / "94_nestorowa_zero_shot.json").write_text(json.dumps({
            "variant": variant, "n_cells_total": int(a.n_obs), "n_cells_labeled": int(mask.sum()),
            "coarse_acc": float(acc), "coarse_f1": float(f1),
            "coarse_per_class": rep,
            "fine_by_gate": fine_by_gate.to_dict(),
            "max_cos_p50": float(np.median(probs.max(axis=1))),
        }, indent=2, default=str))
        pd.DataFrame({"cell_id": a.obs_names, "facs_gate": y_true,
                      "pred_fine": pred, "pred_coarse": pred_coarse,
                      "max_cos": probs.max(axis=1)}).to_csv(out / "94_nestorowa_predictions.csv", index=False)


if __name__ == "__main__":
    main()