249 lines
8.5 KiB
Python
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'}"
|
|
)
|