Download src/pivot/training/train.py from ChatterjeeLab/PIVOT: direct link, hf CLI and curl.
- Browser
- Download file 8.43 kB
-
https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/training/train.py
- Command line
-
hf download hf://ChatterjeeLab/PIVOT/src/pivot/training/train.py
-
curl -L -o train.py https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/training/train.py
8.43 kB
| """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 | |
| 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) | |
| 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 | |