File size: 8,284 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
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
"""labeled zero-shot on nestorowa GSE81682 (hsc validation target)."""
from __future__ import annotations
from pathlib import Path
import sys, warnings, json, argparse, pickle, numpy as np, pandas as pd, anndata as ad, scanpy as sc, scipy.sparse as sp, torch
from sklearn.metrics import accuracy_score, f1_score, classification_report
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 load_nestorowa():
    """load nestorowa GSE81682 htseq counts, map ENSMUSG -> symbol."""
    counts_path = ROOT / "data/raw/GSE81682_HTSeq_counts.txt.gz"
    if not counts_path.exists():
        raise FileNotFoundError(counts_path)
    df = pd.read_csv(counts_path, sep="\t", index_col=0)
    print(f"[nestorowa] raw counts shape: {df.shape}", flush=True)
    # rows=genes, cols=cells; transpose
    if df.shape[0] > df.shape[1]:
        df = df.T
    a = ad.AnnData(X=sp.csr_matrix(df.values.astype(np.float32)),
                   obs=pd.DataFrame(index=df.index.astype(str)),
                   var=pd.DataFrame(index=df.columns.astype(str)))
    if any(g.startswith("ENSMUSG") for g in a.var_names[:100]):
        import mygene
        mg = mygene.MyGeneInfo()
        res = mg.querymany(a.var_names.astype(str).tolist(), scopes="ensembl.gene",
                           fields="symbol", species="mouse", verbose=False)
        id2sym = {r["query"]: r["symbol"] for r in res if "symbol" in r}
        syms = pd.Series(a.var_names.astype(str)).map(id2sym).values
        keep = pd.notna(syms)
        a = a[:, keep].copy(); a.var_names = syms[keep]; a.var_names_make_unique()
    print(f"[nestorowa] {a.shape} after gene symbol conversion", flush=True)
    return a


def score_hsc_labels(a):
    """assign hsc labels by marker scoring, proxy for population_annotation."""
    programs = {
        "LT-HSC":       ["Hlf", "Meis1", "Mecom", "Procr", "Fgd5", "Mllt3"],
        "MPP":          ["Cd48", "Flt3", "Cd34"],
        "LMPP":         ["Flt3", "Irf8", "Satb1"],
        "CMP":          ["Cd34", "Mpo", "Gata2"],
        "MEP":          ["Gata1", "Klf1", "Itga2b"],
        "GMP":          ["Elane", "Mpo", "Prtn3", "Ctsg", "Cebpe"],
        "erythroblast": ["Klf1", "Car1", "Car2", "Blvrb", "Hba-a1"],
        "megakaryocyte":["Itga2b", "Pf4", "Gp1bb"],
        "basophil-mast":["Cpa3", "Ms4a2", "Gata2"],
        "CLP":          ["Il7r", "Rag1", "Dntt"],
    }
    sc.pp.normalize_total(a, target_sum=1e4); sc.pp.log1p(a)
    score_cols = []
    for cls, gs in programs.items():
        present = [g for g in gs if g in a.var_names]
        if present:
            sc.tl.score_genes(a, gene_list=present, score_name=f"s_{cls}", use_raw=False)
        else:
            a.obs[f"s_{cls}"] = 0.0
        score_cols.append(f"s_{cls}")
    scores = a.obs[score_cols].values
    argmax = np.argmax(scores, axis=1)
    labels = [c.replace("s_", "") for c in score_cols]
    a.obs["approx_label"] = np.array(labels)[argmax]
    a.obs["approx_conf"] = scores.max(axis=1)
    return a


def infer(a, variant):
    ckpt = torch.load(ROOT / f"checkpoints/hematopoiesis/{variant}/panda_final.pt",
                      map_location=DEVICE, weights_only=False)
    classes = ckpt["classes"]
    marker_genes = ckpt.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()
    # score_hsc_labels already log-normed; re-check in case it was skipped
    if a_c.X.max() > 20:
        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(ckpt["datasets"])).to(DEVICE).eval()
    model.load_state_dict(ckpt["model"])
    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)
            mc = model.max_sub_cos(out["z"])
            preds.append(mc.argmax(dim=1).cpu().numpy())
            max_cos_list.append(mc.max(dim=1).values.cpu().numpy())
    return np.array([classes[i] for i in np.concatenate(preds)]), np.concatenate(max_cos_list), classes


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--variant", choices=["pca", "marker"], required=True)
    args = ap.parse_args()
    print(f"=== Nestorowa HSC zero-shot ({args.variant}) ===", flush=True)
    a = load_nestorowa()
    a = score_hsc_labels(a)
    pred, max_cos, classes = infer(a, args.variant)
    a.obs["pred_label"] = pred
    a.obs["max_cos"] = max_cos

    y_true = a.obs["approx_label"].values
    y_pred = pred
    # only eval on cells with approx-label confidence > 0.05
    conf_mask = a.obs["approx_conf"] > 0.05
    print(f"[eval] eval on {conf_mask.sum()}/{a.n_obs} cells with approx-label conf>0.05", flush=True)
    if conf_mask.sum() > 20:
        common_lbl = sorted(set(y_true[conf_mask]) & set(y_pred[conf_mask]))
        mask2 = conf_mask & np.isin(y_true, common_lbl) & np.isin(y_pred, common_lbl)
        acc = accuracy_score(y_true[mask2], y_pred[mask2])
        f1 = f1_score(y_true[mask2], y_pred[mask2], average="macro", zero_division=0)
        rep = classification_report(y_true[mask2], y_pred[mask2], zero_division=0, output_dict=True)
    else:
        acc = f1 = float("nan"); rep = {}
    result = {
        "variant": args.variant, "n_cells": int(a.n_obs), "n_classes": len(classes),
        "predicted_dist": pd.Series(pred).value_counts().to_dict(),
        "approx_label_dist": pd.Series(y_true).value_counts().to_dict(),
        "eval_acc_vs_approx": float(acc), "eval_f1_vs_approx": float(f1),
        "max_cos_median": float(np.median(max_cos)),
        "n_low_conf_abstain": int((max_cos < 0.5).sum()),
        "per_class_report": rep,
    }
    out_dir = ROOT / f"discovery/hematopoiesis/{args.variant}"
    out_dir.mkdir(parents=True, exist_ok=True)
    (out_dir / "nestorowa_summary.json").write_text(json.dumps(result, indent=2, default=str))
    pd.DataFrame({
        "cell_id": a.obs_names,
        "pred_label": pred, "approx_label": y_true, "approx_conf": a.obs["approx_conf"].values,
        "max_cos": max_cos,
    }).to_csv(out_dir / "nestorowa_predictions.csv", index=False)
    print(f"[write] {out_dir}/nestorowa_*", flush=True)
    print(f"acc_vs_approx={acc:.4f}  f1={f1:.4f}", flush=True)


if __name__ == "__main__":
    main()