File size: 6,638 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
"""zero-shot HSC PANDA on nestorowa 2016 smart-seq2 as unlabeled discovery target."""
from __future__ import annotations
from pathlib import Path
import warnings, json, sys, pickle
warnings.filterwarnings("ignore")
import numpy as np, pandas as pd, anndata as ad, scanpy as sc, scipy.sparse as sp
import torch
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.model import PANDAEncoder

DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
CKPT = Path(str(PANDA_ROOT / "checkpoints/hematopoiesis"))
HARM = Path(str(PANDA_ROOT / "data/corpus/hematopoiesis/harmonized"))
OUT = Path(str(PANDA_ROOT / "discovery/hematopoiesis/marker"))


def load_nestorowa():
    p = Path(str(PANDA_ROOT / "data/corpus/hematopoiesis/tier_a/nestorowa_GSE81682_counts.txt.gz"))
    df = pd.read_csv(p, sep="\t", index_col=0, compression="gzip")
    # rows = ENSMUSG, cols = cells; convert IDs to symbols via mygene
    import mygene
    print(f"[nestorowa] converting {len(df)} ENSMUSG IDs to symbols via mygene …", flush=True)
    mg = mygene.MyGeneInfo()
    ids = df.index.astype(str).tolist()
    res = mg.querymany(ids, scopes="ensembl.gene", fields="symbol",
                       species="mouse", returnall=False, verbose=False)
    id2sym = {}
    for r in res:
        if "symbol" in r and "query" in r and not r.get("notfound", False):
            id2sym[r["query"]] = r["symbol"]
    syms_series = pd.Series(df.index.astype(str)).map(id2sym).values
    syms = pd.Series(syms_series, index=df.index)
    keep = syms.notna()
    print(f"[nestorowa] mapped {int(keep.sum())}/{len(df)} genes", flush=True)
    df = df.loc[keep]
    df.index = syms[keep].values
    df = df.groupby(df.index).sum()
    X = sp.csr_matrix(df.values.T.astype(np.float32))
    obs = pd.DataFrame(index=df.columns.astype(str))
    obs["dataset"] = "nestorowa_GSE81682"
    obs["sample"] = obs.index.str.split("_").str[0]
    var = pd.DataFrame({"gene_symbol": df.index.astype(str).values},
                       index=df.index.astype(str))
    a = ad.AnnData(X=X, obs=obs, var=var)
    a.var_names_make_unique()
    return a


def project(a, shared_hvgs, mu, sig):
    G = len(shared_hvgs); n = a.n_obs
    hvg2i = {g: i for i, g in enumerate(shared_hvgs)}
    common = [g for g in a.var_names.astype(str) if g in hvg2i]
    frac = len(common) / G
    print(f"[proj] {len(common)}/{G} HVGs present ({frac:.1%})", flush=True)
    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((n, G), dtype=np.float32)
    cols = [hvg2i[g] for g in common]
    Xf[:, cols] = X
    Xz = np.clip((Xf - mu.astype(np.float32)) / sig.astype(np.float32), -10, 10)
    return Xz, frac


def main():
    ck = torch.load(CKPT / "panda_final.pt", map_location=DEVICE, weights_only=False)
    classes = ck["classes"]; datasets = ck["datasets"]
    model = PANDAEncoder(n_pca=50, n_classes=len(classes),
                         n_datasets=len(datasets)).to(DEVICE).eval()
    model.load_state_dict(ck["model"])
    protos = ck["prototypes"]
    protos = protos / (np.linalg.norm(protos, axis=1, keepdims=True) + 1e-8)
    print(f"[model] classes: {classes}", flush=True)

    stats = np.load(HARM / "corpus_stats.npz", allow_pickle=True)
    shared_hvgs = [str(g) for g in stats["shared_hvgs"]]
    mu, sig = stats["mean"], stats["std"]
    with open(HARM / "pca_basis.pkl", "rb") as f: pca = pickle.load(f)

    a = load_nestorowa()
    print(f"[target] Nestorowa shape: {a.shape}", flush=True)

    Xz, frac = project(a, shared_hvgs, mu, sig)
    Xpca = pca.transform(Xz).astype(np.float32)

    all_z = []
    with torch.no_grad():
        for i in range(0, a.n_obs, 4096):
            xb = torch.from_numpy(Xpca[i:i+4096]).to(DEVICE)
            aux = torch.zeros(len(xb), 2, device=DEVICE)
            out = model(xb, aux, lam_dann=0.0)
            all_z.append(out["z"].cpu().numpy())
    Z = np.concatenate(all_z, axis=0)
    cos = Z @ protos.T
    pred_ix = cos.argmax(axis=1)
    conf = cos.max(axis=1)
    entropy = -(np.exp(cos / 0.07) / np.exp(cos / 0.07).sum(axis=1, keepdims=True) *
                np.log(np.exp(cos / 0.07) / np.exp(cos / 0.07).sum(axis=1, keepdims=True) + 1e-12)
                ).sum(axis=1)

    a.obs["pred_label"] = np.array([classes[i] for i in pred_ix], dtype=object)
    a.obs["pred_conf"] = conf.astype(np.float32)
    a.obs["pred_entropy"] = entropy.astype(np.float32)

    print(f"\n[nestorowa] predicted class distribution:")
    print(a.obs["pred_label"].value_counts())
    print(f"\n[nestorowa] pred_conf: p10={np.percentile(conf,10):.3f}, "
          f"p50={np.percentile(conf,50):.3f}, p90={np.percentile(conf,90):.3f}")
    print(f"[nestorowa] shared-HVG fraction: {frac:.1%}")

    a.obs.to_csv(OUT / "63_nestorowa_predictions.csv")
    print(f"\n[nestorowa] wrote predictions to 63_nestorowa_predictions.csv")

    # cluster bottom-decile confidence cells for novel-population DE
    thr = np.percentile(conf, 10)
    mask = conf <= thr
    print(f"\n[nestorowa] bottom-decile confidence: {int(mask.sum())} cells (thr={thr:.3f})",
          flush=True)
    if mask.sum() >= 30:
        sub_low = ad.AnnData(X=Z[mask].astype(np.float32))
        sc.pp.neighbors(sub_low, use_rep="X", n_neighbors=10)
        sc.tl.leiden(sub_low, resolution=0.5, key_added="cluster")
        print(f"[nestorowa] novel clusters: {sub_low.obs['cluster'].nunique()}")
        a_low = a[mask].copy()
        a_low.obs["cluster"] = sub_low.obs["cluster"].values
        try:
            sc.tl.rank_genes_groups(a_low, "cluster", method="wilcoxon", n_genes=10, use_raw=False)
            rows = []
            for cl in sorted(a_low.obs["cluster"].unique()):
                names = a_low.uns["rank_genes_groups"]["names"][cl]
                lfc = a_low.uns["rank_genes_groups"]["logfoldchanges"][cl]
                for g, l in zip(names[:8], lfc[:8]):
                    rows.append({"cluster": cl, "gene": g, "logfc": round(float(l), 3)})
            pd.DataFrame(rows).to_csv(OUT / "63_nestorowa_novel_markers.csv", index=False)
            print(f"[nestorowa] novel markers saved")
        except Exception as e:
            print(f"[nestorowa] DE failed: {e}")


if __name__ == "__main__":
    main()