"""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