Files
giant/giant/train.py
T
2026-06-22 07:27:02 +02:00

249 lines
8.5 KiB
Python

import csv
import math
import os
import signal
import time
from pathlib import Path
from types import FrameType
from typing import Callable
import torch
import torch.optim as optim
from torch.utils.data import DataLoader
from tqdm import tqdm
from giant.model.schedule import CosineSchedule, flow_matching_loss
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,
Callable[[int, FrameType | None], object] | signal.Handlers | int | None,
] = {}
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,
train_loader: DataLoader,
val_loader: DataLoader,
mode: str,
epochs: int,
lr: float,
warmup_epochs: int,
device: torch.device,
out_dir: str | Path,
normalizer_dict: dict | None = None,
pdg_map: dict | None = None,
mat_map: dict | None = None,
model_config: dict | None = None,
resume_path: str | Path | None = None,
validate_every: int = 0,
validate_steps: int = 10,
total_train_batches: int = 0,
) -> None:
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
model = model.to(device)
optimizer = optim.AdamW(model.parameters(), lr=lr)
def _lr_lambda(epoch: int) -> float:
if warmup_epochs > 0 and epoch < warmup_epochs:
return (epoch + 1) / warmup_epochs
t = epoch - warmup_epochs
T = max(epochs - warmup_epochs, 1)
return 0.5 * (1.0 + math.cos(math.pi * t / T))
lr_sched = optim.lr_scheduler.LambdaLR(optimizer, _lr_lambda)
ddpm_schedule = CosineSchedule().to(device) if mode == "ddpm" else None
start_epoch = 1
best_val_loss = float("inf")
if resume_path is not None:
ckpt = torch.load(resume_path, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
lr_sched.load_state_dict(ckpt["lr_sched"])
start_epoch = ckpt.get("epoch", 0) + 1
best_val_loss = ckpt.get("best_val_loss", float("inf"))
metrics_path = out_dir / "metrics.csv"
write_header = not (resume_path is not None and metrics_path.exists())
metrics_file = open(metrics_path, "a", newline="")
metrics_writer = csv.DictWriter(metrics_file, fieldnames=_METRICS_FIELDS)
if write_header:
metrics_writer.writeheader()
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()
current_lr = optimizer.param_groups[0]["lr"]
model.train()
train_loss_sum = 0.0
train_n = 0
ema_loss = 0.0
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 cond_cont, cond_cat, x1 in bar:
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:
assert ddpm_schedule is not None
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()
batch_loss = loss.item()
train_loss_sum += batch_loss * x1.size(0)
train_n += x1.size(0)
ema_loss = (
batch_loss
if train_n == x1.size(0)
else 0.95 * ema_loss + 0.05 * batch_loss
)
bar.set_postfix_str(f"loss={ema_loss:.4f}", refresh=False)
if shutdown.requested:
break
bar.close()
if shutdown.requested:
# Mid-epoch: discard the partial epoch rather than persist an
# inconsistent (lr_sched not stepped, no validation) checkpoint.
break
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:
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:
assert ddpm_schedule is not None
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
is_best = val_loss < best_val_loss
marker = " [best]" if is_best else ""
print(
f"epoch {epoch:{epoch_w}d}/{epochs}"
f" train {train_loss:.4f} val {val_loss:.4f}"
f" lr {current_lr:.2e} {epoch_time:.1f}s{marker}"
)
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,
steps=validate_steps,
)
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'}"
)