# crvs_engine.py -- the shared GPU training engine. # Dual T4 via DataParallel, AMP, cosine schedule, early stopping, and a checkpoint written # EVERY epoch so a killed session costs nothing. Runs are queued and skipped if already done. import json, math, time, os from pathlib import Path from datetime import datetime, timezone import numpy as np import torch import torch.nn as nn from torch.utils.data import DataLoader def _autocast(device_type, enabled): # torch.cuda.amp.autocast / GradScaler are deprecated and warn on every step in # torch >= 2.4. Use the device-typed API where it exists, fall back where it does not. try: return torch.amp.autocast(device_type=device_type, enabled=enabled) except (AttributeError, TypeError): return torch.cuda.amp.autocast(enabled=enabled) def _grad_scaler(device_type, enabled): try: return torch.amp.GradScaler(device_type, enabled=enabled) except (AttributeError, TypeError): return torch.cuda.amp.GradScaler(enabled=enabled) def pick_device(): if torch.cuda.is_available(): n = torch.cuda.device_count() names = [torch.cuda.get_device_name(i) for i in range(n)] return torch.device("cuda"), n, names return torch.device("cpu"), 0, [] def seed_all(s): import random random.seed(s); np.random.seed(s); torch.manual_seed(s) torch.cuda.manual_seed_all(s) class Trainer: def __init__(self, model, loss_fn, out_dir, run_id, sync=None, lr=5e-4, weight_decay=1e-4, epochs=120, patience=20, batch_size=64, num_workers=2, amp=True, multi_gpu=True, grad_clip=1.0, min_lr=1e-6, log_every=50): self.device, self.ngpu, self.gpu_names = pick_device() self.raw_model = model.to(self.device) self.model = self.raw_model if multi_gpu and self.ngpu > 1: self.model = nn.DataParallel(self.raw_model) self.loss_fn = loss_fn self.out = Path(out_dir); self.out.mkdir(parents=True, exist_ok=True) self.run_id = run_id; self.sync = sync self.epochs = epochs; self.patience = patience self.bs = batch_size; self.nw = num_workers self.amp = amp and self.device.type == "cuda" self.grad_clip = grad_clip self.opt = torch.optim.AdamW(self.raw_model.parameters(), lr=lr, weight_decay=weight_decay) self.sched = torch.optim.lr_scheduler.CosineAnnealingLR(self.opt, T_max=epochs, eta_min=min_lr) self.scaler = _grad_scaler(self.device.type, self.amp) self.log_every = log_every self.state = {"epoch": 0, "best": float("inf"), "best_epoch": -1, "history": [], "run_id": run_id, "done": False} @property def ckpt(self): return self.out / "state.pt" def save(self, tag="state"): torch.save({"model": self.raw_model.state_dict(), "opt": self.opt.state_dict(), "sched": self.sched.state_dict(), "scaler": self.scaler.state_dict(), "state": self.state, "torch_rng": torch.get_rng_state(), "np_rng": np.random.get_state()}, self.out / (tag + ".pt")) (self.out / "state.json").write_text(json.dumps(self.state, indent=2, default=str)) def load(self): if not self.ckpt.exists(): return False try: d = torch.load(self.ckpt, map_location=self.device, weights_only=False) self.raw_model.load_state_dict(d["model"]) self.opt.load_state_dict(d["opt"]); self.sched.load_state_dict(d["sched"]) self.scaler.load_state_dict(d["scaler"]); self.state = d["state"] try: torch.set_rng_state(d["torch_rng"].cpu()); np.random.set_state(d["np_rng"]) except Exception: pass print(f" resumed {self.run_id} at epoch {self.state['epoch']}") return True except Exception as e: print(f" checkpoint unreadable ({type(e).__name__}), starting fresh") return False def _loader(self, ds, shuffle): if len(ds) == 0: raise RuntimeError("empty dataset -- check the fold split; training on nothing " "would produce a plausible-looking but untrained checkpoint") # drop_last=True on a split smaller than one batch yields ZERO batches, the optimiser # never steps, and the loss is reported as 0.00000. Guard it explicitly. drop = bool(shuffle) and len(ds) > self.bs if shuffle and not drop: print(f" note: only {len(ds)} train window(s) < batch {self.bs}; keeping the " f"partial batch so the optimiser actually steps") return DataLoader(ds, batch_size=min(self.bs, max(len(ds), 1)), shuffle=shuffle, num_workers=self.nw, pin_memory=(self.device.type == "cuda"), drop_last=drop, persistent_workers=self.nw > 0) def _step(self, batch, train): x, y, pk, rr = [b.to(self.device, non_blocking=True) for b in batch] with _autocast(self.device.type, self.amp): pred = self.model(x) if isinstance(pred, dict) and "aux" in pred and not train: pred = {k: v for k, v in pred.items() if k != "aux"} loss, parts = self.loss_fn(pred, y, pk, rr) return loss, parts, pred, y def fit(self, train_ds, val_ds): tl = self._loader(train_ds, True) vl = self._loader(val_ds, False) start = self.state["epoch"] # Only the epoch counter decides completion. Keying off a sticky `done` flag meant # raising CFG["EPOCHS"] later silently no-opped instead of training further. if start >= self.epochs: print(f" {self.run_id} already complete at epoch {start}/{self.epochs}") return self.state if self.state.get("done"): print(f" extending {self.run_id}: {start} -> {self.epochs} epochs") self.state["done"] = False bad = 0 for ep in range(start, self.epochs): self.model.train(); t0 = time.time(); tot = 0.0; n = 0 for i, batch in enumerate(tl): self.opt.zero_grad(set_to_none=True) loss, parts, _, _ = self._step(batch, True) self.scaler.scale(loss).backward() if self.grad_clip: self.scaler.unscale_(self.opt) torch.nn.utils.clip_grad_norm_(self.raw_model.parameters(), self.grad_clip) self.scaler.step(self.opt); self.scaler.update() tot += float(loss.detach()); n += 1 if n == 0: raise RuntimeError( "the training loader yielded zero batches -- the optimiser never stepped. " "This would write a checkpoint that looks trained and is not.") self.sched.step() tr = tot / n self.model.eval(); vtot = 0.0; vn = 0 with torch.no_grad(): for batch in vl: loss, _, _, _ = self._step(batch, False) vtot += float(loss); vn += 1 va = vtot / max(vn, 1) rec = {"epoch": ep + 1, "train": tr, "val": va, "lr": self.opt.param_groups[0]["lr"], "sec": round(time.time() - t0, 1)} self.state["history"].append(rec); self.state["epoch"] = ep + 1 improved = va < self.state["best"] - 1e-6 if improved: self.state["best"] = va; self.state["best_epoch"] = ep + 1; bad = 0 self.save("best") else: bad += 1 self.save("state") if self.sync: self.sync.log("epoch", run=self.run_id, **rec, best=round(self.state["best"], 6)) if improved: self.sync.stage_done(f"{self.run_id}:best@{ep+1}") print(f" ep {ep+1:>3}/{self.epochs} train {tr:.5f} val {va:.5f}" f"{' *' if improved else ''} {rec['sec']:.0f}s") if bad >= self.patience: print(f" early stop at epoch {ep+1} (no improvement for {self.patience})") break self.state["done"] = True; self.save("state") if self.sync: self.sync.stage_done(f"{self.run_id}:done") return self.state @torch.no_grad() def predict(self, ds, max_keep=200): bp = self.out / "best.pt" if bp.exists(): try: self.raw_model.load_state_dict( torch.load(bp, map_location=self.device, weights_only=False)["model"]) except Exception as e: print(" could not load best.pt:", e) self.model.eval() dl = self._loader(ds, False) Y, P = [], [] for batch in dl: x, y, pk, rr = [b.to(self.device, non_blocking=True) for b in batch] with _autocast(self.device.type, self.amp): out = self.model(x) Y.append(y.squeeze(1).float().cpu().numpy()) P.append(out["wave"].squeeze(1).float().cpu().numpy()) Y = np.concatenate(Y, 0); P = np.concatenate(P, 0) return Y, P