File size: 8,164 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
"""held-out CV for panda. GroupKFold by dataset when possible, else StratifiedKFold."""
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
from sklearn.model_selection import GroupKFold, StratifiedKFold
from sklearn.metrics import accuracy_score, f1_score, roc_auc_score, classification_report
warnings.filterwarnings("ignore")

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,
)
from scripts.common.train_panda import prepare_batches, load_corpus, get_marker_gene_list

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


def train_one_fold(Xpca, Xmark, y, y_dset, log10cz, classes, variant,

                   train_idx, epochs=6, batch=256, lr=1e-3, seed=0):
    n_classes = len(classes)
    n_datasets = int(max(y_dset[train_idx].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)

    train_n = len(train_idx)
    rng = np.random.default_rng(seed)
    Xt = Xpca[train_idx]; yt = y[train_idx]; ydt = y_dset[train_idx]; dt = log10cz[train_idx]
    Xmt = Xmark[train_idx] if Xmark is not None else None

    for epoch in range(epochs):
        stage = 0 if epoch < 1 else 1 if epoch < 3 else 2 if epoch < 6 else 3
        for g in opt.param_groups: g["lr"] = lr * (0.5 if epoch >= epochs - 1 else 1.0)
        perm = rng.permutation(train_n)
        for bstart in range(0, train_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)
            dd = torch.from_numpy(dt[idx]).float().to(DEVICE).unsqueeze(1)
            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) + 0.3 * F.mse_loss(out["depth"], dd) \
                      + 0.05 * hsic_biased(out["repr"], dd)
            if stage >= 3:
                L = L + 0.5 * prototype_repulsion(model.prototypes.detach().clone())
            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_idx, classes):
    model.eval()
    preds, probs = [], []
    Xv = Xpca[val_idx]; Xmv = Xmark[val_idx] if Xmark is not None else None
    with torch.no_grad():
        for i in range(0, len(val_idx), 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)   # (B, K)
            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_idx]
    acc = accuracy_score(yv, preds)
    f1 = f1_score(yv, preds, average="macro", zero_division=0)
    # macro AUROC — nan if val has <2 classes
    n_classes = len(classes)
    try:
        y_onehot = np.eye(n_classes)[yv]
        auc = roc_auc_score(y_onehot, probs, average="macro", multi_class="ovr")
    except Exception:
        auc = float("nan")
    rep = classification_report(yv, preds, labels=list(range(n_classes)),
                                target_names=classes, output_dict=True, zero_division=0)
    return acc, f1, auc, rep


def cv(system, variant, folds=5, epochs=6, split_mode="auto", seed=0):
    a, hvgs, mu, sig, pca = load_corpus(system)
    marker_genes = get_marker_gene_list(system) if variant == "marker" else []
    Xpca, Xmark, y, classes, y_dset, dset_classes, log10cz = prepare_batches(
        a, hvgs, mu, sig, pca, marker_genes, variant
    )
    print(f"[cv] {system}/{variant} n={a.n_obs} K={len(classes)} datasets={len(dset_classes)}", flush=True)

    use_group = (split_mode == "group") or (split_mode == "auto" and len(dset_classes) >= folds)
    if use_group:
        splitter = GroupKFold(n_splits=folds)
        splits = list(splitter.split(np.zeros(len(y)), y, y_dset))
        print(f"[cv] GroupKFold by dataset ({len(dset_classes)} groups)", flush=True)
    else:
        splitter = StratifiedKFold(n_splits=folds, shuffle=True, random_state=seed)
        splits = list(splitter.split(np.zeros(len(y)), y))
        print(f"[cv] StratifiedKFold on labels ({split_mode}) seed={seed}", flush=True)

    per_fold_acc, per_fold_f1, per_fold_auc = [], [], []
    last_report = None
    for fold, (tr, va) in enumerate(splits):
        print(f"[fold {fold+1}/{folds}] train={len(tr)} val={len(va)}", flush=True)
        model = train_one_fold(Xpca, Xmark, y, y_dset, log10cz, classes, variant, tr,
                               epochs=epochs, seed=seed * 100 + fold)
        acc, f1, auc, rep = evaluate(model, Xpca, Xmark, y, va, classes)
        per_fold_acc.append(acc); per_fold_f1.append(f1); per_fold_auc.append(auc)
        last_report = 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(a.n_obs), "n_classes": len(classes),
        "per_fold_acc": per_fold_acc, "per_fold_f1": per_fold_f1, "per_fold_auc": per_fold_auc,
        "mean_acc": float(np.mean(per_fold_acc)), "std_acc": float(np.std(per_fold_acc)),
        "mean_f1":  float(np.mean(per_fold_f1)),  "std_f1":  float(np.std(per_fold_f1)),
        "mean_auc": float(np.nanmean(per_fold_auc)),
        "std_auc":  float(np.nanstd(per_fold_auc)),
        "per_class_report": last_report,
    }
    out_dir = ROOT / f"discovery/{system}/{variant}"
    out_dir.mkdir(parents=True, exist_ok=True)
    suffix = f"_seed{seed}" if seed != 0 else ""
    (out_dir / f"cv_{folds}fold{suffix}.json").write_text(json.dumps(result, indent=2, default=str))
    print(f"\n[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("system", choices=["pan_skin", "hematopoiesis", "pancreas"])
    ap.add_argument("--variant", choices=["pca", "marker"], required=True)
    ap.add_argument("--folds", type=int, default=5)
    ap.add_argument("--epochs", type=int, default=6)
    ap.add_argument("--split", choices=["auto", "group", "stratified"], default="auto")
    ap.add_argument("--seed", type=int, default=0)
    args = ap.parse_args()
    cv(args.system, args.variant, args.folds, args.epochs, args.split, args.seed)