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