File size: 5,376 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
"""PANDA lineage distribution across LARRY time points d2/d9/d16 + clonal purity + per-lineage DE."""
from __future__ import annotations
from pathlib import Path
import warnings, json, sys
warnings.filterwarnings("ignore")
import numpy as np, pandas as pd, anndata as ad, scanpy as sc, torch
from scipy import stats

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"))
CORPUS = Path(str(PANDA_ROOT / "data/corpus/hematopoiesis/harmonized/corpus.h5ad"))
OUT = Path(str(PANDA_ROOT / "discovery/hematopoiesis/marker"))
OUT.mkdir(parents=True, exist_ok=True)


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)

    # predict on all corpus cells (this is training-set prediction — for analysis only)
    a = ad.read_h5ad(CORPUS)
    print(f"[hsc-mech] corpus: {a.shape}", flush=True)
    X = np.asarray(a.obsm["X_pca"]).astype(np.float32)

    all_z = []
    with torch.no_grad():
        for i in range(0, X.shape[0], 8192):
            xb = torch.from_numpy(X[i:i+8192]).to(DEVICE)
            aux = torch.zeros(len(xb), 2, device=DEVICE)
            all_z.append(model(xb, aux, lam_dann=0.0)["z"].cpu().numpy())
    Z = np.concatenate(all_z, axis=0)
    cos = Z @ protos.T
    pred_ix = cos.argmax(axis=1)
    pred = np.array([classes[i] for i in pred_ix], dtype=object)
    a.obs["pred_label"] = pred
    a.obs["pred_conf"] = cos.max(axis=1)

    if "Time point" not in a.obs.columns:
        print("[hsc-mech] no Time point column; skipping time analysis")
    else:
        tp = a.obs["Time point"].astype(int)
        xt = pd.crosstab(pred, tp, normalize="columns")
        print("\n[hsc-mech] Fraction per predicted class per time point:")
        print(xt.round(3))
        xt.to_csv(OUT / "62_time_course_class_fractions.csv")

        rows = []
        n_d2 = int((tp == 2).sum()); n_d16 = int((tp == 16).sum())
        for c in classes:
            n_c_d16 = int(((pred == c) & (tp == 16)).sum())
            n_c_d2  = int(((pred == c) & (tp == 2)).sum())
            contingency = np.array([[n_c_d16, n_d16 - n_c_d16], [n_c_d2, n_d2 - n_c_d2]])
            odds, p = stats.fisher_exact(contingency)
            f16 = (n_c_d16 + 1) / (n_d16 + 2); f2 = (n_c_d2 + 1) / (n_d2 + 2)
            rows.append({"class": c, "n_d16": n_c_d16, "n_d2": n_c_d2,
                         "log2_fold_d16_vs_d2": round(np.log2(f16 / f2), 3),
                         "fisher_p": p})
        df = pd.DataFrame(rows).sort_values("log2_fold_d16_vs_d2", ascending=False)
        print("\n[hsc-mech] class enrichment d16 vs d2 (positive = expanded at late time):")
        print(df.to_string(index=False))
        df.to_csv(OUT / "62_time_course_enrichment.csv", index=False)

    # sibling-fate concordance: Library = clonal barcode
    if "Library" in a.obs.columns:
        libs = a.obs["Library"].astype(str)
        top_lib = libs.value_counts().head(200).index  # top 200 largest clones
        clone_purity = []
        for L in top_lib:
            m = libs == L
            if m.sum() < 3: continue
            pl = pd.Series(pred[m.values]).value_counts(normalize=True)
            clone_purity.append({
                "library": L, "n": int(m.sum()),
                "dominant_class": pl.index[0],
                "purity": float(pl.iloc[0]),
            })
        cp = pd.DataFrame(clone_purity)
        print(f"\n[hsc-mech] clonal purity (dominant-class fraction) — {len(cp)} clones:")
        print(f"  median: {cp['purity'].median():.3f}, mean: {cp['purity'].mean():.3f}, "
              f"n_clones_pure_>0.9: {(cp['purity'] > 0.9).sum()}/{len(cp)}")
        cp.to_csv(OUT / "62_clonal_purity.csv", index=False)

    a.obs["pred_label"] = pd.Categorical(pred)
    keep_classes = [c for c in classes if (pred == c).sum() >= 100]
    a_sub = a[np.isin(pred, keep_classes)].copy()
    if a_sub.n_obs >= 500:
        sc.tl.rank_genes_groups(a_sub, "pred_label", method="wilcoxon",
                                n_genes=30, use_raw=False)
        rows = []
        for cls in keep_classes:
            try:
                names = a_sub.uns["rank_genes_groups"]["names"][cls]
                lfc = a_sub.uns["rank_genes_groups"]["logfoldchanges"][cls]
                for g, l in zip(names[:15], lfc[:15]):
                    rows.append({"class": cls, "gene": g, "logfc": round(float(l), 3)})
            except Exception: pass
        pd.DataFrame(rows).to_csv(OUT / "62_lineage_markers.csv", index=False)
        print(f"\n[hsc-mech] wrote lineage markers to 62_lineage_markers.csv")

    print(f"\n[hsc-mech] complete. Outputs in {OUT}/")


if __name__ == "__main__":
    main()