File size: 3,042 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
"""cluster the top-500 attributed genes by cross-class attribution correlation."""
from __future__ import annotations
from pathlib import Path
import warnings, json, pickle, sys, numpy as np, pandas as pd
warnings.filterwarnings("ignore")
from scipy.cluster.hierarchy import linkage, fcluster
from scipy.spatial.distance import squareform

import os as _os
from pathlib import Path as _Path
PANDA_ROOT = _Path(_os.environ.get("PANDA_ROOT", str(_Path(__file__).resolve().parents[2])))
CORP_ROOT = Path(str(PANDA_ROOT / "data/corpus"))
OUT = Path(str(PANDA_ROOT / "discovery")); OUT.mkdir(exist_ok=True)
CKPT_ROOT = Path(str(PANDA_ROOT / "checkpoints"))

TOP_GENES = 500  # keep top-500 by |att| summed across classes
N_MODULES = 15


def main(sys):
    print(f"\n===== {sys} =====", flush=True)
    stats = np.load(CORP_ROOT / sys / "harmonized/corpus_stats.npz", allow_pickle=True)
    hvgs = [str(g) for g in stats["shared_hvgs"]]
    import torch
    ck = torch.load(CKPT_ROOT / sys / "marker" / "panda_final.pt", map_location="cpu", weights_only=False)
    classes = ck["classes"]

    A = np.load(OUT / f"80_{sys}_gene_attribution_full.npy")  # (K, G)
    print(f"[load] A shape {A.shape}, classes={classes}", flush=True)

    gene_score = np.abs(A).sum(axis=0)
    top_idx = np.argsort(-gene_score)[:TOP_GENES]
    A_top = A[:, top_idx]
    top_genes = [hvgs[i] for i in top_idx]

    Xn = A_top - A_top.mean(axis=0, keepdims=True)
    Xn = Xn / (Xn.std(axis=0, keepdims=True) + 1e-8)
    C = np.corrcoef(Xn.T)
    print(f"[corr] gene-gene C shape {C.shape}, diag mean={np.diag(C).mean():.3f}", flush=True)

    D = 1 - C
    np.fill_diagonal(D, 0)
    D = np.clip(D, 0, 2)
    Z = linkage(squareform(D, checks=False), method="average")
    labels = fcluster(Z, t=N_MODULES, criterion="maxclust")

    rows = []
    for mod in sorted(set(labels)):
        members = np.where(labels == mod)[0]
        if len(members) < 3: continue
        member_genes = [top_genes[i] for i in members]
        mod_att = A_top[:, members].mean(axis=1)
        dom_ci = int(np.argmax(mod_att))
        rows.append({
            "module_id": int(mod),
            "size": int(len(members)),
            "dominant_class": classes[dom_ci],
            "dom_class_mean_att": float(mod_att[dom_ci]),
            "member_genes": ",".join(member_genes[:30]),
            "n_shown": min(30, len(member_genes)),
        })

    df = pd.DataFrame(rows).sort_values(["dominant_class", "dom_class_mean_att"], ascending=[True, False])
    df.to_csv(OUT / f"82_{sys}_coatt_modules.csv", index=False)
    print(f"[wrote] {len(rows)} modules to {OUT}/82_{sys}_coatt_modules.csv", flush=True)
    print(df.head(15).to_string(index=False)[:2000], flush=True)


for sys in ["pan_skin", "hematopoiesis", "pancreas"]:
    try:
        main(sys)
    except Exception as e:
        import traceback; traceback.print_exc()
        print(f"[!] {sys}: {e}", flush=True)
print("\n=== DONE ===", flush=True)