File size: 8,954 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
"""5-fold stratified CV, 3 systems x 2 variants. 5 epochs/fold (shorter than train_panda)."""
from __future__ import annotations
import argparse, sys, json, pickle, warnings, numpy as np, pandas as pd, torch, torch.nn.functional as F
from pathlib import Path
import anndata as ad, scanpy as sc, scipy.sparse as sp, yaml
warnings.filterwarnings("ignore"); sc.settings.verbosity = 0
from sklearn.model_selection import StratifiedKFold
from sklearn.metrics import accuracy_score, f1_score, roc_auc_score, classification_report

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, supcon_loss, vicreg_loss, hsic_biased,
                   subcenter_angular_infonce, prototype_repulsion)

ROOT = Path(str(PANDA_ROOT))
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")


def load_and_prepare(system, variant):
    """canonical corpus + features (PCA, optional marker channel)."""
    a = ad.read_h5ad(ROOT / f"data/corpus/{system}/harmonized/corpus.h5ad")
    stats = np.load(ROOT / f"data/corpus/{system}/harmonized/corpus_stats.npz", allow_pickle=True)
    pca = pickle.load(open(ROOT / f"data/corpus/{system}/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)

    Xmark = None; marker_genes = []
    if variant == "marker":
        marker_genes = yaml.safe_load(open(ROOT / "panda/markers.yaml"))[system]
        mv = 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()
                mv[:, j] = col.flatten().astype(np.float32)
        mmu = mv.mean(axis=0, keepdims=True); msig = mv.std(axis=0, keepdims=True) + 1e-6
        Xmark = np.clip((mv - mmu) / msig, -5, 5).astype(np.float32)

    labels = a.obs["canonical_label"].astype(str).values
    classes = sorted(set(labels))
    y = np.array([classes.index(l) for l in labels], dtype=np.int64)
    datasets = sorted(set(a.obs["dataset"].astype(str).values))
    y_dset = np.array([datasets.index(d) for d in a.obs["dataset"].astype(str).values], dtype=np.int64)
    return Xpca, Xmark, y, classes, y_dset, datasets, marker_genes


def train_fold(Xpca, Xmark, y, y_dset, classes, variant, tr_ix, epochs=5, batch=256, lr=1e-3, seed=0):
    n_classes = len(classes)
    n_datasets = int(max(y_dset[tr_ix].max() + 1, 1))
    n_markers = Xmark.shape[1] if Xmark is not None else 0
    torch.manual_seed(seed); np.random.seed(seed)
    model = PANDAEncoder(variant=variant, n_pca=50, n_markers=n_markers,
                         n_classes=n_classes, n_sub=3, n_datasets=n_datasets, dropout=0.2).to(DEVICE)
    opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
    Xt = Xpca[tr_ix]; yt = y[tr_ix]; ydt = y_dset[tr_ix]
    Xmt = Xmark[tr_ix] if Xmark is not None else None
    rng = np.random.default_rng(seed)
    n = len(tr_ix)
    for epoch in range(epochs):
        stage = 0 if epoch < 1 else 1 if epoch < 3 else 2
        perm = rng.permutation(n)
        for bstart in range(0, n, batch):
            idx = perm[bstart:bstart+batch]
            x = torch.from_numpy(Xt[idx]).to(DEVICE)
            xm = torch.from_numpy(Xmt[idx]).to(DEVICE) if Xmt is not None else None
            yy = torch.from_numpy(yt[idx]).to(DEVICE)
            yd = torch.from_numpy(ydt[idx]).to(DEVICE)
            aux = torch.zeros(len(idx), 2, device=DEVICE)
            lam = 0.1 if stage >= 2 else 0.0
            out = model(x, aux, x_markers=xm, lam_dann=lam)
            z = out["z"]
            L = supcon_loss(z, yy, 0.1) + 1.0 * vicreg_loss(z) + 0.4 * F.cross_entropy(out["logits"], yy)
            if stage >= 1:
                L = L + 0.6 * subcenter_angular_infonce(z, yy, model.prototypes.detach().clone(),
                                                        margin=0.15, temperature=0.07)
            if stage >= 2:
                L = L + F.cross_entropy(out["dom"], yd)
            opt.zero_grad(); L.backward(); opt.step()
            if stage >= 1:
                with torch.no_grad(): model.update_prototypes(z.detach(), yy)
    return model


def evaluate(model, Xpca, Xmark, y, val_ix, classes):
    model.eval()
    preds, probs = [], []
    Xv = Xpca[val_ix]; Xmv = Xmark[val_ix] if Xmark is not None else None
    with torch.no_grad():
        for i in range(0, len(val_ix), 2048):
            xb = torch.from_numpy(Xv[i:i+2048]).to(DEVICE)
            xmb = torch.from_numpy(Xmv[i:i+2048]).to(DEVICE) if Xmv is not None else None
            aux = torch.zeros(len(xb), 2, device=DEVICE)
            out = model(xb, aux, x_markers=xmb, lam_dann=0.0)
            z = out["z"]
            mc = model.max_sub_cos(z)
            preds.append(mc.argmax(dim=1).cpu().numpy())
            probs.append(F.softmax(mc / 0.07, dim=1).cpu().numpy())
    preds = np.concatenate(preds); probs = np.concatenate(probs)
    yv = y[val_ix]
    acc = accuracy_score(yv, preds)
    f1 = f1_score(yv, preds, average="macro", zero_division=0)
    try:
        auc = roc_auc_score(np.eye(len(classes))[yv], probs, average="macro", multi_class="ovr")
    except Exception:
        auc = float("nan")
    rep = classification_report(yv, preds, labels=list(range(len(classes))),
                                target_names=classes, output_dict=True, zero_division=0)
    return acc, f1, auc, rep


def cv(system, variant, folds=5, epochs=5, seed=0):
    print(f"\n=== CV {system}/{variant} ({folds}-fold, {epochs} epochs) ===", flush=True)
    Xpca, Xmark, y, classes, y_dset, datasets, _ = load_and_prepare(system, variant)
    print(f"[cv] n={len(y):,} K={len(classes)}", flush=True)
    skf = StratifiedKFold(n_splits=folds, shuffle=True, random_state=seed)
    accs, f1s, aucs = [], [], []
    last_rep = None
    for fold, (tr, va) in enumerate(skf.split(np.zeros(len(y)), y)):
        model = train_fold(Xpca, Xmark, y, y_dset, classes, variant, tr,
                           epochs=epochs, seed=seed * 100 + fold)
        acc, f1, auc, rep = evaluate(model, Xpca, Xmark, y, va, classes)
        accs.append(acc); f1s.append(f1); aucs.append(auc)
        last_rep = rep
        print(f"[fold {fold+1}] acc={acc:.4f} F1={f1:.4f} AUC={auc:.4f}", flush=True)

    result = {
        "system": system, "variant": variant, "folds": folds, "epochs": epochs, "seed": seed,
        "n_cells": int(len(y)), "n_classes": len(classes),
        "per_class_report_note": "per_class_report is from the LAST fold only, not aggregated",
        "per_fold_acc": accs, "per_fold_f1": f1s, "per_fold_auc": aucs,
        "mean_acc": float(np.mean(accs)), "std_acc": float(np.std(accs)),
        "mean_f1": float(np.mean(f1s)), "std_f1": float(np.std(f1s)),
        "mean_auc": float(np.nanmean(aucs)), "std_auc": float(np.nanstd(aucs)),
        "per_class_report": last_rep,
    }
    out_dir = ROOT / f"discovery/{system}/{variant}"
    out_dir.mkdir(parents=True, exist_ok=True)
    # seed 0 is the canonical file; other seeds get a suffix so true seed-replicates
    # (same corpus, same script, same curriculum) can be compared side by side
    fname = "cv_5fold.json" if seed == 0 else f"cv_5fold_seed{seed}.json"
    (out_dir / fname).write_text(json.dumps(result, indent=2, default=str))
    print(f"[cv] mean acc={result['mean_acc']:.4f}±{result['std_acc']:.4f} "
          f"F1={result['mean_f1']:.4f} AUC={result['mean_auc']:.4f}", flush=True)


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--systems", nargs="*", default=["pan_skin", "hematopoiesis", "pancreas"])
    ap.add_argument("--variants", nargs="*", default=["pca", "marker"])
    ap.add_argument("--folds", type=int, default=5)
    ap.add_argument("--epochs", type=int, default=5)
    ap.add_argument("--seed", type=int, default=0,
                    help="fold-assignment + model-init seed; non-zero seeds write cv_5fold_seed{N}.json")
    args = ap.parse_args()
    for s in args.systems:
        for v in args.variants:
            cv(s, v, args.folds, args.epochs, seed=args.seed)