Files
giant/scripts/hparam_scan.py
T
lars 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
Refactor train.py into giant/training/ around a metrics collector
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>
2026-08-07 17:03:20 +02:00

182 lines
5.3 KiB
Python

"""Hyperparameter scan over dropout x n_blocks x hidden_dim.
Runs `giant train` sequentially (this machine has a single GPU) for every
combination, plus one extra run at the default architecture with a higher
learning rate. Runs are shuffled so the parameter space gets coarse coverage
early rather than exhausting one corner of the grid first.
See `uv run dwarf hparam-scan --help` for the CLI.
"""
import csv
import itertools
import os
import random
import subprocess
import time
from pathlib import Path
DATA_DEFAULT = "/home/lars/geant_steps/train"
SCAN_DIR_DEFAULT = "checkpoints/scan"
EPOCHS = 50
DEFAULT_LR = 3e-4
DROPOUTS = [0.0, 0.1]
N_BLOCKS = [5, 6, 7, 8]
HIDDEN_DIMS = [256, 512, 1024]
EXTRA_LR_RUN = {"hidden_dim": 256, "n_blocks": 6, "dropout": 0.1, "lr": 1e-3}
_SUMMARY_FIELDS = [
"name",
"hidden_dim",
"n_blocks",
"dropout",
"lr",
"epochs_completed",
"final_val_loss",
"best_val_loss",
"wall_time_s",
]
def build_runs(seed: int) -> list[dict]:
runs = [
{"hidden_dim": h, "n_blocks": n, "dropout": d, "lr": DEFAULT_LR}
for d, n, h in itertools.product(DROPOUTS, N_BLOCKS, HIDDEN_DIMS)
]
runs.append(dict(EXTRA_LR_RUN))
random.Random(seed).shuffle(runs)
return runs
def run_name(run: dict) -> str:
return f"h{run['hidden_dim']}_n{run['n_blocks']}_d{run['dropout']}_lr{run['lr']}"
def last_completed_epoch(metrics_path: Path) -> int:
if not metrics_path.exists():
return 0
with open(metrics_path, newline="") as f:
rows = list(csv.DictReader(f))
if not rows:
return 0
return int(rows[-1]["epoch"])
def final_metrics(metrics_path: Path) -> tuple[int, float, float]:
with open(metrics_path, newline="") as f:
rows = list(csv.DictReader(f))
epochs_completed = int(rows[-1]["epoch"])
final_val_loss = float(rows[-1]["val/loss"])
best_val_loss = min(float(r["val/loss"]) for r in rows)
return epochs_completed, final_val_loss, best_val_loss
def append_summary(summary_path: Path, row: dict) -> None:
write_header = not summary_path.exists()
with open(summary_path, "a", newline="") as f:
writer = csv.DictWriter(f, fieldnames=_SUMMARY_FIELDS)
if write_header:
writer.writeheader()
writer.writerow(row)
def run_hparam_scan(
data: str = DATA_DEFAULT,
scan_dir: str = SCAN_DIR_DEFAULT,
seed: int = 0,
dry_run: bool = False,
) -> None:
runs = build_runs(seed)
scan_dir_path = Path(scan_dir)
if dry_run:
for i, run in enumerate(runs, 1):
print(f"[{i}/{len(runs)}] {run_name(run)}")
return
scan_dir_path.mkdir(parents=True, exist_ok=True)
summary_path = scan_dir_path / "scan_summary.csv"
env = os.environ.copy()
env["TQDM_DISABLE"] = "1"
for i, run in enumerate(runs, 1):
name = run_name(run)
out_dir = scan_dir_path / name
metrics_path = out_dir / "metrics.csv"
last_ckpt = out_dir / "last.pt"
completed = last_completed_epoch(metrics_path)
if completed >= EPOCHS:
print(f"[{i}/{len(runs)}] {name} — already complete, skipping")
continue
out_dir.mkdir(parents=True, exist_ok=True)
cmd = [
"giant",
"train",
data,
"--mode",
"flow",
"--epochs",
str(EPOCHS),
"--batch-size",
"auto",
"--hidden-dim",
str(run["hidden_dim"]),
"--n-blocks",
str(run["n_blocks"]),
"--dropout",
str(run["dropout"]),
"--lr",
str(run["lr"]),
"--out",
str(out_dir),
]
if last_ckpt.exists():
cmd += ["--resume", str(last_ckpt)]
print(f"[{i}/{len(runs)}] {name} — resuming from epoch {completed}")
else:
print(f"[{i}/{len(runs)}] {name} — starting")
start = time.monotonic()
try:
with open(out_dir / "train.log", "a") as log:
subprocess.run(cmd, env=env, stdout=log, stderr=subprocess.STDOUT)
except KeyboardInterrupt:
print(
f"\ninterrupted during {name} — re-run this script to resume "
f"(checkpoint/resume is handled by `giant train` itself)"
)
return
wall_time_s = time.monotonic() - start
if metrics_path.exists():
epochs_completed, final_val_loss, best_val_loss = final_metrics(
metrics_path
)
append_summary(
summary_path,
{
"name": name,
"hidden_dim": run["hidden_dim"],
"n_blocks": run["n_blocks"],
"dropout": run["dropout"],
"lr": run["lr"],
"epochs_completed": epochs_completed,
"final_val_loss": final_val_loss,
"best_val_loss": best_val_loss,
"wall_time_s": round(wall_time_s, 1),
},
)
print(
f"[{i}/{len(runs)}] {name} — val_loss {final_val_loss:.4f} "
f"({wall_time_s:.1f}s)"
)
else:
print(
f"[{i}/{len(runs)}] {name} — no metrics.csv produced, check train.log"
)