File size: 8,624 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
"""train panda on the canonical paper-labeled corpus for one system + variant."""
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

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_corpus(system):
    """canonical loader; alias for load_corpus_v3 after finalize_rename."""
    return load_corpus_v3(system)


def load_corpus_v3(system):
    # kept for backward-compat with older scripts that import load_corpus_v3
    p = 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"))
    a = ad.read_h5ad(p)
    hvgs = [str(g) for g in stats["shared_hvgs"]]
    return a, hvgs, stats["mean"], stats["std"], pca


def get_marker_gene_list(system):
    y = yaml.safe_load(open(ROOT / "panda/markers.yaml"))
    return y[system]


def prepare_batches(adata, hvgs, mu, sig, pca, marker_genes=None, variant="pca",

                    legacy_double_norm=True):
    # corpus.h5ad X is already normalize_total+log1p'd by the corpus builders, and the
    # published checkpoints were trained with a second normalize/log1p applied on top
    # (against per-gene stats computed from singly-normalized data). legacy_double_norm=True
    # reproduces that behaviour bit-for-bit; pass False to train on the corpus values the
    # stats were actually computed from. Do not mix: a checkpoint must be evaluated under
    # the same setting it was trained with.
    hvg2i = {g: i for i, g in enumerate(hvgs)}
    common = [g for g in adata.var_names.astype(str) if g in hvg2i]
    a_c = adata[:, common].copy()
    if legacy_double_norm:
        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((adata.n_obs, len(hvgs)), dtype=np.float32)
    cols = np.array([hvg2i[g] for g in common])
    Xf[:, cols] = X_
    Xz = np.clip((Xf - mu.astype(np.float32)) / sig.astype(np.float32), -10, 10)
    Xpca = pca.transform(Xz).astype(np.float32)

    Xmark = None
    if variant == "marker" and marker_genes:
        mvals = np.zeros((adata.n_obs, len(marker_genes)), dtype=np.float32)
        for j, g in enumerate(marker_genes):
            if g in adata.var_names:
                col = adata[:, 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)
        # expose the training-corpus marker stats so the checkpoint can carry them;
        # inference must reuse these rather than refit on the target dataset
        prepare_batches.last_marker_stats = (mmu, msig)

    labels = adata.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(adata.obs["dataset"].astype(str).values))
    y_dset = np.array([datasets.index(d) for d in adata.obs["dataset"].astype(str).values], dtype=np.int64)
    counts = np.asarray(adata.X.sum(axis=1)).ravel()
    log10cz = ((np.log10(counts + 1) - np.log10(counts + 1).mean()) /
               (np.log10(counts + 1).std() + 1e-6)).astype(np.float32)
    return Xpca, Xmark, y, classes, y_dset, datasets, log10cz


def train(system, variant, epochs=8, batch=256, lr=1e-3):
    a, hvgs, mu, sig, pca = load_corpus_v3(system)
    marker_genes = get_marker_gene_list(system) if variant == "marker" else []
    Xpca, Xmark, y, classes, y_dset, datasets, log10cz = prepare_batches(
        a, hvgs, mu, sig, pca, marker_genes, variant
    )
    print(f"[train] {system}/{variant} n={a.n_obs} K={len(classes)} datasets={len(datasets)}", flush=True)
    print(f"[train] classes: {classes}", flush=True)
    n_markers = Xmark.shape[1] if Xmark is not None else 0

    model = PANDAEncoder(variant=variant, n_pca=50, n_markers=n_markers,
                         n_classes=len(classes), n_sub=3, n_datasets=len(datasets), dropout=0.2).to(DEVICE)
    opt = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
    rng = np.random.default_rng(0)
    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(a.n_obs)
        losses = []
        for bstart in range(0, a.n_obs, batch):
            idx = perm[bstart:bstart+batch]
            x = torch.from_numpy(Xpca[idx]).to(DEVICE)
            xm = torch.from_numpy(Xmark[idx]).to(DEVICE) if Xmark is not None else None
            yy = torch.from_numpy(y[idx]).to(DEVICE)
            yd = torch.from_numpy(y_dset[idx]).to(DEVICE)
            dd = torch.from_numpy(log10cz[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)
            # NOTE: earlier revisions added 0.5 * prototype_repulsion(model.prototypes.detach())
            # at stage 3. prototypes is a gradient-free EMA buffer and the tensor was detached,
            # so the term contributed exactly zero gradient — it only inflated the printed loss.
            # Removed from the objective (behaviour-preserving); prototype_repulsion() remains
            # in panda.model as an analysis metric.
            opt.zero_grad(); L.backward(); opt.step()
            if stage >= 1:
                with torch.no_grad(): model.update_prototypes(z.detach(), yy)
            losses.append(float(L))
        print(f"[train {system}/{variant}] epoch {epoch}/{epochs} stage={stage} loss={np.mean(losses):.4f}", flush=True)

    # save to the same path every consumer (zero_shot, run_all_zero_shot, extract_prototypes)
    # loads from. Earlier revisions wrote to checkpoints/{system}_v3/ while consumers read
    # checkpoints/{system}/ — a retrain silently never propagated.
    ck_dir = ROOT / f"checkpoints/{system}/{variant}"
    ck_dir.mkdir(parents=True, exist_ok=True)
    mstats = getattr(prepare_batches, "last_marker_stats", None) if variant == "marker" else None
    torch.save({"model": model.state_dict(), "classes": classes, "datasets": datasets,
                "marker_genes": marker_genes if variant == "marker" else [],
                "marker_mu": mstats[0] if mstats else None,
                "marker_sig": mstats[1] if mstats else None,
                "legacy_double_norm": True,
                "prototypes": model.prototypes.detach().cpu().numpy()},
               ck_dir / "panda_final.pt")
    print(f"[save] {ck_dir}/panda_final.pt", 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("--epochs", type=int, default=8)
    ap.add_argument("--batch", type=int, default=256)
    ap.add_argument("--lr", type=float, default=1e-3)
    args = ap.parse_args()
    train(args.system, args.variant, args.epochs, args.batch, args.lr)