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:
+125
-67
@@ -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'}"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user