8019a80563
CI / Format (ruff format) (push) Successful in 27s
CI / Lint (ruff check) (push) Successful in 27s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 39s
CI / Type check (ty) (push) Successful in 43s
CI / Format (ruff format) (pull_request) Successful in 31s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 31s
CI / Tests (push) Successful in 2m35s
CI / Tests (pull_request) Successful in 2m36s
Every metric name used to exist in four places: the dict keys each
StageTrainer returned, the hardcoded _metrics_fields() column list, the
~110-line metrics_row assembly in train(), and the tqdm/summary
formatting. The two had to be kept in exact correspondence by hand or
csv.DictWriter would raise.
Each metric is now declared once, as a MetricSpec on the trainer that
computes it. MetricsCollector derives the CSV header and W&B payload from
those declarations and owns all accumulation, so train() no longer carries
a running sum, and every isinstance(tr, WGANStageTrainer) branch is gone —
replaced by four trainer hooks (batch_loss, summary, val_objective,
supports_val_loss).
giant/train.py (1875 lines) becomes giant/training/:
trainers.py StageSpec + shared StageTrainer base + the two subclasses
metrics.py MetricSpec, MetricsCollector
stage2_inputs.py the pure AR/teacher-forcing tensor helpers, moved verbatim
loop.py train() (225 lines, was ~514) + graceful shutdown
checkpoint.py build/load, lifted out of train()'s closures
The trainers shared ~15 identical constructor arguments and copy-pasted
their cosine-warmup lambda, EMA setup, state_dict/load_state_dict,
resume_lr and train_mode/eval_mode. StageSpec resolves one stage's config
once (constructors go from 24 and 22 keyword arguments to (spec, model,
device)), the base class holds the rest, and build_stage_trainers drops
from ~100 lines to 15.
Metric columns are renamed to a uniform stage/split/metric scheme
(stage1/train/loss, stage2/train/d_loss, stage1/lr, stage1/router/entropy,
val/loss, ...). Old metrics.csv files and W&B history are not comparable.
The checkpoint format is unchanged.
BEHAVIOR CHANGE — WGAN best-checkpoint selection. The old code meant to
score a WGAN stage on its marginal KL, but the guard
`{n: kl for n in wgan_names if n not in val_loss_per_stage}` could never
fire: val_loss_per_stage was pre-seeded with 0.0 for every stage, so a
WGAN stage contributed a flat 0.0 and the KL was written to metrics.csv
without ever influencing best.pt. val_objective now returns it as
intended. On the test harness's default flow+wgan config val_loss went
from 2.182 (stage 1 only) to 15.137 (stage 1 + KL 12.954), and which epoch
won changed. Runs before this commit picked their best checkpoint on the
non-adversarial stages alone. Written up in docs/v0.3.0-followups.md.
Verified: 699 tests pass; ruff, ruff format and ty clean. Baseline-vs-
refactor metrics.csv compared across five configs (flow+wgan, AR+onehot,
routed, both-flow, AR-flow) — every comparable value bit-identical except
val/loss where the fix applies. Resume appends without a duplicate header
and reproduces a HEAD worktree's per-epoch losses and LRs exactly across
the resume boundary. A refactored last.pt loads through
cli.py:_load_model_weights in both raw and ema modes.
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
69 lines
2.6 KiB
Python
69 lines
2.6 KiB
Python
"""Checkpoint assembly and restore.
|
|
|
|
The on-disk layout is unchanged from v0.2/v0.3.0 and is read by
|
|
`giant/cli.py`, `giant/rollout.py`, `giant/sample.py` and
|
|
`giant/analysis/router_gating.py` — stage 1's weights live under `model`,
|
|
stage 2's under `sec_decoder`, with `_ema`/`critic`/`sec_critic` companions
|
|
and per-stage `optimizer_<stage>` / `optimizer_d_<stage>` / `lr_sched_<stage>`
|
|
entries.
|
|
"""
|
|
|
|
from giant.training.trainers import StageTrainer
|
|
|
|
#: Stage name -> the checkpoint key its weights live under. Historical: stage
|
|
#: 1 predates the two-stage split, so it kept the bare "model" key.
|
|
_STAGE_KEY = {"stage1": "model", "stage2": "sec_decoder"}
|
|
_CRITIC_KEY = {"stage1": "critic", "stage2": "sec_critic"}
|
|
|
|
|
|
def build_checkpoint(
|
|
trainers: dict[str, StageTrainer],
|
|
epoch: int,
|
|
global_step: int,
|
|
best_val_loss: float,
|
|
extras: dict,
|
|
) -> dict:
|
|
"""`extras` carries the dataset-level sidecars (normalizer, vocab maps,
|
|
model_config) that `train()` receives as arguments; `None` values are
|
|
omitted so an absent sidecar leaves no key behind."""
|
|
ckpt: dict = {
|
|
"epoch": epoch,
|
|
"best_val_loss": best_val_loss,
|
|
"global_step": global_step,
|
|
}
|
|
for name, trainer in trainers.items():
|
|
sd = trainer.state_dict()
|
|
key = _STAGE_KEY[name]
|
|
ckpt[key] = sd["model"]
|
|
if "model_ema" in sd:
|
|
ckpt[f"{key}_ema"] = sd["model_ema"]
|
|
if "critic" in sd:
|
|
ckpt[_CRITIC_KEY[name]] = sd["critic"]
|
|
ckpt[f"optimizer_d_{name}"] = sd["optimizer_d"]
|
|
ckpt[f"optimizer_{name}"] = sd["optimizer"]
|
|
ckpt[f"lr_sched_{name}"] = sd["lr_sched"]
|
|
ckpt.update({k: v for k, v in extras.items() if v is not None})
|
|
return ckpt
|
|
|
|
|
|
def load_checkpoint(trainers: dict[str, StageTrainer], ckpt: dict, lr: float) -> None:
|
|
"""Restore every active stage, then hand `lr`'s authority back to the
|
|
config — `load_state_dict` would otherwise leave the checkpoint's own
|
|
base LR in place, silently ignoring `--lr` on resume."""
|
|
for name, trainer in trainers.items():
|
|
key = _STAGE_KEY[name]
|
|
sd = {
|
|
"model": ckpt[key],
|
|
"optimizer": ckpt[f"optimizer_{name}"],
|
|
"lr_sched": ckpt[f"lr_sched_{name}"],
|
|
}
|
|
ema_key = f"{key}_ema"
|
|
if ema_key in ckpt:
|
|
sd["model_ema"] = ckpt[ema_key]
|
|
crit_key = _CRITIC_KEY[name]
|
|
if crit_key in ckpt:
|
|
sd["critic"] = ckpt[crit_key]
|
|
sd["optimizer_d"] = ckpt[f"optimizer_d_{name}"]
|
|
trainer.load_state_dict(sd)
|
|
trainer.resume_lr(lr)
|