Add graceful shutdown on SIGINT/SIGTERM

A Ctrl-C or job-scheduler kill signal during training used to crash with a
raw KeyboardInterrupt mid-batch, abandoning whatever checkpoint state was
in flight. Now a signal sets a flag instead: the loop discards an
in-progress epoch's partial work (since lr_sched hasn't stepped and there's
no validation pass yet for it), but lets an epoch that's already past its
training loop finish normally — checkpoint, metrics row, and all — before
stopping. A second signal force-kills immediately for an unresponsive run.

Verified against a backgrounded run: SIGINT mid-training stopped cleanly
with a consistent last.pt/metrics.csv, and --resume picked up exactly at
the next epoch.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-18 14:00:01 +02:00
parent 893d91e749
commit 5d161a52b2
+125 -67
View File
@@ -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'}"
)