"""The training loop. `train()` owns the epoch structure and nothing else: the per-stage step is `giant.training.trainers`' job, every number reported is `giant.training.metrics`' job, and the on-disk checkpoint is `giant.training.checkpoint`'s. """ import os import signal import time from collections.abc import Callable from pathlib import Path from types import FrameType from typing import Self import numpy as np import torch from torch.utils.data import DataLoader from tqdm import tqdm from giant.data.loader import TopNMap from giant.data.setup_cache import topnmap_to_json from giant.training.checkpoint import build_checkpoint, init_stages_from_checkpoints, load_checkpoint from giant.training.metrics import MetricsCollector from giant.training.trainers import ( FlowDDPMStageTrainer, StageTrainer, build_stage_trainers, ) from giant.validate import validate_marginals _CATCHABLE_SIGNALS = (signal.SIGINT, signal.SIGTERM) class _GracefulShutdown: """Turns SIGINT/SIGTERM into a flag check instead of an immediate crash. A second signal while already shutting down restores the default handler and re-sends the signal, so an unresponsive run can still be force-killed. """ def __init__(self) -> None: self.requested = False self._previous: dict[ int, Callable[[int, FrameType | None], object] | signal.Handlers | int | None, ] = {} def __enter__(self) -> Self: for sig in _CATCHABLE_SIGNALS: self._previous[sig] = signal.getsignal(sig) signal.signal(sig, self._handle) return self def __exit__(self, *exc_info) -> None: for sig, handler in self._previous.items(): signal.signal(sig, handler) def _handle(self, signum: int, frame) -> None: if self.requested: signal.signal(signum, self._previous[signum]) os.kill(os.getpid(), signum) return self.requested = True print( f"\nreceived {signal.Signals(signum).name} — finishing the current " "batch, then saving a checkpoint and exiting (send again to force-quit)" ) def _try_validate_marginals(trainer: StageTrainer, val_loader, device, **kwargs): """Runs `validate_marginals` on `trainer`'s sampling model (EMA model if present, else the raw model). `validate_marginals` itself dispatches through `giant.sample.sample_stage1`/`sample_stage2`/`resolve_n_sec`, so this is generator- and one-shot-vs-autoregressive-agnostic.""" model = trainer.sampling_model() return validate_marginals(model, val_loader, device=device, **kwargs) def _marginal_kl(trainers: dict[str, StageTrainer], val_loader, device, **kwargs) -> float: """Mean marginal KL over the stage-1 sampling chain, or NaN when stage 1 is inactive or `validate_marginals` declined to produce a result.""" stage1 = trainers.get("stage1") if stage1 is None: return float("nan") result = _try_validate_marginals( stage1, val_loader, device, sec_decoder=trainers["stage2"].sampling_model() if "stage2" in trainers else None, **kwargs, ) if result is None: return float("nan") return float(np.mean(result["kl_divergence"])) def train( cfg: dict, models: dict[str, torch.nn.Module | None], critics: dict[str, torch.nn.Module | None], train_loader: DataLoader, val_loader: DataLoader, device: torch.device, out_dir: str | Path, normalizer_dict: dict | None = None, pdg_map: dict | None = None, mat_map: dict | None = None, proc_map: dict | None = None, pdg_topn_map: TopNMap | None = None, sec_type_topn_map: TopNMap | None = None, mat_topn_map: TopNMap | None = None, model_config: dict | None = None, resume_path: str | Path | None = None, total_train_batches: int = 0, use_wandb: bool = False, wandb_project: str = "giant", wandb_run_name: str = "", wandb_log_every: int = 50, ) -> None: """Train whichever of stage1/stage2 are active, each through its own `StageTrainer`. `models`/`critics` are the dicts `giant.model.network.build_models`/`build_critics` return — a `None` entry means that stage is `active = false`. """ out_dir = Path(out_dir) out_dir.mkdir(parents=True, exist_ok=True) t = cfg["train"] epochs = t["epochs"] validate_every = t.get("validate_every", 0) validate_steps = t.get("validate_steps", 10) max_val_batches = t.get("max_val_batches", 0) sec_type_class_counts = sec_type_topn_map.class_counts if sec_type_topn_map is not None else None trainers = build_stage_trainers(cfg, models, critics, device, total_train_batches, sec_type_class_counts) if not trainers: raise ValueError("no active stage — stage1_model.active and stage2_model.active are both false") for line in init_stages_from_checkpoints(trainers): print(line) has_adversarial = any(not tr.supports_val_loss for tr in trainers.values()) checkpoint_extras = { "normalizer": normalizer_dict, "pdg_map": pdg_map, "mat_map": mat_map, "proc_map": proc_map, "pdg_topn_map": topnmap_to_json(pdg_topn_map) if pdg_topn_map is not None else None, "sec_type_topn_map": topnmap_to_json(sec_type_topn_map) if sec_type_topn_map is not None else None, "mat_topn_map": topnmap_to_json(mat_topn_map) if mat_topn_map is not None else None, "model_config": model_config, } start_epoch = 1 best_val_loss = float("inf") global_step = 0 if resume_path is not None: ckpt = torch.load(resume_path, map_location=device, weights_only=False) load_checkpoint(trainers, ckpt, t["lr"]) start_epoch = ckpt.get("epoch", 0) + 1 best_val_loss = ckpt.get("best_val_loss", float("inf")) global_step = ckpt.get("global_step", 0) if start_epoch > epochs: print(f"checkpoint already completed epoch {start_epoch - 1} (>= --epochs {epochs}) — nothing to train") return collector = MetricsCollector.create( trainers, out_dir, cfg, model_config, resume=resume_path is not None, use_wandb=use_wandb, wandb_project=wandb_project, wandb_run_name=wandb_run_name, wandb_log_every=wandb_log_every, ) epoch_w = len(str(epochs)) last_completed_epoch = start_epoch - 1 with _GracefulShutdown() as shutdown: for epoch in range(start_epoch, epochs + 1): epoch_start = time.monotonic() if device.type == "cuda": torch.cuda.reset_peak_memory_stats(device) collector.start_epoch(epoch) for trainer in trainers.values(): trainer.train_mode() bar = tqdm( train_loader, desc=f" epoch {epoch:{epoch_w}d}/{epochs}", total=total_train_batches or None, leave=False, unit="batch", dynamic_ncols=True, ) for batch in bar: B = batch[0].size(0) collector.add_train_batch( {name: trainer.step(batch, device, global_step) for name, trainer in trainers.items()}, B, ) bar.set_postfix_str(collector.postfix(), refresh=False) global_step += 1 collector.log_batch(global_step, batch, device) if shutdown.requested: break bar.close() if shutdown.requested: ckpt = build_checkpoint(trainers, epoch - 1, global_step, best_val_loss, checkpoint_extras) torch.save(ckpt, out_dir / "last.pt") last_completed_epoch = epoch - 1 print( f"saved in-progress weights from partway through epoch " f"{epoch} to {out_dir / 'last.pt'} " f"(resume will restart epoch {epoch})" ) break for trainer in trainers.values(): trainer.eval_mode() # --- per-stage validation --- scored = {name: tr for name, tr in trainers.items() if tr.supports_val_loss} if scored: with torch.no_grad(): for val_batch_idx, batch in enumerate(val_loader): if max_val_batches > 0 and val_batch_idx >= max_val_batches: break B = batch[0].size(0) collector.add_val_batch( {name: tr.val_loss(batch, device) for name, tr in scored.items()}, B, ) collector.observe_routers(batch[0].to(device), batch[1].to(device), B) # An adversarial stage has no averageable validation loss, so it # needs the marginal-KL signal every epoch to pick a best # checkpoint at all; a purely non-adversarial run only pays for # it every `validate_every` epochs. marginal_kl = float("nan") if has_adversarial: marginal_kl = _marginal_kl(trainers, val_loader, device) elif validate_every > 0 and epoch % validate_every == 0: stage1 = trainers.get("stage1") ddpm_steps = 1000 if isinstance(stage1, FlowDDPMStageTrainer) and stage1.ddpm_schedule is not None: ddpm_steps = stage1.ddpm_schedule.T marginal_kl = _marginal_kl( trainers, val_loader, device, steps=validate_steps, ddpm_steps=ddpm_steps, ) val_loss = sum( trainer.val_objective( collector.train_means(name), collector.val_means(name), marginal_kl, ) for name, trainer in trainers.items() ) epoch_time = time.monotonic() - epoch_start is_best = val_loss < best_val_loss collector.set("val/loss", val_loss) collector.set("val/marginal_kl", marginal_kl) collector.set( "gpu_mem_mb", torch.cuda.max_memory_allocated(device) / (1024 * 1024) if device.type == "cuda" else 0.0, ) collector.set("samples_per_sec", collector.train_samples / max(epoch_time, 1e-8)) collector.set("is_best", int(is_best)) collector.set("epoch_time_s", epoch_time) print(collector.summary_line(val_loss, epoch_time, is_best)) collector.write_epoch(global_step) ckpt = build_checkpoint(trainers, epoch, global_step, best_val_loss, checkpoint_extras) if is_best: best_val_loss = val_loss ckpt["best_val_loss"] = best_val_loss torch.save(ckpt, out_dir / "best.pt") torch.save(ckpt, out_dir / "last.pt") last_completed_epoch = epoch if shutdown.requested: break collector.close() if shutdown.requested: print( f"stopped after epoch {last_completed_epoch} due to shutdown signal — " f"resume with --resume {out_dir / 'last.pt'}" )