File size: 7,498 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
"""dahlin marker deep-dive: wilcoxon per predicted class + Kit-W41 vs WT enrichment."""
from pathlib import Path
import warnings, numpy as np, pandas as pd, anndata as ad, scanpy as sc, scipy.sparse as sp, torch, pickle
warnings.filterwarnings("ignore"); sc.settings.verbosity = 0
import sys
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")
OUT = ROOT / "discovery/hematopoiesis/marker"
OUT.mkdir(parents=True, exist_ok=True)

PANELS = {
    "LT-HSC":         ["Hlf", "Meis1", "Mecom", "Procr", "Fgd5", "Mllt3", "Kit"],
    "MPP":            ["Cd48", "Flt3", "Cd34", "Sell", "Slamf1"],
    "erythroid":      ["Klf1", "Car1", "Car2", "Blvrb", "Hba-a1", "Hba-a2", "Kit"],
    "megakaryocyte":  ["Itga2b", "Pf4", "Gp1bb", "Gata1"],
    "myeloid":        ["Elane", "Mpo", "Prtn3", "Ctsg", "Cebpe", "Wfdc17", "Mmp8", "Ctss"],
    "basophil-mast":  ["Cpa3", "Ms4a2", "Gata2", "Mcpt8", "Hdc"],
    "lymphoid":       ["Il7r", "Rag1", "Dntt", "Vpreb1"],
    "Kit-signaling":  ["Kit", "Kitl", "Sox4"],
    "MYC-targets":    ["Myc", "Nolc1", "Nop58"],
    "ISR":            ["Atf4", "Ddit3", "Ppp1r15a"],
    "Apoptosis-pro":  ["Bax", "Bak1", "Bid"],
}


def load_dahlin():
    from pathlib import Path as _P
    D_DIR = _P(str(PANDA_ROOT / "data/corpus/hematopoiesis/held_out_unlabeled/dahlin_extract"))
    GT = {"SIGAB1":"WT","SIGAC1":"WT","SIGAD1":"WT","SIGAF1":"WT","SIGAG1":"WT",
          "SIGAH1":"WT","SIGAG8":"Kit_W41","SIGAH8":"Kit_W41"}
    parts = []
    for f in sorted(D_DIR.glob("*.txt.gz")):
        sample = f.name.split("_")[1].split(".")[0]
        df = pd.read_csv(f, sep="\t", compression="gzip", index_col=0)
        X = sp.csr_matrix(df.values.T.astype(np.float32))
        obs = pd.DataFrame(index=[f"{sample}_{bc}" for bc in df.columns.astype(str)])
        obs["sample"] = sample; obs["genotype"] = GT.get(sample, "unknown")
        var = pd.DataFrame(index=df.index.astype(str))
        parts.append(ad.AnnData(X=X, obs=obs, var=var))
    a = ad.concat(parts, join="outer", label="_batch")
    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()
    return a


def project_dahlin(a):
    ck = torch.load(ROOT / "checkpoints/hematopoiesis/marker/panda_final.pt",
                    map_location=DEVICE, weights_only=False)
    classes = ck["classes"]; marker_genes = ck["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)

    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="marker", n_pca=50, n_markers=len(marker_genes),
                         n_classes=len(classes), n_sub=3, n_datasets=len(ck["datasets"])).to(DEVICE).eval()
    model.load_state_dict(ck["model"])

    preds = []
    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)
            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())
    preds = np.concatenate(preds)
    return np.array([classes[i] for i in preds])


print("[load] Dahlin + predict", flush=True)
raw = load_dahlin()
raw.obs["pred_label"] = pd.Categorical(project_dahlin(raw))
print(f"[align] {raw.n_obs} cells across {raw.obs['pred_label'].nunique()} classes", flush=True)

counts_s = raw.obs["pred_label"].value_counts()
keep_cls = counts_s[counts_s >= 50].index.tolist()
raw = raw[raw.obs["pred_label"].isin(keep_cls)].copy()
raw.obs["pred_label"] = raw.obs["pred_label"].astype(str).astype("category")
print(f"[filter] {raw.n_obs} cells × {len(keep_cls)} classes", flush=True)

sc.pp.normalize_total(raw, target_sum=1e4); sc.pp.log1p(raw)
print("[wilcoxon] running...", flush=True)
sc.tl.rank_genes_groups(raw, groupby="pred_label", method="wilcoxon", n_genes=25, use_raw=False)

rows = []
for cls in raw.uns["rank_genes_groups"]["names"].dtype.names:
    mask = raw.obs["pred_label"] == cls
    if mask.sum() < 50: continue
    genes = list(raw.uns["rank_genes_groups"]["names"][cls][:20])
    pvals = [float(x) for x in raw.uns["rank_genes_groups"]["pvals_adj"][cls][:20]]
    logfc = [float(x) for x in raw.uns["rank_genes_groups"]["logfoldchanges"][cls][:20]]

    top_str = ",".join([f"{g}(LFC{lf:+.1f})" for g, lf in zip(genes[:10], logfc[:10])])
    panel_hits = {}
    for pname, plist in PANELS.items():
        hits = [g for g in plist if g in genes[:20]]
        panel_hits[pname] = f"{len(hits)}/{len(plist)}: {','.join(hits)}"
    best_panel = max(panel_hits.items(),
                     key=lambda x: int(x[1].split("/")[0]) / (int(x[1].split(":")[0].split("/")[1]) + 1e-6))

    gt = raw.obs["genotype"][mask].astype(str)
    nwt = int((gt == "WT").sum()); nkit = int((gt == "Kit_W41").sum())
    frac_wt = nwt / max(1, nwt + nkit)

    rows.append({
        "predicted_class": cls,
        "n_cells": int(mask.sum()),
        "top_wilcoxon_markers": top_str,
        "min_p_adj_top5": min(pvals[:5], default=float("nan")),
        "best_canonical_panel_match": best_panel[0],
        "recovery": best_panel[1],
        "n_WT": nwt,
        "n_Kit_W41": nkit,
        "frac_WT": frac_wt,
    })

df = pd.DataFrame(rows).sort_values("n_cells", ascending=False)
df.to_csv(OUT / "92_dahlin_marker_deep_dive.csv", index=False)
print(f"[write] {OUT}/92_dahlin_marker_deep_dive.csv ({len(df)} classes)", flush=True)
print()
print(df[["predicted_class", "n_cells", "best_canonical_panel_match", "recovery", "frac_WT"]].to_string(index=False))