File size: 8,425 Bytes
6fa9282
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
"""Train the approved residual map and retain validation-selected checkpoints."""

from __future__ import annotations
import copy, json, time, platform
from pathlib import Path
from dataclasses import dataclass, asdict
import numpy as np
import torch
from pivot.models.pivot import PIVOT
from pivot.models.encoders import build_pert_tensors
from pivot.training.losses import compute_losses
from pivot.evaluation.rewards import rbf_mmd2
from pivot.utils.common import set_seed


@dataclass
class TrainConfig:
    d_pert: int = 64
    hidden: int = 512
    depth: int = 4
    epochs: int = 60
    batch_size: int = 1024
    lr: float = 1e-3
    weight_decay: float = 1e-5
    lam_tan: float = 1.0
    lam_semi: float = 0.5
    lam_reg: float = 1e-4
    lam_dist: float = 0.0
    n_dist_perts: int = 4
    dist_n: int = 64
    grad_clip: float = 5.0
    match: str = "batch"
    rep_mode: str = "gene_op"
    seed: int = 0
    device: str = "cpu"
    threads: int = 4
    endpoint_weight: float = 0.0


def make_model(data, cfg):
    if cfg.rep_mode not in ("gene_op", "gene_only", "op_only", "random_id"):
        raise ValueError("Supported encoders use gene/operation metadata only")
    return PIVOT(
        data.d,
        len(data.genes_vocab),
        len(data.op_vocab),
        len(data.perturbations),
        d_pert=cfg.d_pert,
        hidden=cfg.hidden,
        depth=cfg.depth,
        rep_mode=cfg.rep_mode,
    ).to(cfg.device)


@torch.no_grad()
def validation_loss(model, data, cfg) -> float:
    """Mean condition-level endpoint MSE against validation centroids in PCA space."""
    rng = np.random.default_rng(cfg.seed + 910)
    ctr = data.indices("val", True)
    vid = data.indices("val", False)
    c0 = torch.as_tensor(
        data.emb[rng.choice(ctr, min(128, len(ctr)), replace=False)], device=cfg.device
    )
    model.eval()
    loss = []
    for label in data.labels("val"):
        ids = np.intersect1d(data.pert_to_idx[label], vid)
        g, o, m, pid = build_pert_tensors(data, [label], cfg.device)
        pred = model.endpoint_from_pert(c0, g, o, m, pid).mean(0)
        truth = torch.as_tensor(data.emb[ids].mean(0), device=cfg.device)
        loss.append((pred - truth).square().mean().item())
    return float(np.mean(loss))


def train(data, cfg: TrainConfig, output: str, resume: str | None = None) -> dict:
    """Fit using training cells/controls; write best.pt, last.pt, config and history.

    Checkpoints retain the vocabulary and cache fingerprint. Resume restores
    optimizer, scheduler, NumPy RNG, and Torch RNG state before the next epoch.
    """
    set_seed(cfg.seed)
    torch.set_num_threads(cfg.threads)
    if cfg.device.startswith("cuda") and not torch.cuda.is_available():
        raise RuntimeError("Requested CUDA is unavailable")
    out = Path(output)
    out.mkdir(parents=True, exist_ok=True)
    model = make_model(data, cfg)
    opt = torch.optim.AdamW(
        model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay
    )
    sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=cfg.epochs)
    rng = np.random.default_rng(cfg.seed)
    start = 0
    history = []
    best = float("inf")
    if resume:
        ck = torch.load(resume, map_location=cfg.device, weights_only=False)
        if ck["data_meta"] != data.meta or ck["config"] != asdict(cfg):
            raise ValueError("Resume configuration or data mismatch")
        model.load_state_dict(ck["model"])
        opt.load_state_dict(ck["optimizer"])
        sched.load_state_dict(ck["scheduler"])
        rng.bit_generator.state = ck["numpy_rng"]
        torch.set_rng_state(ck["torch_rng"].cpu())
        if ck.get("cuda_rng") is not None and torch.cuda.is_available():
            torch.cuda.set_rng_state_all(ck["cuda_rng"])
        start = ck["epoch"] + 1
        history = ck["history"]
        best = ck["best_val"]
    train_ids = data.indices("train", False)
    ctrl = data.indices("train", True)
    labels = data.obs.perturbation.to_numpy()
    z = torch.as_tensor(data.emb, device=cfg.device)
    groups = {p: train_ids[labels[train_ids] == p] for p in data.labels("train")}
    lam = {"map": 1.0, "tan": cfg.lam_tan, "semi": cfg.lam_semi, "reg": cfg.lam_reg}
    t0 = time.perf_counter()
    for epoch in range(start, cfg.epochs):
        model.train()
        terms = []
        for ids in np.array_split(
            rng.permutation(train_ids), int(np.ceil(len(train_ids) / cfg.batch_size))
        ):
            ci = data.sample_controls(ids, cfg.match, rng, ctrl)
            g, o, m, pid = build_pert_tensors(data, labels[ids], cfg.device)
            e = model.encode(g, o, m, pid)
            total, parts = compute_losses(model.flow, e, z[ci], z[ids], lam)
            if cfg.endpoint_weight:
                le = (model.flow.endpoint(z[ci], e) - z[ids]).square().sum(-1).mean()
                total = total + cfg.endpoint_weight * le
                parts["endpoint"] = le.item()
            if cfg.lam_dist:
                ds = []
                for p in rng.choice(
                    list(groups), min(cfg.n_dist_perts, len(groups)), replace=False
                ):
                    yi = rng.choice(
                        groups[p], min(cfg.dist_n, len(groups[p])), replace=False
                    )
                    xi = data.sample_controls(yi, cfg.match, rng, ctrl)
                    gd, od, md, pd = build_pert_tensors(data, [p], cfg.device)
                    yp = model.endpoint_from_pert(z[xi], gd, od, md, pd)
                    ds.append(rbf_mmd2(yp, z[yi], data.meta["mmd_gamma"]))
                ld = torch.stack(ds).mean()
                total = total + cfg.lam_dist * ld
                parts["dist"] = ld.item()
            opt.zero_grad()
            total.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.grad_clip)
            opt.step()
            parts["total"] = total.item()
            terms.append(parts)
        sched.step()
        val = validation_loss(model, data, cfg)
        record = {
            "epoch": epoch,
            "validation_mse": val,
            **{k: float(np.mean([t[k] for t in terms])) for k in terms[0]},
        }
        history.append(record)
        improved = val < best
        best = min(best, val)
        ck = {
            "protocol": "split-first-v1",
            "model": model.state_dict(),
            "optimizer": opt.state_dict(),
            "scheduler": sched.state_dict(),
            "epoch": epoch,
            "best_val": best,
            "config": asdict(cfg),
            "data_meta": data.meta,
            "gene_vocab": data.genes_vocab,
            "perturbations": data.perturbations,
            "numpy_rng": rng.bit_generator.state,
            "torch_rng": torch.get_rng_state(),
            "cuda_rng": (
                torch.cuda.get_rng_state_all() if torch.cuda.is_available() else None
            ),
            "history": history,
        }
        torch.save(ck, out / "last.pt")
        if improved:
            torch.save(ck, out / "best.pt")
        print(
            f"epoch {epoch+1}/{cfg.epochs} train={record['total']:.4f} val={val:.4f}",
            flush=True,
        )
    info = {
        "protocol": "split-first-v1",
        "config": asdict(cfg),
        "history": history,
        "duration_s": time.perf_counter() - t0,
        "n_train_cells": len(train_ids),
        "n_parameters": sum(p.numel() for p in model.parameters()),
        "software": {
            "python": platform.python_version(),
            "torch": torch.__version__,
            "numpy": np.__version__,
        },
    }
    (out / "training.json").write_text(json.dumps(info, indent=2))
    return info


def load_checkpoint(path, data, device="cpu"):
    """Load trusted local checkpoint with an exact cache/vocabulary match."""
    ck = torch.load(path, map_location=device, weights_only=False)
    if ck.get("protocol") != "split-first-v1":
        raise ValueError(
            "Historical weights need their original preprocessing and archived loader"
        )
    if ck["data_meta"] != data.meta or ck["gene_vocab"] != data.genes_vocab:
        raise ValueError("Checkpoint and cache do not match")
    cfg = TrainConfig(**ck["config"])
    cfg.device = device
    torch.set_num_threads(cfg.threads)
    model = make_model(data, cfg)
    model.load_state_dict(ck["model"])
    model.eval()
    return model, cfg