diff --git a/giant/train.py b/giant/train.py index 3266cf8..3869154 100644 --- a/giant/train.py +++ b/giant/train.py @@ -1,4 +1,6 @@ import csv +import os +import signal import time from pathlib import Path @@ -11,6 +13,42 @@ from giant.validate import validate_marginals _METRICS_FIELDS = ["epoch", "train_loss", "val_loss", "lr", "epoch_time_s"] +_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, object] = {} + + def __enter__(self) -> "_GracefulShutdown": + 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 train( model: torch.nn.Module, @@ -54,84 +92,104 @@ def train( if write_header: metrics_writer.writeheader() - for epoch in range(start_epoch, epochs + 1): - epoch_start = time.monotonic() - current_lr = optimizer.param_groups[0]["lr"] - model.train() - train_loss_sum = 0.0 - train_n = 0 - for cond_cont, cond_cat, x1 in train_loader: - cond_cont = cond_cont.to(device) - cond_cat = cond_cat.to(device) - x1 = x1.to(device) - - if mode == "flow": - loss = flow_matching_loss(model, x1, cond_cont, cond_cat) - else: - loss = ddpm_schedule.loss(model, x1, cond_cont, cond_cat) - - optimizer.zero_grad() - loss.backward() - torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) - optimizer.step() - train_loss_sum += loss.item() * x1.size(0) - train_n += x1.size(0) - - train_loss = train_loss_sum / max(train_n, 1) - lr_sched.step() - - model.eval() - val_loss_sum = 0.0 - val_n = 0 - with torch.no_grad(): - for cond_cont, cond_cat, x1 in val_loader: + last_completed_epoch = start_epoch - 1 + with _GracefulShutdown() as shutdown: + for epoch in range(start_epoch, epochs + 1): + epoch_start = time.monotonic() + current_lr = optimizer.param_groups[0]["lr"] + model.train() + train_loss_sum = 0.0 + train_n = 0 + for cond_cont, cond_cat, x1 in train_loader: cond_cont = cond_cont.to(device) cond_cat = cond_cat.to(device) x1 = x1.to(device) + if mode == "flow": loss = flow_matching_loss(model, x1, cond_cont, cond_cat) else: loss = ddpm_schedule.loss(model, x1, cond_cont, cond_cat) - val_loss_sum += loss.item() * x1.size(0) - val_n += x1.size(0) - val_loss = val_loss_sum / max(val_n, 1) - epoch_time = time.monotonic() - epoch_start - print( - f"epoch {epoch:4d} train {train_loss:.4f} val {val_loss:.4f} " - f"lr {current_lr:.2e} {epoch_time:.1f}s" - ) - metrics_writer.writerow({ - "epoch": epoch, "train_loss": train_loss, "val_loss": val_loss, - "lr": current_lr, "epoch_time_s": epoch_time, - }) - metrics_file.flush() + optimizer.zero_grad() + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) + optimizer.step() + train_loss_sum += loss.item() * x1.size(0) + train_n += x1.size(0) - if validate_every > 0 and epoch % validate_every == 0: - print(f"[epoch {epoch}] marginal validation:") - validate_marginals(model, val_loader, mode=mode, schedule=ddpm_schedule, device=device) + if shutdown.requested: + break - ckpt: dict = { - "model": model.state_dict(), - "optimizer": optimizer.state_dict(), - "lr_sched": lr_sched.state_dict(), - "epoch": epoch, - "best_val_loss": best_val_loss, - } - if normalizer_dict is not None: - ckpt["normalizer"] = normalizer_dict - if pdg_map is not None: - ckpt["pdg_map"] = pdg_map - if mat_map is not None: - ckpt["mat_map"] = mat_map - if model_config is not None: - ckpt["model_config"] = model_config + if shutdown.requested: + # Mid-epoch: discard the partial epoch rather than persist an + # inconsistent (lr_sched not stepped, no validation) checkpoint. + break - if val_loss < best_val_loss: - best_val_loss = val_loss - ckpt["best_val_loss"] = best_val_loss - torch.save(ckpt, out_dir / "best.pt") + train_loss = train_loss_sum / max(train_n, 1) + lr_sched.step() - torch.save(ckpt, out_dir / "last.pt") + model.eval() + val_loss_sum = 0.0 + val_n = 0 + with torch.no_grad(): + for cond_cont, cond_cat, x1 in val_loader: + cond_cont = cond_cont.to(device) + cond_cat = cond_cat.to(device) + x1 = x1.to(device) + if mode == "flow": + loss = flow_matching_loss(model, x1, cond_cont, cond_cat) + else: + loss = ddpm_schedule.loss(model, x1, cond_cont, cond_cat) + val_loss_sum += loss.item() * x1.size(0) + val_n += x1.size(0) + val_loss = val_loss_sum / max(val_n, 1) + epoch_time = time.monotonic() - epoch_start + + print( + f"epoch {epoch:4d} train {train_loss:.4f} val {val_loss:.4f} " + f"lr {current_lr:.2e} {epoch_time:.1f}s" + ) + metrics_writer.writerow({ + "epoch": epoch, "train_loss": train_loss, "val_loss": val_loss, + "lr": current_lr, "epoch_time_s": epoch_time, + }) + metrics_file.flush() + + if validate_every > 0 and epoch % validate_every == 0: + print(f"[epoch {epoch}] marginal validation:") + validate_marginals(model, val_loader, mode=mode, schedule=ddpm_schedule, device=device) + + ckpt: dict = { + "model": model.state_dict(), + "optimizer": optimizer.state_dict(), + "lr_sched": lr_sched.state_dict(), + "epoch": epoch, + "best_val_loss": best_val_loss, + } + if normalizer_dict is not None: + ckpt["normalizer"] = normalizer_dict + if pdg_map is not None: + ckpt["pdg_map"] = pdg_map + if mat_map is not None: + ckpt["mat_map"] = mat_map + if model_config is not None: + ckpt["model_config"] = model_config + + if val_loss < best_val_loss: + 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 metrics_file.close() + + if shutdown.requested: + print( + f"stopped after epoch {last_completed_epoch} due to shutdown signal — " + f"resume with --resume {out_dir / 'last.pt'}" + )