Refactor train.py into giant/training/ around a metrics collector
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>
This commit is contained in:
2026-08-07 17:03:20 +02:00
parent da7cde3ef9
commit 8019a80563
18 changed files with 2216 additions and 1944 deletions
+7 -2
View File
@@ -72,7 +72,12 @@ giant/
│ ├── materials.py # material name → (Z_eff, A_eff, density, X0, λ_int)
│ ├── config.py # default hyperparameters, TOML config merging, device autodetect
│ ├── pipeline.py # builds datasets/normalizers and kicks off a training run (with setup-stage caching)
│ ├── train.py # two-stage training loop, checkpointing, graceful shutdown, W&B logging
│ ├── training/ # two-stage training: loop, per-stage trainers, metrics, checkpointing
│ │ ├── loop.py # epoch loop, graceful shutdown, best-checkpoint selection
│ │ ├── trainers.py # StageSpec + flow/ddpm and WGAN-GP per-stage trainers
│ │ ├── stage2_inputs.py# ground-truth stage-2 targets + autoregressive/teacher-forcing inputs
│ │ ├── metrics.py # MetricsCollector: metrics.csv columns, W&B logging, progress/summary
│ │ └── checkpoint.py # checkpoint assembly/restore (format unchanged since v0.2)
│ ├── sample.py # DDPM / DDIM / flow matching / WGAN samplers + secondary sampling
│ ├── geometry.py # GeometryOracle: position → (material, layer_id, escaped) for rollout
│ ├── rollout.py # autoregressive shower rollout driver
@@ -125,7 +130,7 @@ dwarf build-geometry-oracle path/to/steps.parquet --out oracle.pkl
giant rollout path/to/steps.parquet --checkpoint checkpoints/.../best.pt --geometry oracle.pkl
```
`train`/`predict` accept a TOML config file (`--config`) and CLI overrides for hyperparameters; see `--help` on any command for the full option list. `giant train --wandb` logs per-epoch metrics (the same ones written to `metrics.csv`) to Weights & Biases; requires `uv sync --extra wandb`. A repeat `giant train` against the same dataset (e.g. a hyperparameter sweep) reuses a cached setup-stage sidecar (vocab maps, event split, normalizer stats) unless `--no-cache-setup`/`--rebuild-setup-cache`; `dwarf warm-cache` precomputes it ahead of time. `giant rollout` seeds showers from the highest-energy entry step per event, then autoregressively steps the two-stage model to completion — pushing secondaries as new tracks and looking up `material`/`layer_id` from the oracle at each step. Tracks terminate on energy cutoff, per-track max steps, detector escape, or natural end; energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction.
`train`/`predict` accept a TOML config file (`--config`) and CLI overrides for hyperparameters; see `--help` on any command for the full option list. `giant train --wandb` logs per-epoch metrics (the same ones written to `metrics.csv`) to Weights & Biases; requires `uv sync --extra wandb`. Metric names are `<stage>/<split>/<metric>` (e.g. `stage1/train/loss`, `stage2/train/d_loss`, `stage1/lr`), plus an unprefixed run-level tail (`val/loss`, `grad_norm`, `epoch_time_s`, ...); each is declared once on the `StageTrainer` that computes it (`giant/training/trainers.py`), so the CSV header and the W&B panel names are derived, never hand-maintained. A repeat `giant train` against the same dataset (e.g. a hyperparameter sweep) reuses a cached setup-stage sidecar (vocab maps, event split, normalizer stats) unless `--no-cache-setup`/`--rebuild-setup-cache`; `dwarf warm-cache` precomputes it ahead of time. `giant rollout` seeds showers from the highest-energy entry step per event, then autoregressively steps the two-stage model to completion — pushing secondaries as new tracks and looking up `material`/`layer_id` from the oracle at each step. Tracks terminate on energy cutoff, per-track max steps, detector escape, or natural end; energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction.
## Validation and analysis
+14 -7
View File
@@ -322,7 +322,7 @@ stage1_context = "truth"
| `lambda` | float | `1.0` | Weight of stage 2's loss in the total. Was `train.lambda_s2`. |
| `k_max` | int | `15` | Maximum secondary slots. Under `"one_shot"` this is the fixed output width; under `"autoregressive"` it is a safety cap on the generation loop. Was the global constant `K_MAX` in `giant/constants.py` (max observed `n_sec` is 14 in the PbWO4 dataset, so 15 covers it with one spare). |
| `context_dim` | int | `64` | Width of the projected stage-1 outcome fed into stage 2's conditioning. Was the hardcoded `stage1_proj_dim = 64`. |
| `stage1_context` | `"truth"` \| `"sampled"` | `"truth"` | What stage 2 conditions on during training. `"truth"`: the ground-truth stage-1 target vector, detached — v0.2 behaviour (`train.py:254` passes `x1_s1.detach()`), i.e. stage-level teacher forcing. `"sampled"`: stage 1's own sampled output, closing the train/inference gap at the cost of a sampling pass per batch and a moving target early in training. |
| `stage1_context` | `"truth"` \| `"sampled"` | `"truth"` | What stage 2 conditions on during training. `"truth"`: the ground-truth stage-1 target vector, detached — v0.2 behaviour (v0.2 `train.py:254` passes `x1_s1.detach()`), i.e. stage-level teacher forcing. `"sampled"`: stage 1's own sampled output, closing the train/inference gap at the cost of a sampling pass per batch and a moving target early in training. |
#### `[stage2_model.n_sec]`
@@ -748,7 +748,7 @@ configuration that plausibly meets the budget; flow AR is for quality comparison
## 7. Training loop
Decision 2 (full mixed per-stage objectives) makes `train.py` one trainer object
Decision 2 (full mixed per-stage objectives) makes `giant/training/` one trainer object
per active stage:
```python
@@ -765,8 +765,15 @@ FlowTrainer, DDPMTrainer, WGANTrainer(critic, n_critic, gp_weight)
critic updates then a generator update. Independent optimizers, independent EMA.
- Router auxiliary losses (`lambda_balance` / `lambda_proc` / `lambda_entropy`)
become per-stage, summed over whichever stages are routed. Today's
`hasattr(model, "router")` check (`train.py:262`) generalizes cleanly.
- `metrics.csv` and W&B metric names gain a stage prefix.
`hasattr(model, "router")` check (v0.2 `train.py:262`) generalizes cleanly.
- `metrics.csv` and W&B metric names gain a stage prefix. **Implemented as**
`<stage>/<split>/<metric>` (`stage1/train/loss`, `stage2/train/d_loss`,
`stage1/lr`, `stage1/router/entropy`) plus an unprefixed run-level tail
(`val/loss`, `val/marginal_kl`, `grad_norm`, `gpu_mem_mb`,
`samples_per_sec`, `is_best`, `epoch_time_s`). Each name is declared once,
as a `MetricSpec` on the `StageTrainer` that computes it; the CSV header
and W&B payload are derived from those declarations
(`giant/training/metrics.py`).
- With `stage1_model.active = false`, stage 2 still needs its stage-1 context: it
comes from the ground-truth target already in the batch (`x1_s1`), which is
exactly what v0.2 does anyway. **Stage-2-only training is therefore a cheap
@@ -855,7 +862,7 @@ All in `giant/config.py`:
| File | Why |
|------|-----|
| `giant/pipeline.py` | Builds `model_config`; now per-stage. Delete the wgan+router rejection at `:275`. Router `centers_init` seeding becomes per-stage. |
| `giant/train.py` | Per-stage trainers (§7). |
| `giant/training/` | Per-stage trainers (§7), metric collection, checkpointing, the epoch loop. |
| `giant/cli.py` | Stage-prefixed flags for `train` and `new-run`; `build_models` now returns a dict (`:871`, `:1254`); `:1339` writes `model_config` into the rollout sidecar. |
| `giant/sample.py` | Sampler picked per stage from `stage*_model.generator`; new AR sampling loop with KV cache under `history = "attention"`. |
| `giant/rollout.py` | AR secondary generation; categorical class -> PDG decode; `other_policy` handling. |
@@ -894,7 +901,7 @@ All in `giant/config.py`:
later is a config addition, not a break.
- **`stage2_model.generator = "ddpm"`.** The value is **accepted by the
schema** (§3.3 lists `"flow" | "ddpm" | "wgan"` with no caveat) but
`FlowDDPMStageTrainer.__init__` (`giant/train.py`) raises
`FlowDDPMStageTrainer.__init__` (`giant/training/trainers.py`) raises
`NotImplementedError` for stage 2 — only `"flow"` and `"wgan"` have a
stage-2 secondary-decoder loss implemented. `stage1_model.generator =
"ddpm"` is unaffected; this restriction is stage-2-only. Landing stage-2
@@ -968,7 +975,7 @@ or a non-adversarial CE head, both of which already exist as config options.
2. **`network.py`** — the §5 decomposition, with `Stage2OneShot` reproducing v0.2
exactly. Gate on the §4.3 migration test: load a v0.2 checkpoint through the
shim and diff outputs against v0.2 code.
3. **`train.py`** — per-stage trainers; `active = false` paths. At this point
3. **`giant/training/`** — per-stage trainers; `active = false` paths. At this point
Stage-2-only training works and the meeting's step 2 (one-shot WGAN baseline,
trained standalone) is runnable.
4. **Type map** — `loader.py` + `setup_cache.py` + checkpoint persistence +
+17 -1
View File
@@ -30,8 +30,24 @@ turned out fine:
- `analysis/render.py`/`analysis/router_gating.py` correctly branch
old-flat vs new-nested `model_config["router"]` location — doc flagged this
as a likely stale spot (§10) but it's actually fine.
- `train.py`'s per-stage trainers, mixed flow+wgan runs, WGAN critic cadence,
- `giant/training/`'s per-stage trainers, mixed flow+wgan runs, WGAN critic cadence,
per-stage router auxiliary losses, stage-prefixed metrics, stage-2-only
training via ground-truth `x1_s1` (§7).
## Behavior change: WGAN best-checkpoint selection
The `giant/training/` split fixed a dead guard in WGAN validation scoring. A
WGAN stage was *meant* to contribute its marginal KL to the `val_loss` that
drives `best.pt`, but the guard `if n not in val_loss_per_stage` could never
fire (every stage was pre-seeded to `0.0`), so the stage contributed a flat
`0.0` and the KL was written to `metrics.csv` without ever being used.
`WGANStageTrainer.val_objective` now returns the KL as intended.
**Consequence:** any checkpoint selected before this commit under a config
with a WGAN stage — including the v0.3.0 default (`stage2_model.generator =
"wgan"`) — picked its best epoch on the non-adversarial stages alone. Measured
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. Do not compare `val/loss` or `best.pt` choice across this commit.
- `pipeline.py`'s deleted wgan+router rejection, per-stage `centers_init`
seeding, removed stale expert-size warning (§9, §10).
+2 -2
View File
@@ -316,7 +316,7 @@ def _load_model_weights(
) -> None:
"""Load either the raw or EMA state dicts from a training checkpoint.
EMA weights (giant.train's shadow copy, see --ema-decay) only exist in
EMA weights (giant.training's shadow copy, see --ema-decay) only exist in
checkpoints written after that feature landed, so `ema` fails loudly
rather than silently falling back to raw weights a caller didn't ask for.
"""
@@ -877,7 +877,7 @@ def train(
# Name only encodes what's non-default (see default_out_dir_name), so
# two runs with identical hyperparams in the same to-the-minute
# timestamp would otherwise collide on this name — which also
# doubles as the W&B run id (giant.train) — hence the suffix loop in
# doubles as the W&B run id (giant.training) — hence the suffix loop in
# resolve_default_out_dir.
out_dir = gconfig.resolve_default_out_dir(cfg)
+1 -1
View File
@@ -894,7 +894,7 @@ def default_out_dir_name(cfg: dict, now: datetime | None = None) -> str:
Beyond `_OUT_DIR_NAME_MAX_FIELDS` non-default fields, the remainder
collapse into a short deterministic hash suffix rather than growing the
name unboundedly. This name doubles as the run's W&B id (see
giant.train), which is the reason a timestamp is always included.
giant.training), which is the reason a timestamp is always included.
"""
now = now or datetime.now()
tokens = []
+1 -1
View File
@@ -886,7 +886,7 @@ class AttentionHistory(HistoryEncoder):
`Stage2Autoregressive._token_cond` use: `feat[:, i]` is token `i - 1`'s
own `(energy_fraction, direction, type_representation)`, with a learned
start vector substituted at `has_prev == False` positions (only slot 0 in
practice see `giant.train._ar_has_prev`). Causal masking then makes
practice see `giant.training.stage2_inputs._ar_has_prev`). Causal masking then makes
position `i`'s output a function of `feat[:, 1:i+1]` — i.e. tokens
`0..i-1` exactly the prefix available when predicting token `i`.
+1 -1
View File
@@ -159,7 +159,7 @@ def flow_matching_loss_secondary_ar(
Under teacher forcing (docs/v0.3.0-design.md §6.2 point 3) this is still a
single parallel pass over all K_MAX tokens `x1`/`history_feat`/etc. are
already built from ground truth for every slot by the caller
(`giant.train._assemble_stage2_ar_inputs`/`_assemble_stage2_ar_target`).
(`giant.training.stage2_inputs._assemble_stage2_ar_inputs`/`_assemble_stage2_ar_target`).
x1: (B, K_MAX, CONT_SLOT_DIM + type_dim) per-token flattened target
(stick_logit, dir, then a `type_dim`-wide type slice)
+1 -1
View File
@@ -32,7 +32,7 @@ from giant.data.transforms import (
)
from giant.data.dataset import make_event_split, StreamingStepsDataset
from giant.model.network import build_models, build_critics
from giant.train import train as run_training
from giant.training import train as run_training
@dataclass
+1 -1
View File
@@ -245,7 +245,7 @@ def sample_secondaries_ar(
one token at a time, in descending-energy slot order, `k_max` sequential
calls. Unlike training (teacher forcing, §6.2 point 3 a single
parallel pass over ground-truth tokens, see
`giant.train._assemble_stage2_ar_inputs`), there is no ground truth at
`giant.training.stage2_inputs._assemble_stage2_ar_inputs`), there is no ground truth at
inference: each token's conditioning is built free-running, from the
PREVIOUS TOKEN'S OWN just-generated output — the train/inference gap
§6.2 point 4 explicitly flags as the cost of markov history's
-1875
View File
File diff suppressed because it is too large Load Diff
+30
View File
@@ -0,0 +1,30 @@
"""Training: per-stage trainers, metric collection, checkpointing, the loop.
Split out of the former single-module `giant/train.py`. The public surface is
`train` (the entry point `giant.pipeline` calls) plus the trainer/spec types
that tests and tooling construct directly.
"""
from giant.training.checkpoint import build_checkpoint, load_checkpoint
from giant.training.metrics import MetricsCollector, MetricSpec
from giant.training.loop import train
from giant.training.trainers import (
FlowDDPMStageTrainer,
StageSpec,
StageTrainer,
WGANStageTrainer,
build_stage_trainers,
)
__all__ = [
"FlowDDPMStageTrainer",
"MetricSpec",
"MetricsCollector",
"StageSpec",
"StageTrainer",
"WGANStageTrainer",
"build_checkpoint",
"build_stage_trainers",
"load_checkpoint",
"train",
]
+68
View File
@@ -0,0 +1,68 @@
"""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)
+328
View File
@@ -0,0 +1,328 @@
"""The training loop.
`train()` owns the epoch structure and nothing else: the per-stage step is
`giant.training.trainers`' job, every number reported is
`giant.training.metrics`' job, and the on-disk checkpoint is
`giant.training.checkpoint`'s.
"""
import os
import signal
import time
from pathlib import Path
from types import FrameType
from typing import Callable
import numpy as np
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm
from giant.data.loader import TopNMap
from giant.data.setup_cache import topnmap_to_json
from giant.training.checkpoint import build_checkpoint, load_checkpoint
from giant.training.metrics import MetricsCollector
from giant.training.trainers import (
FlowDDPMStageTrainer,
StageTrainer,
build_stage_trainers,
)
from giant.validate import validate_marginals
_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 _try_validate_marginals(trainer: StageTrainer, val_loader, device, **kwargs):
"""Runs `validate_marginals` on `trainer`'s sampling model (EMA model if
present, else the raw model). `validate_marginals` itself dispatches
through `giant.sample.sample_stage1`/`sample_stage2`/`resolve_n_sec`
(docs/v0.3.0-design.md §10), so this is generator- and
one-shot-vs-autoregressive-agnostic."""
model = trainer.sampling_model()
return validate_marginals(model, val_loader, device=device, **kwargs)
def _marginal_kl(
trainers: dict[str, StageTrainer], val_loader, device, **kwargs
) -> float:
"""Mean marginal KL over the stage-1 sampling chain, or NaN when stage 1
is inactive or `validate_marginals` declined to produce a result."""
stage1 = trainers.get("stage1")
if stage1 is None:
return float("nan")
result = _try_validate_marginals(
stage1,
val_loader,
device,
sec_decoder=trainers["stage2"].sampling_model()
if "stage2" in trainers
else None,
**kwargs,
)
if result is None:
return float("nan")
return float(np.mean(result["kl_divergence"]))
def train(
cfg: dict,
models: dict[str, torch.nn.Module | None],
critics: dict[str, torch.nn.Module | None],
train_loader: DataLoader,
val_loader: DataLoader,
device: torch.device,
out_dir: str | Path,
normalizer_dict: dict | None = None,
pdg_map: dict | None = None,
mat_map: dict | None = None,
proc_map: dict | None = None,
pdg_topn_map: TopNMap | None = None,
mat_topn_map: TopNMap | None = None,
model_config: dict | None = None,
resume_path: str | Path | None = None,
total_train_batches: int = 0,
use_wandb: bool = False,
wandb_project: str = "giant",
wandb_run_name: str = "",
wandb_log_every: int = 50,
) -> None:
"""Train whichever of stage1/stage2 are active, each through its own
`StageTrainer` (design doc §7). `models`/`critics` are the dicts
`giant.model.network.build_models`/`build_critics` return a `None`
entry means that stage is `active = false`.
"""
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
t = cfg["train"]
epochs = t["epochs"]
validate_every = t.get("validate_every", 0)
validate_steps = t.get("validate_steps", 10)
max_val_batches = t.get("max_val_batches", 0)
trainers = build_stage_trainers(cfg, models, critics, device, total_train_batches)
if not trainers:
raise ValueError(
"no active stage — stage1_model.active and stage2_model.active are both false"
)
has_adversarial = any(not tr.supports_val_loss for tr in trainers.values())
checkpoint_extras = {
"normalizer": normalizer_dict,
"pdg_map": pdg_map,
"mat_map": mat_map,
"proc_map": proc_map,
"pdg_topn_map": topnmap_to_json(pdg_topn_map)
if pdg_topn_map is not None
else None,
"mat_topn_map": topnmap_to_json(mat_topn_map)
if mat_topn_map is not None
else None,
"model_config": model_config,
}
start_epoch = 1
best_val_loss = float("inf")
global_step = 0
if resume_path is not None:
ckpt = torch.load(resume_path, map_location=device, weights_only=False)
load_checkpoint(trainers, ckpt, t["lr"])
start_epoch = ckpt.get("epoch", 0) + 1
best_val_loss = ckpt.get("best_val_loss", float("inf"))
global_step = ckpt.get("global_step", 0)
if start_epoch > epochs:
print(
f"checkpoint already completed epoch {start_epoch - 1} "
f"(>= --epochs {epochs}) — nothing to train"
)
return
collector = MetricsCollector.create(
trainers,
out_dir,
cfg,
model_config,
resume=resume_path is not None,
use_wandb=use_wandb,
wandb_project=wandb_project,
wandb_run_name=wandb_run_name,
wandb_log_every=wandb_log_every,
)
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()
if device.type == "cuda":
torch.cuda.reset_peak_memory_stats(device)
collector.start_epoch(epoch)
for trainer in trainers.values():
trainer.train_mode()
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 batch in bar:
B = batch[0].size(0)
collector.add_train_batch(
{
name: trainer.step(batch, device, global_step)
for name, trainer in trainers.items()
},
B,
)
bar.set_postfix_str(collector.postfix(), refresh=False)
global_step += 1
collector.log_batch(global_step, batch, device)
if shutdown.requested:
break
bar.close()
if shutdown.requested:
ckpt = build_checkpoint(
trainers, epoch - 1, global_step, best_val_loss, checkpoint_extras
)
torch.save(ckpt, out_dir / "last.pt")
last_completed_epoch = epoch - 1
print(
f"saved in-progress weights from partway through epoch "
f"{epoch} to {out_dir / 'last.pt'} "
f"(resume will restart epoch {epoch})"
)
break
for trainer in trainers.values():
trainer.eval_mode()
# --- per-stage validation ---
scored = {name: tr for name, tr in trainers.items() if tr.supports_val_loss}
if scored:
with torch.no_grad():
for val_batch_idx, batch in enumerate(val_loader):
if max_val_batches > 0 and val_batch_idx >= max_val_batches:
break
B = batch[0].size(0)
collector.add_val_batch(
{
name: tr.val_loss(batch, device)
for name, tr in scored.items()
},
B,
)
collector.observe_routers(
batch[0].to(device), batch[1].to(device), B
)
# An adversarial stage has no averageable validation loss, so it
# needs the marginal-KL signal every epoch to pick a best
# checkpoint at all; a purely non-adversarial run only pays for
# it every `validate_every` epochs.
marginal_kl = float("nan")
if has_adversarial:
marginal_kl = _marginal_kl(trainers, val_loader, device)
elif validate_every > 0 and epoch % validate_every == 0:
stage1 = trainers.get("stage1")
ddpm_steps = 1000
if (
isinstance(stage1, FlowDDPMStageTrainer)
and stage1.ddpm_schedule is not None
):
ddpm_steps = stage1.ddpm_schedule.T
marginal_kl = _marginal_kl(
trainers,
val_loader,
device,
steps=validate_steps,
ddpm_steps=ddpm_steps,
)
val_loss = sum(
trainer.val_objective(
collector.train_means(name),
collector.val_means(name),
marginal_kl,
)
for name, trainer in trainers.items()
)
epoch_time = time.monotonic() - epoch_start
is_best = val_loss < best_val_loss
collector.set("val/loss", val_loss)
collector.set("val/marginal_kl", marginal_kl)
collector.set(
"gpu_mem_mb",
torch.cuda.max_memory_allocated(device) / (1024 * 1024)
if device.type == "cuda"
else 0.0,
)
collector.set(
"samples_per_sec", collector.train_samples / max(epoch_time, 1e-8)
)
collector.set("is_best", int(is_best))
collector.set("epoch_time_s", epoch_time)
print(collector.summary_line(val_loss, epoch_time, is_best))
collector.write_epoch(global_step)
ckpt = build_checkpoint(
trainers, epoch, global_step, best_val_loss, checkpoint_extras
)
if is_best:
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
collector.close()
if shutdown.requested:
print(
f"stopped after epoch {last_completed_epoch} due to shutdown signal — "
f"resume with --resume {out_dir / 'last.pt'}"
)
+405
View File
@@ -0,0 +1,405 @@
"""Per-epoch metric accumulation, `metrics.csv`, and W&B logging.
Every scalar a training run reports is declared exactly once, as a
`MetricSpec` on the `StageTrainer` that computes it (see
`giant.training.trainers`). `MetricsCollector` derives the CSV/W&B column set
from those declarations, so adding a metric means adding one line next to the
code that produces it there is no second list to keep in sync.
Column naming is uniform: `<stage>/train/<key>`, `<stage>/val/<key>`,
`<stage>/<key>` for point-in-time values (`lr`, `critic_lr`),
`<stage>/router/<key>` for routing diagnostics, and an unprefixed run-level
tail (`val/loss`, `grad_norm`, `epoch_time_s`, ...). W&B groups panels on
`/`, so the same names read well there.
"""
import csv
from dataclasses import dataclass
from pathlib import Path
import torch
_ROUTER_KEYS = ("entropy", "util_min", "util_max", "util_std")
# Written after every stage's columns, by `MetricsCollector` itself rather
# than by any one trainer — these describe the run, not a stage.
_RUN_COLUMNS = (
"val/loss",
"val/marginal_kl",
"grad_norm",
"gpu_mem_mb",
"samples_per_sec",
"is_best",
"epoch_time_s",
)
# tqdm/W&B batch-granularity smoothing, matching v0.2/v0.3.0's inline EMA.
_EMA_ALPHA = 0.05
@dataclass(frozen=True)
class MetricSpec:
"""One scalar a trainer emits per batch, and how it is reported.
`key` indexes the dict `StageTrainer.step()` / `.val_loss()` returns;
`column` is the CSV/W&B column suffix, joined to the stage name with
"/". `reduce` is either "mean" (batch-size-weighted average over the
epoch) or "last" (the most recent value for point-in-time quantities
like the learning rate, which is a schedule readout, not a statistic).
"""
key: str
column: str
reduce: str = "mean"
def train_metric(key: str, column: str | None = None) -> MetricSpec:
return MetricSpec(key, column or f"train/{key}")
def val_metric(key: str, column: str | None = None) -> MetricSpec:
return MetricSpec(key, column or f"val/{key}")
def stage_metric(key: str, column: str | None = None) -> MetricSpec:
"""A point-in-time stage-level readout (`lr`, `critic_lr`) — reported
unprefixed by split, as `<stage>/<key>`."""
return MetricSpec(key, column or key, reduce="last")
def _wandb_run_config(cfg: dict, model_config: dict | None, param_counts: dict) -> dict:
return {
"train": cfg["train"],
"conditioning": cfg["conditioning"],
"stage1_model": cfg["stage1_model"],
"stage2_model": cfg["stage2_model"],
"model_config": model_config or {},
"param_counts": param_counts,
}
class _Accumulator:
"""Batch-size-weighted sums for one stage and one split."""
def __init__(self) -> None:
self.sums: dict[str, float] = {}
self.n = 0
self.last: dict[str, float] = {}
def add(self, stats: dict, keys: set[str], batch_size: int) -> None:
for key in keys:
if key in stats:
self.sums[key] = self.sums.get(key, 0.0) + stats[key] * batch_size
self.last.update(stats)
self.n += batch_size
def mean(self, key: str) -> float:
return self.sums.get(key, 0.0) / max(self.n, 1)
def means(self) -> dict[str, float]:
return {key: self.mean(key) for key in self.sums}
def reset(self) -> None:
self.sums.clear()
self.last.clear()
self.n = 0
class _RouterAccumulator:
"""Gate-diagnostic sums for one routed stage."""
def __init__(self, n_experts: int) -> None:
self.n_experts = n_experts
self.entropy = 0.0
self.importance: torch.Tensor | None = None
self.n = 0
def add(self, entropy: torch.Tensor, importance: torch.Tensor, n: int) -> None:
self.entropy += entropy.item() * n
self.importance = (
importance.clone()
if self.importance is None
else self.importance + importance
)
self.n += n
def stats(self) -> dict[str, float]:
if self.importance is None or self.n == 0:
return dict.fromkeys(_ROUTER_KEYS, 0.0)
util = self.importance / self.importance.sum().clamp_min(1e-8)
return {
"entropy": self.entropy / self.n,
"util_min": util.min().item(),
"util_max": util.max().item(),
"util_std": util.std().item() if self.n_experts > 1 else 0.0,
}
def reset(self) -> None:
self.entropy = 0.0
self.importance = None
self.n = 0
class MetricsCollector:
"""Owns every number a training run reports.
Accumulates per-batch stats from each stage, writes one `metrics.csv` row
per epoch, mirrors it to W&B, and formats the tqdm postfix and the epoch
summary line so `giant.training.loop.train` never carries a running
sum, a column name, or a W&B call of its own.
"""
def __init__(
self,
trainers: dict,
out_dir: Path,
*,
epochs: int,
resume: bool = False,
wandb_run=None,
wandb_log_every: int = 50,
) -> None:
self.trainers = trainers
self.epochs = epochs
self.wandb_run = wandb_run
self.wandb_log_every = wandb_log_every
self.epoch_width = len(str(epochs))
self._train = {name: _Accumulator() for name in trainers}
self._val = {name: _Accumulator() for name in trainers}
self._routers = {
name: _RouterAccumulator(tr.router.n_experts)
for name, tr in trainers.items()
if tr.router is not None
}
# Only "mean" specs need summing; "last" specs are read straight off
# the accumulator's most recent stats dict. "grad_norm" is always
# summed — it feeds the run-level `grad_norm` column whether or not
# a trainer reports it per stage.
self._train_keys = {
name: {spec.key for spec in tr.train_metrics if spec.reduce == "mean"}
| {"grad_norm"}
for name, tr in trainers.items()
}
self._val_keys = {
name: {spec.key for spec in tr.val_metrics if spec.reduce == "mean"}
for name, tr in trainers.items()
}
self._run_values: dict[str, float] = {}
self._epoch = 0
self._ema_loss = 0.0
self._ema_grad_norm = 0.0
self._ema_seeded = False
self._batch_loss = 0.0
self._batch_grad_norm = 0.0
self.fieldnames = self._build_fieldnames()
metrics_path = out_dir / "metrics.csv"
append = resume and metrics_path.exists()
self._file = open(metrics_path, "a" if append else "w", newline="")
self._writer = csv.DictWriter(self._file, fieldnames=self.fieldnames)
if not append:
self._writer.writeheader()
# --- construction ---------------------------------------------------
@classmethod
def create(
cls,
trainers: dict,
out_dir: Path,
cfg: dict,
model_config: dict | None,
*,
resume: bool = False,
use_wandb: bool = False,
wandb_project: str = "giant",
wandb_run_name: str = "",
wandb_log_every: int = 50,
) -> "MetricsCollector":
"""Build the collector, starting a W&B run first when enabled."""
wandb_run = None
if use_wandb:
try:
import wandb
except ImportError as exc:
raise RuntimeError(
"train.wandb = true (--wandb) requires the 'wandb' package — "
"install it via `uv sync --extra wandb`"
) from exc
param_counts = {
name: sum(p.numel() for p in tr.model.parameters())
for name, tr in trainers.items()
}
param_counts["total"] = sum(param_counts.values())
wandb_run = wandb.init(
project=wandb_project,
name=wandb_run_name or out_dir.name,
id=out_dir.name,
resume="allow",
config=_wandb_run_config(cfg, model_config, param_counts),
)
return cls(
trainers,
out_dir,
epochs=cfg["train"]["epochs"],
resume=resume,
wandb_run=wandb_run,
wandb_log_every=wandb_log_every,
)
def _build_fieldnames(self) -> list[str]:
fields = ["epoch"]
for name, trainer in self.trainers.items():
for spec in trainer.train_metrics:
fields.append(f"{name}/{spec.column}")
for spec in trainer.val_metrics:
fields.append(f"{name}/{spec.column}")
if trainer.router is not None:
fields += [f"{name}/router/{key}" for key in _ROUTER_KEYS]
for spec in trainer.stage_metrics:
fields.append(f"{name}/{spec.column}")
fields += list(_RUN_COLUMNS)
return fields
def close(self) -> None:
self._file.close()
if self.wandb_run is not None:
self.wandb_run.finish()
# --- per-batch ------------------------------------------------------
def start_epoch(self, epoch: int) -> None:
self._epoch = epoch
for acc in self._train.values():
acc.reset()
for acc in self._val.values():
acc.reset()
for acc in self._routers.values():
acc.reset()
self._run_values.clear()
self._ema_seeded = False
def add_train_batch(self, stats: dict[str, dict], batch_size: int) -> None:
"""`stats` maps stage name -> the dict that stage's `step()` returned."""
self._batch_loss = 0.0
self._batch_grad_norm = 0.0
for name, stage_stats in stats.items():
self._train[name].add(stage_stats, self._train_keys[name], batch_size)
self._batch_loss += self.trainers[name].batch_loss(stage_stats)
self._batch_grad_norm += stage_stats.get("grad_norm", 0.0)
if self._ema_seeded:
self._ema_loss += _EMA_ALPHA * (self._batch_loss - self._ema_loss)
self._ema_grad_norm += _EMA_ALPHA * (
self._batch_grad_norm - self._ema_grad_norm
)
else:
self._ema_loss = self._batch_loss
self._ema_grad_norm = self._batch_grad_norm
self._ema_seeded = True
def add_val_batch(self, stats: dict[str, dict], batch_size: int) -> None:
for name, stage_stats in stats.items():
self._val[name].add(stage_stats, self._val_keys[name], batch_size)
@torch.no_grad()
def observe_routers(
self, cond_cont: torch.Tensor, cond_cat: torch.Tensor, batch_size: int
) -> None:
"""Record gate diagnostics for every routed stage on this batch.
Called from the validation pass only (as in v0.2/v0.3.0), so a stage
whose trainer has no validation pass i.e. WGAN reports zeros.
"""
for name, acc in self._routers.items():
router = self.trainers[name].router
entropy, importance = router.gate_stats(cond_cont, cond_cat)
acc.add(entropy, importance, batch_size)
def postfix(self) -> str:
"""tqdm postfix for the training bar."""
return f"loss={self._ema_loss:.4f} gnorm={self._ema_grad_norm:.3f}"
def log_batch(self, global_step: int, batch: tuple, device: torch.device) -> None:
"""Batch-granularity W&B log, throttled to every `wandb_log_every`
optimizer steps (a single epoch can be tens of thousands). `batch` is
the raw training batch, needed only to re-derive routing entropy for
routed stages it is never moved to `device` otherwise."""
if self.wandb_run is None or self.wandb_log_every <= 0:
return
if global_step % self.wandb_log_every != 0:
return
payload = {
"batch/epoch": self._epoch,
"batch/loss": self._batch_loss,
"batch/loss_ema": self._ema_loss,
"batch/grad_norm": self._batch_grad_norm,
}
for name, trainer in self.trainers.items():
payload[f"batch/{name}/lr"] = trainer.optimizer.param_groups[0]["lr"]
if trainer.router is not None:
with torch.no_grad():
entropy, _ = trainer.router.gate_stats(
batch[0].to(device), batch[1].to(device)
)
payload[f"batch/{name}/router/entropy"] = entropy.item()
self.wandb_run.log(payload, step=global_step)
# --- per-epoch ------------------------------------------------------
def train_means(self, stage: str) -> dict[str, float]:
return self._train[stage].means()
def val_means(self, stage: str) -> dict[str, float]:
return self._val[stage].means()
@property
def train_samples(self) -> int:
"""Samples seen this epoch — identical across stages (every stage
steps on every batch), so any one accumulator's count will do."""
return max((acc.n for acc in self._train.values()), default=0)
def set(self, column: str, value: float) -> None:
"""Record a run-level value for this epoch's row (`val/loss`,
`gpu_mem_mb`, ...). Must name a column in `_RUN_COLUMNS`."""
if column not in _RUN_COLUMNS:
raise KeyError(f"{column!r} is not a run-level metrics column")
self._run_values[column] = value
def summary_line(self, val_loss: float, epoch_time: float, is_best: bool) -> str:
bits = [
trainer.summary(self.train_means(name))
for name, trainer in self.trainers.items()
]
marker = " [best]" if is_best else ""
return (
f"epoch {self._epoch:{self.epoch_width}d}/{self.epochs} "
+ " ".join(bits)
+ f" val {val_loss:.4f} {epoch_time:.1f}s{marker}"
)
def write_epoch(self, global_step: int) -> None:
"""Assemble, write, and flush this epoch's row; mirror it to W&B."""
row: dict = {"epoch": self._epoch}
grad_norm_total = 0.0
for name, trainer in self.trainers.items():
train_acc, val_acc = self._train[name], self._val[name]
for spec in trainer.train_metrics:
row[f"{name}/{spec.column}"] = train_acc.mean(spec.key)
for spec in trainer.val_metrics:
row[f"{name}/{spec.column}"] = val_acc.mean(spec.key)
if trainer.router is not None:
for key, value in self._routers[name].stats().items():
row[f"{name}/router/{key}"] = value
for spec in trainer.stage_metrics:
row[f"{name}/{spec.column}"] = train_acc.last.get(spec.key, 0.0)
grad_norm_total += train_acc.mean("grad_norm")
for column in _RUN_COLUMNS:
row[column] = self._run_values.get(column, float("nan"))
row["grad_norm"] = grad_norm_total
self._writer.writerow(row)
self._file.flush()
if self.wandb_run is not None:
self.wandb_run.log(row, step=global_step)
+321
View File
@@ -0,0 +1,321 @@
"""Ground-truth tensor assembly for stage-2 training.
Pure functions, no optimizer/model state: they turn a batch's ground-truth
secondary tensors into the per-token targets and autoregressive conditioning
inputs `giant.training.trainers` feeds to `Stage2OneShot` /
`Stage2Autoregressive`. Split out of the trainers so the (target, generator,
decoder) width rules the fiddliest part of docs/v0.3.0-design.md §2.1/§6
live in one place and stay unit-testable on their own.
"""
import torch
import torch.nn.functional as F
from giant.constants import CONT_SLOT_DIM, PARTICLE_PHYS_DIM
from giant.sample import sample_secondaries_ar
def _gumbel_tau(step: int, total_steps: int, tau_start: float, tau_end: float) -> float:
"""Linear anneal of the straight-through Gumbel-softmax temperature.
Deterministic in `step`/`total_steps` alone (no extra state), so it
recomputes correctly on `--resume` from a checkpoint's saved `global_step`
without needing to persist anything new (see
giant.model.network.Router.combine_weights).
"""
progress = min(step / max(total_steps, 1), 1.0)
return tau_start + (tau_end - tau_start) * progress
def _type_repr(
sec_type_idx: torch.Tensor,
sec_cont: torch.Tensor,
particle_type_cfg: dict,
cond_enc: torch.nn.Module,
emb_dim: int,
) -> torch.Tensor:
"""(B, K_MAX, type_dim) ground-truth type representation, generator-
independent (unlike `_assemble_stage2_ar_target`'s training *target*,
which varies by generator/objective see its docstring): `"physical"` ->
`(log_mass, charge)`; `"onehot"` -> one-hot of the true class;
`"embedding"` -> the conditioning's own detached embedding-table row.
Used both to build `_assemble_stage2_ar_target`'s wgan+onehot/embedding
branch and as the AR history features' previous-secondary identity — the
latter must always reflect the true physical secondary that came before,
regardless of what the *current* token's own training objective is.
"""
target = particle_type_cfg.get("target", "physical")
if target == "physical":
return sec_cont[..., CONT_SLOT_DIM : CONT_SLOT_DIM + PARTICLE_PHYS_DIM]
if target == "onehot":
return F.one_hot(sec_type_idx, num_classes=emb_dim).float()
return cond_enc.pdg_emb(sec_type_idx).detach()
def _assemble_stage2_ar_target(
sec_cont: torch.Tensor,
sec_type_idx: torch.Tensor,
particle_type_cfg: dict,
generator: str,
cond_enc: torch.nn.Module,
emb_dim: int,
) -> torch.Tensor:
"""(B, K_MAX, token_dim) ground-truth per-token target — the unflattened
analogue of `_assemble_stage2_real` (defined below in terms of this),
matching whatever width `Stage2Autoregressive`'s (or `Stage2OneShot`'s)
own trunk produces for this (target, generator) combination
(`giant.model.network.stage2_trunk_sec_dim`; docs/v0.3.0-design.md
decision 2/3):
- `target = "physical"`: unchanged from v0.2 `sec_cont` (stick_logit,
dir, log_mass, charge) as-is.
- `target` in `("onehot", "embedding")` + `generator in ("flow", "ddpm")`:
just the continuous stick/dir slots the type slice isn't part of
this tensor at all (`type_head` handles it separately).
- `target` in `("onehot", "embedding")` + `generator == "wgan"`: stick/dir
slots concatenated with the per-slot type representation (a one-hot of
the true class, relaxed on the *generated* side only, by the caller;
or the conditioning's own detached embedding-table row).
"""
target = particle_type_cfg.get("target", "physical")
if target == "physical":
return sec_cont
cont = sec_cont[..., :CONT_SLOT_DIM]
if generator != "wgan":
return cont
type_repr = _type_repr(sec_type_idx, sec_cont, particle_type_cfg, cond_enc, emb_dim)
return torch.cat([cont, type_repr], dim=-1)
def _assemble_stage2_real(
sec_cont: torch.Tensor,
sec_type_idx: torch.Tensor,
particle_type_cfg: dict,
generator: str,
cond_enc: torch.nn.Module,
emb_dim: int,
) -> torch.Tensor:
"""Ground-truth flattened stage-2 vector for `Stage2OneShot` — the
flattened form of `_assemble_stage2_ar_target`, which
`Stage2Autoregressive`'s per-token target also uses; the two must stay in
lockstep. See `_assemble_stage2_ar_target`'s docstring for the
(target, generator) width rules."""
return _assemble_stage2_ar_target(
sec_cont, sec_type_idx, particle_type_cfg, generator, cond_enc, emb_dim
).flatten(1)
def _stick_fraction(sec_cont: torch.Tensor) -> torch.Tensor:
"""(B, K_MAX) — sigmoid of each slot's own stick-breaking logit
(`sec_cont[...,0]`); scale-free (see `giant.data.transforms.
encode_secondaries`), so this needs no absolute `e_sec`."""
return torch.sigmoid(sec_cont[..., 0])
def _remaining_energy_fraction(fraction: torch.Tensor) -> torch.Tensor:
"""(B, K_MAX) — fraction of the original e_sec budget unclaimed entering
slot i: `1.0` at `i=0`, `prod_{j<i}(1-fraction_j)` for `i>=1`
(docs/v0.3.0-design.md §6.3 "no re-derivation needed": the existing
stick-breaking encoding is already scale-free, so this is derivable from
the batch's ground-truth stick logits alone, no `e_sec` required)."""
cumprod = torch.cumprod(1.0 - fraction, dim=1)
return torch.cat([torch.ones_like(cumprod[:, :1]), cumprod[:, :-1]], dim=1)
def _shift_prev(x: torch.Tensor) -> torch.Tensor:
"""`(B, K, ...)` -> same shape, slot i holds slot i-1's value; slot 0 gets
an arbitrary zero placeholder (never read as-is see `_ar_has_prev`;
`MarkovHistory` substitutes its own learned start vector there instead)."""
return torch.cat([torch.zeros_like(x[:, :1]), x[:, :-1]], dim=1)
def _ar_has_prev(k_max: int, device: torch.device) -> torch.Tensor:
"""`(1, K_MAX)` bool: True for slot index `>= 1`. Correct without
`n_sec`: `sec_mask` is a prefix mask, so any *valid* token at `k>=1`
always has a valid predecessor at `k-1`; the only wrong cases are tokens
that are themselves padding, already masked out of every loss."""
return (torch.arange(k_max, device=device) >= 1).unsqueeze(0)
def _assemble_stage2_ar_inputs(
sec_cont: torch.Tensor,
sec_type_idx: torch.Tensor,
particle_type_cfg: dict,
cond_enc: torch.nn.Module,
emb_dim: int,
) -> dict[str, torch.Tensor]:
"""Ground-truth per-token AR conditioning tensors — all `(B, K_MAX, ...)`
or `(B, K_MAX)`, built in one vectorized pass (teacher forcing means
every token's input is ground truth, docs/v0.3.0-design.md §6.2 point 3).
Keys match `Stage2Autoregressive.forward`'s trailing kwargs."""
device = sec_cont.device
B, K = sec_cont.shape[0], sec_cont.shape[1]
fraction = _stick_fraction(sec_cont)
type_repr = _type_repr(sec_type_idx, sec_cont, particle_type_cfg, cond_enc, emb_dim)
history_feat = torch.cat(
[
_shift_prev(fraction).unsqueeze(-1),
_shift_prev(sec_cont[..., 1:CONT_SLOT_DIM]),
_shift_prev(type_repr),
],
dim=-1,
)
slot_idx = (torch.arange(K, device=device).float() / max(K - 1, 1)).unsqueeze(0)
return {
"history_feat": history_feat,
"has_prev": _ar_has_prev(K, device).expand(B, -1),
"remaining_frac": _remaining_energy_fraction(fraction),
"slot_idx": slot_idx.expand(B, -1),
}
def _stage2_tf_prob(
mode: str, p_start: float, p_end: float, epoch: int, total_epochs: int
) -> float:
"""P(condition slot k+1 on the TRUE token k rather than the model's own
prediction), for the current epoch (docs/v0.3.0-design.md §3.3
`stage2_model.autoregressive.teacher_forcing`). `"always"`/`"never"` are
the two degenerate constants; `"scheduled"` linearly interpolates
`p_start` (epoch 0) to `p_end` (the final epoch) standard scheduled
sampling (Bengio et al. 2015)."""
if mode == "always":
return 1.0
if mode == "never":
return 0.0
frac = epoch / max(total_epochs - 1, 1)
frac = min(max(frac, 0.0), 1.0)
return p_start + (p_end - p_start) * frac
def _history_repr_from_ar_sample(
sec_cont_pred: torch.Tensor,
sec_type_pred: torch.Tensor,
particle_type_cfg: dict,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""`(fraction, direction, type_repr)` — the same triple `_type_repr` /
`_stick_fraction` derive from ground truth, but from a free-running
`sample_secondaries_ar` self-sample instead, so the two can be mixed
slot-by-slot under scheduled sampling (`_assemble_stage2_ar_inputs_scheduled`).
`target="onehot"` collapses the raw per-slot type logits to a hard
one-hot of `argmax` `sample_secondaries_ar`'s own history convention
(see its docstring), matching what `MarkovHistory`/`AttentionHistory`
were trained on; the other two targets are already the right
representation."""
fraction = torch.sigmoid(sec_cont_pred[..., 0])
direction = sec_cont_pred[..., 1:4]
if particle_type_cfg.get("target", "physical") == "onehot":
type_dim = sec_type_pred.size(-1)
type_repr = F.one_hot(sec_type_pred.argmax(-1), num_classes=type_dim).float()
else:
type_repr = sec_type_pred
return fraction, direction, type_repr
def _assemble_stage2_ar_inputs_scheduled(
model: torch.nn.Module,
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
stage1_ctx: torch.Tensor,
sec_cont: torch.Tensor,
sec_type_idx: torch.Tensor,
n_sec: torch.Tensor,
particle_type_cfg: dict,
cond_enc: torch.nn.Module,
emb_dim: int,
p_tf: float,
sample_steps: int,
) -> dict[str, torch.Tensor]:
"""Scheduled-sampling counterpart of `_assemble_stage2_ar_inputs`
(docs/v0.3.0-design.md §3.3 `teacher_forcing` = "scheduled"/"never"):
each slot's history is the TRUE previous token with probability `p_tf`
(an independent per-example, per-slot Bernoulli draw) and the model's own
free-running prediction otherwise closing the train/inference gap that
`teacher_forcing="always"` (ground truth throughout training) never sees.
`p_tf >= 1.0` degenerates exactly to `_assemble_stage2_ar_inputs` (and
skips self-sampling entirely), so callers can call this unconditionally.
The free-running estimate is a REAL autoregressive self-sample
`giant.sample.sample_secondaries_ar` under `torch.no_grad()` not a
cheap one-step proxy, so building it costs the same `k_max` (`* steps`
for flow) sequential forwards `sample.py` pays at inference, EVERY batch
this is called on (§6.4's cost note, paid at train time too whenever
teacher_forcing != "always"). Fully detached: gradient only ever flows
through the "real" target path each stage trainer already uses
(`_assemble_stage2_ar_target`), never through this self-sample.
"""
device = sec_cont.device
B, K = sec_cont.shape[0], sec_cont.shape[1]
if p_tf >= 1.0:
return _assemble_stage2_ar_inputs(
sec_cont, sec_type_idx, particle_type_cfg, cond_enc, emb_dim
)
was_training = model.training
sec_cont_pred, sec_type_pred, _ = sample_secondaries_ar(
model, cond_cont, cond_cat, stage1_ctx, n_sec, steps=sample_steps
)
if was_training:
model.train()
fraction_gt = _stick_fraction(sec_cont)
dir_gt = sec_cont[..., 1:4]
type_repr_gt = _type_repr(
sec_type_idx, sec_cont, particle_type_cfg, cond_enc, emb_dim
)
fraction_pred, dir_pred, type_repr_pred = _history_repr_from_ar_sample(
sec_cont_pred, sec_type_pred, particle_type_cfg
)
use_gt = torch.rand(B, K, device=device) < p_tf
fraction = torch.where(use_gt, fraction_gt, fraction_pred)
direction = torch.where(use_gt.unsqueeze(-1), dir_gt, dir_pred)
type_repr = torch.where(use_gt.unsqueeze(-1), type_repr_gt, type_repr_pred)
own_feat = torch.cat([fraction.unsqueeze(-1), direction, type_repr], dim=-1)
return {
"history_feat": _shift_prev(own_feat),
"has_prev": _ar_has_prev(K, device).expand(B, -1),
"remaining_frac": _remaining_energy_fraction(fraction),
"slot_idx": (torch.arange(K, device=device).float() / max(K - 1, 1))
.unsqueeze(0)
.expand(B, -1),
}
def _relax_onehot_type_slice(
x_flat: torch.Tensor,
k_max: int,
cont_dim: int,
type_dim: int,
tau: float,
grad_probe: dict[str, float] | None = None,
) -> torch.Tensor:
"""Straight-through Gumbel-softmax relaxation of the per-slot type slice
inside a flattened `(B, k_max * (cont_dim + type_dim))` WGAN generator
output decision 5 (docs/v0.3.0-design.md §2.1): the forward pass is a
hard one-hot (matching what the critic sees from real data), the
backward pass flows smooth gradient. Continuous slots (stick/dir, and
the type slice itself under `target = "embedding"`, which never calls
this) pass through unchanged.
`grad_probe`, if given, gets `["cont"]`/`["type"]` populated with the L2
norm of the gradient reaching this split point during the next
`.backward()` call that touches it a backward hook, not a second
backward pass. This is the §11.4 differentiability validation-obligation
instrumentation (docs/v0.3.0-design.md): the trunk-gradient contribution
from the type slice vs. the continuous slices, for
`particle_type.target="onehot"` + `generator="wgan"`. Only ever populated
on a `did_g_step` batch the critic step backprops through
`fake.detach()`, which never reaches these hooks so it stays empty
(callers default to `0.0`) otherwise."""
B = x_flat.size(0)
x = x_flat.view(B, k_max, cont_dim + type_dim)
cont, type_logits = x[..., :cont_dim], x[..., cont_dim:]
if grad_probe is not None:
cont.register_hook(lambda g: grad_probe.__setitem__("cont", g.norm().item()))
type_logits.register_hook(
lambda g: grad_probe.__setitem__("type", g.norm().item())
)
type_soft = F.gumbel_softmax(type_logits, tau=tau, hard=True, dim=-1)
return torch.cat([cont, type_soft], dim=-1).reshape(B, -1)
+987
View File
@@ -0,0 +1,987 @@
"""Per-stage trainers: optimizer(s), EMA, LR schedule, and the per-batch step.
`StageSpec` resolves one stage's slice of the config once, so the two
concrete trainers share a single constructor shape instead of ~24 keyword
arguments each, and `StageTrainer` carries every piece that used to be
copy-pasted between them (cosine warmup, EMA, checkpoint state, LR resume,
train/eval toggling).
Each trainer also *declares* the metrics it emits, as `MetricSpec` lists
that declaration is the single source of truth for `metrics.csv` and W&B
columns (see `giant.training.metrics`) and exposes the three small hooks
(`batch_loss`, `summary`, `val_objective`) that let the epoch loop treat
adversarial and non-adversarial stages identically.
"""
import copy
import math
from dataclasses import dataclass, field
import torch
import torch.nn.functional as F
import torch.optim as optim
from giant.constants import CONT_SLOT_DIM
from giant.model.network import Router, stage2_type_dim
from giant.model.schedule import (
CosineSchedule,
flow_matching_loss,
flow_matching_loss_secondary,
flow_matching_loss_secondary_ar,
)
from giant.model.wgan import generator_loss, gradient_penalty
from giant.training.metrics import MetricSpec, stage_metric, train_metric, val_metric
from giant.training.stage2_inputs import (
_assemble_stage2_ar_inputs_scheduled,
_assemble_stage2_ar_target,
_assemble_stage2_real,
_gumbel_tau,
_relax_onehot_type_slice,
_stage2_tf_prob,
)
@torch.no_grad()
def _update_ema(
ema_model: torch.nn.Module, model: torch.nn.Module, decay: float
) -> None:
for ema_p, p in zip(ema_model.parameters(), model.parameters()):
ema_p.mul_(decay).add_(p, alpha=1 - decay)
def _stage_router(model: torch.nn.Module) -> Router | None:
"""A stage model's Router, if its trunk is routed — else None.
Post-step-2 refactor the router lives at `model.trunk.router`
(`giant.model.network.RoutedTrunk`), not `model.router` directly.
"""
trunk = getattr(model, "trunk", None)
return getattr(trunk, "router", None)
def _cosine_warmup_lambda(warmup_steps: int, total_steps: int):
"""Linear warmup for `warmup_steps`, then cosine decay to zero over the
remainder the LR schedule both trainers use, in their own step units
(optimizer steps for flow/ddpm, generator steps for WGAN)."""
def _lr_lambda(step: int) -> float:
if warmup_steps > 0 and step < warmup_steps:
return (step + 1) / warmup_steps
t = step - warmup_steps
T = max(total_steps - warmup_steps, 1)
return 0.5 * (1.0 + math.cos(math.pi * min(t, T) / T))
return _lr_lambda
def _batch_to_device(batch: tuple, device: torch.device) -> tuple:
return tuple(t.to(device) for t in batch)
@dataclass(frozen=True)
class StageSpec:
"""One stage's resolved training configuration.
Built once by `StageSpec.from_config`, which is the only place that reads
the `cfg` dict so a new config key means one new field and one new read,
not another argument threaded through two constructors.
"""
name: str
is_stage2: bool
generator: str
decoder: str = "one_shot"
# loss weights
lambda_weight: float = 1.0
n_sec_lambda: float = 0.1
# particle-type target (stage 2 only)
particle_type: dict = field(default_factory=lambda: {"target": "physical"})
particle_type_emb_dim: int = 16
# optimization
lr: float = 3e-4
weight_decay: float = 0.01
ema_decay: float = 0.9999
warmup_epochs: int = 0
epochs: int = 1
steps_per_epoch: int = 1
# routing auxiliaries
lambda_balance: float = 0.0
lambda_proc: float = 0.0
lambda_entropy: float = 0.0
gumbel_tau_start: float = 1.0
gumbel_tau_end: float = 0.1
# autoregressive stage 2
teacher_forcing: str = "always"
tf_p_start: float = 1.0
tf_p_end: float = 1.0
ar_sample_steps: int = 10
# generator-specific
ddpm_n_steps: int = 1000
n_critic: int = 5
gp_weight: float = 10.0
critic_lr: float = 0.0
type_gumbel_tau_start: float = 1.0
type_gumbel_tau_end: float = 0.1
@classmethod
def from_config(
cls, cfg: dict, name: str, is_stage2: bool, steps_per_epoch: int
) -> "StageSpec":
t = cfg["train"]
stage_cfg = cfg[f"{name}_model"]
router_cfg = stage_cfg.get("router") or {}
wgan_cfg = stage_cfg.get("wgan") or {}
ar_cfg = (stage_cfg.get("autoregressive") or {}) if is_stage2 else {}
return cls(
name=name,
is_stage2=is_stage2,
generator=stage_cfg["generator"],
decoder=stage_cfg.get("decoder", "one_shot") if is_stage2 else "one_shot",
lambda_weight=stage_cfg.get("lambda", 1.0),
n_sec_lambda=cfg["stage2_model"].get("n_sec", {}).get("lambda", 0.1),
particle_type=cfg["stage2_model"].get("particle_type")
or {"target": "physical"},
particle_type_emb_dim=cfg["conditioning"]["particle"]["emb_dim"],
lr=t["lr"],
weight_decay=t.get("weight_decay", 0.01),
ema_decay=t.get("ema_decay", 0.9999),
warmup_epochs=t["warmup_epochs"],
epochs=t["epochs"],
steps_per_epoch=max(steps_per_epoch, 1),
lambda_balance=router_cfg.get("lambda_balance", 0.0),
lambda_proc=router_cfg.get("lambda_proc", 0.0),
lambda_entropy=router_cfg.get("lambda_entropy", 0.0),
gumbel_tau_start=router_cfg.get("gumbel_tau_start", 1.0),
gumbel_tau_end=router_cfg.get("gumbel_tau_end", 0.1),
teacher_forcing=ar_cfg.get("teacher_forcing", "always"),
tf_p_start=ar_cfg.get("tf_p_start", 1.0),
tf_p_end=ar_cfg.get("tf_p_end", 1.0),
# AR self-sampling under scheduled/never teacher forcing reuses
# train.validate_steps as its flow-matching ODE step count — no
# dedicated config key for this (docs/v0.3.0-design.md §3.3 lists
# tf_p_start/tf_p_end/attn_n_heads/attn_n_layers only).
ar_sample_steps=t.get("validate_steps", 10),
ddpm_n_steps=stage_cfg.get("ddpm", {}).get("n_steps", 1000),
n_critic=wgan_cfg.get("n_critic", 5),
gp_weight=wgan_cfg.get("gp_weight", 10.0),
critic_lr=wgan_cfg.get("critic_lr", 0.0),
type_gumbel_tau_start=wgan_cfg.get("gumbel_tau_start", 1.0),
type_gumbel_tau_end=wgan_cfg.get("gumbel_tau_end", 0.1),
)
class StageTrainer:
"""One active stage's optimizer(s), EMA, and per-batch step.
Reads only the shared batch tuple `(cond_cont, cond_cat, x1_s1, n_sec,
sec_cont, proc_idx, sec_type_idx)` stage 2 always conditions on the
ground-truth `x1_s1` (`stage2_model.stage1_context = "truth"`,
stage-level teacher forcing; `"sampled"` is not implemented see
docs/v0.3.0-design.md §3.3), so stage trainers never need each other's
output at train time. This means "stage-2-only training is a cheap
ablation, not new plumbing" (design doc §7) falls out for free: a trainer
only exists for active stages, and inactive stages are simply never
constructed.
Grad-norm clipping is per-stage here v0.2's single shared optimizer
clipped both stages' gradients jointly; splitting per stage is a small,
disclosed behavior change. It doesn't affect Adam's per-parameter update
math itself (no cross-parameter coupling), only the clip threshold's
scope.
"""
#: Metrics this trainer emits, declared once — `giant.training.metrics`
#: derives every CSV/W&B column from these. Instance attributes rather
#: than class constants because some are conditional on the stage's own
#: configuration (see `WGANStageTrainer.__init__`).
train_metrics: list[MetricSpec]
val_metrics: list[MetricSpec]
stage_metrics: list[MetricSpec]
#: False for adversarial stages, which have no monotone per-batch
#: validation loss worth averaging (see `val_objective`).
supports_val_loss: bool = True
#: Built by the subclass (the optimizer flavour differs) and wired to the
#: schedule via `_init_lr_schedule`.
optimizer: optim.Optimizer
lr_sched: optim.lr_scheduler.LambdaLR
total_steps: int
def __init__(
self,
spec: StageSpec,
model: torch.nn.Module,
device: torch.device,
extra_modules: tuple[torch.nn.Module, ...] = (),
) -> None:
self.spec = spec
self.name = spec.name
self.is_stage2 = spec.is_stage2
self.generator = spec.generator
self.decoder = spec.decoder
self.device = device
self.model = model.to(device)
self.router = _stage_router(self.model)
self._modules = (self.model, *extra_modules)
self.particle_type_cfg = dict(spec.particle_type or {"target": "physical"})
self.particle_type_emb_dim = spec.particle_type_emb_dim
self.ema_decay = spec.ema_decay
self.ema_model: torch.nn.Module | None = None
if spec.ema_decay > 0:
self.ema_model = copy.deepcopy(self.model).eval()
for p in self.ema_model.parameters():
p.requires_grad_(False)
# --- schedule -------------------------------------------------------
def _init_lr_schedule(
self, optimizer: optim.Optimizer, warmup_steps: int, total_steps: int
) -> None:
self._lr_lambda = _cosine_warmup_lambda(warmup_steps, total_steps)
self.total_steps = total_steps
self.lr_sched = optim.lr_scheduler.LambdaLR(optimizer, self._lr_lambda)
# --- per-batch (subclass responsibility) ----------------------------
def step(self, batch: tuple, device: torch.device, global_step: int) -> dict:
raise NotImplementedError
def val_loss(self, batch: tuple, device: torch.device) -> dict:
raise NotImplementedError
# --- reporting hooks ------------------------------------------------
def batch_loss(self, stats: dict) -> float:
"""The single number this stage contributes to the progress bar's
smoothed loss."""
raise NotImplementedError
def summary(self, means: dict) -> str:
"""This stage's fragment of the end-of-epoch console line."""
raise NotImplementedError
def val_objective(
self, train_means: dict, val_means: dict, marginal_kl: float
) -> float:
"""This stage's contribution to the best-checkpoint selection score."""
raise NotImplementedError
# --- mode / state ---------------------------------------------------
def sampling_model(self) -> torch.nn.Module:
return self.ema_model if self.ema_model is not None else self.model
def train_mode(self) -> None:
for module in self._modules:
module.train()
def eval_mode(self) -> None:
for module in self._modules:
module.eval()
def _extra_state(self) -> dict:
"""Subclass state beyond model/optimizer/lr_sched/EMA."""
return {}
def _load_extra_state(self, sd: dict) -> None:
return None
def state_dict(self) -> dict:
sd = {
"model": self.model.state_dict(),
"optimizer": self.optimizer.state_dict(),
"lr_sched": self.lr_sched.state_dict(),
}
if self.ema_model is not None:
sd["model_ema"] = self.ema_model.state_dict()
sd.update(self._extra_state())
return sd
def load_state_dict(self, sd: dict) -> None:
self.model.load_state_dict(sd["model"])
self.optimizer.load_state_dict(sd["optimizer"])
self.lr_sched.load_state_dict(sd["lr_sched"])
if self.ema_model is not None:
self.ema_model.load_state_dict(sd.get("model_ema", sd["model"]))
self._load_extra_state(sd)
def _resume_extra_lr(self, lr: float) -> None:
return None
def resume_lr(self, lr: float) -> None:
"""Restore the configured `lr`'s authority after `load_state_dict`
restored the checkpoint's own base LR."""
self.lr_sched.base_lrs = [lr for _ in self.lr_sched.base_lrs]
resumed_lr = lr * self._lr_lambda(self.lr_sched.last_epoch)
for group in self.optimizer.param_groups:
group["lr"] = resumed_lr
self._resume_extra_lr(lr)
class FlowDDPMStageTrainer(StageTrainer):
"""flow or ddpm generator for a single stage."""
def __init__(
self, spec: StageSpec, model: torch.nn.Module, device: torch.device
) -> None:
if spec.is_stage2 and spec.generator not in ("flow",):
raise NotImplementedError(
f"stage2_model.generator={spec.generator!r} is accepted by the "
"schema but not implemented in v0.3.0 for stage 2 (only "
"'flow' and 'wgan' have a stage-2 secondary-decoder loss — "
"see docs/v0.3.0-design.md §11.2)"
)
super().__init__(spec, model, device)
self.particle_type_lambda = self.particle_type_cfg.get("lambda", 1.0)
# Width of the type slice actually folded into x1_s2 by
# _assemble_stage2_real, under this trainer's generator (flow/ddpm
# only — see the NotImplementedError above): "physical" keeps it
# folded in (PARTICLE_PHYS_DIM wide, unchanged from v0.2); "onehot"/
# "embedding" pull it out into model.type_head instead (0 here).
self._flow_type_dim = (
None
if self.particle_type_cfg.get("target", "physical") == "physical"
else 0
)
self.params = list(self.model.parameters())
self.optimizer = optim.AdamW(
self.params, lr=spec.lr, weight_decay=spec.weight_decay
)
self._init_lr_schedule(
self.optimizer,
warmup_steps=spec.warmup_epochs * spec.steps_per_epoch,
total_steps=max(spec.epochs * spec.steps_per_epoch, 1),
)
self.ddpm_schedule = (
CosineSchedule(T=spec.ddpm_n_steps).to(device)
if spec.generator == "ddpm"
else None
)
self.train_metrics = [
train_metric(key)
for key in (
"loss",
"loss_gen",
"loss_nsec",
"loss_balance",
"loss_proc",
"loss_entropy",
"nsec_acc",
"loss_type",
"type_acc",
"grad_norm",
)
]
self.val_metrics = [
val_metric(key)
for key in (
"loss",
"loss_gen",
"loss_nsec",
"nsec_acc",
"loss_type",
"type_acc",
)
]
self.stage_metrics = [stage_metric("lr")]
def _predict_n_sec(self, cond_cont, cond_cat, stage1_ctx):
if self.model.n_sec_head is None:
return None
if self.is_stage2:
return self.model.predict_n_sec(cond_cont, cond_cat, stage1_ctx)
return self.model.predict_n_sec(cond_cont, cond_cat)
def _generator_loss(
self, cond_cont, cond_cat, x1_s1, x1_s2, sec_mask, stage1_ctx, ar_inputs=None
):
if not self.is_stage2:
if self.generator == "flow":
return flow_matching_loss(self.model, x1_s1, cond_cont, cond_cat)
assert self.ddpm_schedule is not None
return self.ddpm_schedule.loss(self.model, x1_s1, cond_cont, cond_cat)
if self.decoder == "autoregressive":
assert ar_inputs is not None
return flow_matching_loss_secondary_ar(
self.model,
x1_s2,
cond_cont,
cond_cat,
stage1_ctx,
ar_inputs["history_feat"],
ar_inputs["has_prev"],
ar_inputs["remaining_frac"],
ar_inputs["slot_idx"],
sec_mask,
type_dim=self._flow_type_dim,
)
return flow_matching_loss_secondary(
self.model,
x1_s2,
cond_cont,
cond_cat,
stage1_ctx,
sec_mask,
type_dim=self._flow_type_dim,
)
def _type_loss(
self,
cond_cont,
cond_cat,
stage1_ctx,
sec_type_idx,
sec_mask,
device,
ar_inputs=None,
):
"""CE (`target="onehot"`) or MSE (`target="embedding"`) loss for the
stage-2 model's `type_head` — the non-adversarial counterpart to
WGANStageTrainer's ST-Gumbel-into-the-critic path (decision 2/5).
Zero when this stage has no `type_head` (stage 1, or
`particle_type.target = "physical"`)."""
l_type = torch.zeros((), device=device)
type_acc = torch.zeros((), device=device)
type_head = getattr(self.model, "type_head", None)
if not self.is_stage2 or type_head is None:
return l_type, type_acc
if self.decoder == "autoregressive":
assert ar_inputs is not None
type_out = self.model.predict_type(
cond_cont,
cond_cat,
stage1_ctx,
ar_inputs["history_feat"],
ar_inputs["has_prev"],
ar_inputs["remaining_frac"],
ar_inputs["slot_idx"],
)
else:
type_out = self.model.predict_type(cond_cont, cond_cat, stage1_ctx)
mask = sec_mask.float()
denom = mask.sum().clamp(min=1)
if self.particle_type_cfg.get("target") == "onehot":
ce = F.cross_entropy(
type_out.transpose(1, 2), sec_type_idx, reduction="none"
)
l_type = (ce * mask).sum() / denom
type_acc = (
(type_out.argmax(-1) == sec_type_idx).float() * mask
).sum() / denom
else: # "embedding"
target_vec = self.model.cond_enc.pdg_emb(sec_type_idx).detach()
se = ((type_out - target_vec) ** 2).mean(-1)
l_type = (se * mask).sum() / denom
return l_type, type_acc
def _compute(
self, batch: tuple, device: torch.device, epoch: int | None = None
) -> dict:
"""`epoch=None` (the `val_loss` path) always uses full teacher
forcing (`p_tf=1.0`) regardless of `spec.teacher_forcing` validation
should stay a stable, non-stochastic ground-truth comparison; only
the training `step` path schedules `p_tf` by epoch."""
(
cond_cont,
cond_cat,
x1_s1,
n_sec,
sec_cont,
proc_idx,
sec_type_idx,
) = _batch_to_device(batch, device)
sec_mask = torch.arange(sec_cont.size(1), device=device).unsqueeze(
0
) < n_sec.unsqueeze(1)
stage1_ctx = x1_s1.detach()
x1_s2 = None
ar_inputs = None
if self.is_stage2 and self.decoder == "autoregressive":
p_tf = (
1.0
if epoch is None
else _stage2_tf_prob(
self.spec.teacher_forcing,
self.spec.tf_p_start,
self.spec.tf_p_end,
epoch,
self.spec.epochs,
)
)
ar_inputs = _assemble_stage2_ar_inputs_scheduled(
self.model,
cond_cont,
cond_cat,
stage1_ctx,
sec_cont,
sec_type_idx,
n_sec,
self.particle_type_cfg,
self.model.cond_enc,
self.particle_type_emb_dim,
p_tf,
self.spec.ar_sample_steps,
)
x1_s2 = _assemble_stage2_ar_target(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
self.generator,
self.model.cond_enc,
self.particle_type_emb_dim,
)
elif self.is_stage2:
x1_s2 = _assemble_stage2_real(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
self.generator,
self.model.cond_enc,
self.particle_type_emb_dim,
)
l_gen = self._generator_loss(
cond_cont, cond_cat, x1_s1, x1_s2, sec_mask, stage1_ctx, ar_inputs=ar_inputs
)
n_sec_logits = self._predict_n_sec(cond_cont, cond_cat, stage1_ctx)
l_nsec = torch.zeros((), device=device)
nsec_acc = torch.zeros((), device=device)
if n_sec_logits is not None:
l_nsec = F.cross_entropy(n_sec_logits, n_sec)
nsec_acc = (n_sec_logits.argmax(dim=-1) == n_sec).float().mean()
l_type, type_acc = self._type_loss(
cond_cont,
cond_cat,
stage1_ctx,
sec_type_idx,
sec_mask,
device,
ar_inputs=ar_inputs,
)
l_balance = l_proc = l_entropy = torch.zeros((), device=device)
if self.router is not None:
l_balance = self.router.balance_loss(cond_cont, cond_cat)
l_proc = self.router.classify_loss(cond_cont, cond_cat, proc_idx)
l_entropy = self.router.entropy_loss(cond_cont, cond_cat)
total = (
self.spec.lambda_weight * l_gen
+ self.spec.n_sec_lambda * l_nsec
+ self.particle_type_lambda * l_type
)
if self.spec.lambda_balance > 0:
total = total + self.spec.lambda_balance * l_balance
if self.spec.lambda_proc > 0:
total = total + self.spec.lambda_proc * l_proc
if self.spec.lambda_entropy > 0:
total = total + self.spec.lambda_entropy * l_entropy
return {
"loss": total,
"loss_gen": l_gen,
"loss_nsec": l_nsec,
"loss_type": l_type,
"type_acc": type_acc,
"loss_balance": l_balance,
"loss_proc": l_proc,
"loss_entropy": l_entropy,
"nsec_acc": nsec_acc,
}
def step(self, batch: tuple, device: torch.device, global_step: int) -> dict:
if self.router is not None:
self.router.gumbel_tau = _gumbel_tau(
global_step,
self.total_steps,
self.spec.gumbel_tau_start,
self.spec.gumbel_tau_end,
)
epoch = global_step // self.spec.steps_per_epoch
out = self._compute(batch, device, epoch=epoch)
self.optimizer.zero_grad()
out["loss"].backward()
grad_norm = torch.nn.utils.clip_grad_norm_(self.params, 1.0)
self.optimizer.step()
self.lr_sched.step()
if self.ema_model is not None:
_update_ema(self.ema_model, self.model, self.ema_decay)
stats = {key: value.item() for key, value in out.items()}
stats["grad_norm"] = grad_norm.item()
stats["lr"] = self.optimizer.param_groups[0]["lr"]
return stats
@torch.no_grad()
def val_loss(self, batch: tuple, device: torch.device) -> dict:
return {
key: value.item() for key, value in self._compute(batch, device).items()
}
# --- reporting ------------------------------------------------------
def batch_loss(self, stats: dict) -> float:
return stats["loss"]
def summary(self, means: dict) -> str:
return f"{self.name}[loss={means.get('loss', 0.0):.3f}]"
def val_objective(
self, train_means: dict, val_means: dict, marginal_kl: float
) -> float:
return val_means.get("loss", 0.0)
class WGANStageTrainer(StageTrainer):
"""WGAN-GP generator+critic for a single stage (see giant/model/wgan.py).
Ports `_wgan_train_step` to operate on one stage instead of two fused
together the critic updates every batch; every `n_critic`-th batch
additionally updates the generator (`did_g_step`). The (non-adversarial)
n_sec classifier, when this stage's model owns it, updates every batch
regardless folded into whichever generator optimizer step happens this
batch, same precedent as v0.2.
"""
supports_val_loss = False
def __init__(
self,
spec: StageSpec,
model: torch.nn.Module,
critic: torch.nn.Module,
device: torch.device,
) -> None:
self.critic = critic.to(device)
super().__init__(spec, model, device, extra_modules=(self.critic,))
self.n_critic = max(spec.n_critic, 1)
self.gp_weight = spec.gp_weight
self.critic_lr = spec.critic_lr
self.g_params = list(self.model.parameters())
self.d_params = list(self.critic.parameters())
# WGAN-GP recipe (Gulrajani et al. 2017): Adam, beta1=0, no weight decay.
self.optimizer = optim.Adam(self.g_params, lr=spec.lr, betas=(0.0, 0.9))
self.optimizer_d = optim.Adam(
self.d_params,
lr=spec.critic_lr if spec.critic_lr > 0 else spec.lr,
betas=(0.0, 0.9),
)
# Generator steps fire every n_critic-th batch, so warmup/decay must
# be counted in those units, matching v0.2.
gen_steps_per_epoch = max(spec.steps_per_epoch // self.n_critic, 1)
self._init_lr_schedule(
self.optimizer,
warmup_steps=spec.warmup_epochs * gen_steps_per_epoch,
total_steps=max(spec.epochs * gen_steps_per_epoch, 1),
)
train_keys = [
"d_loss",
"g_loss",
"wasserstein",
"gp_loss",
"loss_nsec",
"nsec_acc",
"grad_norm_d",
"grad_norm_g",
]
if self.is_stage2 and self.particle_type_cfg.get("target") == "onehot":
# §11.4 differentiability instrumentation — only meaningful when
# the type slice is a straight-through Gumbel relaxation.
train_keys += ["grad_norm_type_slice", "grad_norm_cont_slice"]
self.train_metrics = [train_metric(key) for key in train_keys]
self.val_metrics = []
self.stage_metrics = [stage_metric("lr"), stage_metric("critic_lr")]
def _stage2_real_and_fake(self, batch_tensors, stage1_ctx, global_step, device):
"""Build `(real, fake_raw, mask, critic_fn)` for stage 2, covering
both decoders and all three particle-type targets. `fake_raw` still
needs the caller's straight-through relaxation under
`particle_type.target = "onehot"`, and neither tensor is masked-and-
multiplied on the fake side yet."""
cond_cont, cond_cat, n_sec, sec_cont, sec_type_idx = batch_tensors
B = cond_cont.size(0)
type_dim = stage2_type_dim(self.particle_type_cfg, self.particle_type_emb_dim)
slot_width = CONT_SLOT_DIM + type_dim
k_max = sec_cont.size(1)
sec_mask = torch.arange(k_max, device=device).unsqueeze(0) < n_sec.unsqueeze(1)
mask = sec_mask.unsqueeze(-1).expand(-1, -1, slot_width).reshape(B, -1).float()
def critic_fn(x):
return self.critic(x, cond_cont, cond_cat, stage1_ctx)
if self.decoder == "autoregressive":
epoch = global_step // self.spec.steps_per_epoch
p_tf = _stage2_tf_prob(
self.spec.teacher_forcing,
self.spec.tf_p_start,
self.spec.tf_p_end,
epoch,
self.spec.epochs,
)
ar = _assemble_stage2_ar_inputs_scheduled(
self.model,
cond_cont,
cond_cat,
stage1_ctx,
sec_cont,
sec_type_idx,
n_sec,
self.particle_type_cfg,
self.model.cond_enc,
self.particle_type_emb_dim,
p_tf,
self.spec.ar_sample_steps,
)
real = (
_assemble_stage2_ar_target(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
"wgan",
self.model.cond_enc,
self.particle_type_emb_dim,
).reshape(B, -1)
* mask
)
z = torch.randn(B, k_max, self.model.noise_dim, device=device)
fake_raw = self.model(
z,
cond_cont,
cond_cat,
stage1_ctx,
ar["history_feat"],
ar["has_prev"],
ar["remaining_frac"],
ar["slot_idx"],
).reshape(B, -1)
else:
real = (
_assemble_stage2_real(
sec_cont,
sec_type_idx,
self.particle_type_cfg,
"wgan",
self.model.cond_enc,
self.particle_type_emb_dim,
)
* mask
)
z = torch.randn(B, self.model.noise_dim, device=device)
fake_raw = self.model(z, cond_cont, cond_cat, stage1_ctx)
return real, fake_raw, mask, critic_fn
def step(self, batch: tuple, device: torch.device, global_step: int) -> dict:
(
cond_cont,
cond_cat,
x1_s1,
n_sec,
sec_cont,
_proc_idx,
sec_type_idx,
) = _batch_to_device(batch, device)
B = cond_cont.size(0)
stage1_ctx = x1_s1.detach()
grad_probe: dict[str, float] = {}
if not self.is_stage2:
real = x1_s1
def critic_fn(x):
return self.critic(x, cond_cont, cond_cat)
z = torch.randn(B, self.model.noise_dim, device=device)
fake = self.model(z, cond_cont, cond_cat)
mask = None
else:
real, fake_raw, mask, critic_fn = self._stage2_real_and_fake(
(cond_cont, cond_cat, n_sec, sec_cont, sec_type_idx),
stage1_ctx,
global_step,
device,
)
if self.particle_type_cfg.get("target", "physical") == "onehot":
# Straight-through Gumbel-softmax relaxation of the type
# slice only (decision 5) — the critic must see a hard
# one-hot forward (matching what "real" data looks like)
# while gradient still flows smoothly to the generator.
# grad_probe captures the §11.4 gradient-magnitude
# instrumentation — see _relax_onehot_type_slice's docstring.
tau = _gumbel_tau(
global_step,
self.total_steps,
self.spec.type_gumbel_tau_start,
self.spec.type_gumbel_tau_end,
)
fake_raw = _relax_onehot_type_slice(
fake_raw,
sec_cont.size(1),
CONT_SLOT_DIM,
stage2_type_dim(self.particle_type_cfg, self.particle_type_emb_dim),
tau,
grad_probe=grad_probe,
)
fake = fake_raw * mask
# --- critic step (every batch) ---
fake_detached = fake.detach()
real_score = critic_fn(real)
fake_score = critic_fn(fake_detached)
gp = gradient_penalty(critic_fn, real, fake_detached, mask=mask)
d_loss = fake_score.mean() - real_score.mean() + self.gp_weight * gp
wasserstein = (real_score.mean() - fake_score.mean()).detach()
self.optimizer_d.zero_grad()
d_loss.backward()
grad_norm_d = torch.nn.utils.clip_grad_norm_(self.d_params, 1.0)
self.optimizer_d.step()
# --- generator (+ n_sec) step ---
did_g_step = global_step % self.n_critic == 0
n_sec_logits = None
if self.model.n_sec_head is not None:
n_sec_logits = (
self.model.predict_n_sec(cond_cont, cond_cat)
if not self.is_stage2
else self.model.predict_n_sec(cond_cont, cond_cat, stage1_ctx)
)
l_nsec = torch.zeros((), device=device)
nsec_acc = torch.zeros((), device=device)
if n_sec_logits is not None:
l_nsec = F.cross_entropy(n_sec_logits, n_sec)
nsec_acc = (n_sec_logits.argmax(dim=-1) == n_sec).float().mean()
# On a non-generator-step batch with no n_sec_head on this stage
# (n_sec now defaults to stage 2, decision 1), there's nothing for
# the generator optimizer to do this batch — g_loss would otherwise
# be a graph-less zero tensor, which .backward() rejects outright.
skip_g_step = not did_g_step and n_sec_logits is None
if did_g_step:
g_loss_adv = generator_loss(critic_fn, fake)
g_loss = (
self.spec.lambda_weight * g_loss_adv + self.spec.n_sec_lambda * l_nsec
)
else:
g_loss_adv = torch.zeros((), device=device)
g_loss = self.spec.n_sec_lambda * l_nsec
if skip_g_step:
grad_norm_g = torch.zeros(())
else:
self.optimizer.zero_grad()
g_loss.backward()
grad_norm_g = torch.nn.utils.clip_grad_norm_(self.g_params, 1.0)
self.optimizer.step()
if did_g_step:
self.lr_sched.step()
if self.ema_model is not None:
_update_ema(self.ema_model, self.model, self.ema_decay)
return {
"d_loss": d_loss.item(),
"g_loss": g_loss_adv.item(),
"wasserstein": wasserstein.item(),
"gp_loss": gp.item(),
"loss_nsec": l_nsec.item(),
"nsec_acc": nsec_acc.item(),
"did_g_step": did_g_step,
"grad_norm": grad_norm_d.item() + grad_norm_g.item(),
"grad_norm_d": grad_norm_d.item(),
"grad_norm_g": grad_norm_g.item(),
"grad_norm_type_slice": grad_probe.get("type", 0.0),
"grad_norm_cont_slice": grad_probe.get("cont", 0.0),
"lr": self.optimizer.param_groups[0]["lr"],
"critic_lr": self.optimizer_d.param_groups[0]["lr"],
}
# --- reporting ------------------------------------------------------
def batch_loss(self, stats: dict) -> float:
return stats["d_loss"] + stats["g_loss"]
def summary(self, means: dict) -> str:
return (
f"{self.name}[d={means.get('d_loss', 0.0):.3f} "
f"g={means.get('g_loss', 0.0):.3f}]"
)
def val_objective(
self, train_means: dict, val_means: dict, marginal_kl: float
) -> float:
"""No monotone per-batch WGAN loss fit for averaging, so
best-checkpoint selection uses the real marginal-KL signal when
`validate_marginals` produced one, and falls back to this epoch's own
Wasserstein-distance magnitude otherwise.
Behavior change vs. every run up to v0.3.0: the pre-refactor code
meant to do exactly this, but its 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 marginal KL
was recorded in the metrics row without ever influencing `best.pt`.
Runs from before this commit therefore selected their best checkpoint
on the non-adversarial stages alone."""
if math.isfinite(marginal_kl):
return marginal_kl
return abs(train_means.get("wasserstein", 0.0))
# --- state ----------------------------------------------------------
def _extra_state(self) -> dict:
return {
"critic": self.critic.state_dict(),
"optimizer_d": self.optimizer_d.state_dict(),
}
def _load_extra_state(self, sd: dict) -> None:
self.critic.load_state_dict(sd["critic"])
self.optimizer_d.load_state_dict(sd["optimizer_d"])
def _resume_extra_lr(self, lr: float) -> None:
resumed_critic_lr = self.critic_lr if self.critic_lr > 0 else lr
for group in self.optimizer_d.param_groups:
group["lr"] = resumed_critic_lr
def build_stage_trainers(
cfg: dict,
models: dict[str, torch.nn.Module | None],
critics: dict[str, torch.nn.Module | None],
device: torch.device,
total_train_batches: int,
) -> dict[str, StageTrainer]:
"""One trainer per active stage — `models[name] is None` means that stage
is `active = false` and is simply never constructed."""
trainers: dict[str, StageTrainer] = {}
for name, is_stage2 in (("stage1", False), ("stage2", True)):
model = models.get(name)
if model is None:
continue
spec = StageSpec.from_config(cfg, name, is_stage2, max(total_train_batches, 1))
if spec.generator == "wgan":
critic = critics.get(name)
assert critic is not None, (
f"{name}_model.generator='wgan' requires a critic (see "
"giant.model.network.build_critics)"
)
trainers[name] = WGANStageTrainer(spec, model, critic, device)
else:
trainers[name] = FlowDDPMStageTrainer(spec, model, device)
return trainers
+2 -2
View File
@@ -68,8 +68,8 @@ 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)
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
+30 -50
View File
@@ -1,4 +1,4 @@
"""Tests for giant/train.py."""
"""Tests for giant/training/."""
import copy
import csv
@@ -18,14 +18,19 @@ from giant.constants import (
X_DIM,
)
from giant.model.network import build_critics, build_models
from giant.train import (
from giant.training import (
FlowDDPMStageTrainer,
StageSpec,
WGANStageTrainer,
build_stage_trainers,
train,
)
from giant.training.metrics import _wandb_run_config
from giant.training.stage2_inputs import (
_ar_has_prev,
_assemble_stage2_ar_inputs,
_assemble_stage2_ar_target,
_assemble_stage2_real,
_build_stage_trainers,
_gumbel_tau,
_relax_onehot_type_slice,
_remaining_energy_fraction,
@@ -33,8 +38,6 @@ from giant.train import (
_stage2_tf_prob,
_stick_fraction,
_type_repr,
_wandb_run_config,
train,
)
PDG_VOCAB = 6
@@ -558,9 +561,9 @@ def test_metrics_csv_columns_are_stage_prefixed():
out_dir = Path(tmp) / "run"
_run_train(cfg, out_dir)
header = (out_dir / "metrics.csv").read_text().splitlines()[0].split(",")
assert "stage1_train_loss" in header
assert "stage2_train_d_loss" in header
assert "val_loss" in header
assert "stage1/train/loss" in header
assert "stage2/train/d_loss" in header
assert "val/loss" in header
assert "epoch" in header
@@ -574,22 +577,16 @@ def test_wgan_stage_trainer_skips_generator_step_when_no_grad_this_batch():
models = build_models(model_config)
critics = build_critics(model_config)
assert models["stage1"] is not None and critics["stage1"] is not None
trainer = WGANStageTrainer(
spec = StageSpec(
name="stage1",
model=models["stage1"],
critic=critics["stage1"],
is_stage2=False,
lambda_weight=1.0,
n_sec_lambda=0.1,
generator="wgan",
n_critic=1000, # never a generator step in this test
gp_weight=10.0,
lr=3e-4,
critic_lr=0.0,
ema_decay=0.0,
warmup_epochs=0,
epochs=1,
steps_per_epoch=4,
device=torch.device("cpu"),
)
trainer = WGANStageTrainer(
spec, models["stage1"], critics["stage1"], torch.device("cpu")
)
assert trainer.model.n_sec_head is None
batch = _fake_batches(1, 8)[0]
@@ -598,28 +595,11 @@ def test_wgan_stage_trainer_skips_generator_step_when_no_grad_this_batch():
def test_flow_stage_trainer_ddpm_not_implemented_for_stage2():
spec = StageSpec(
name="stage2", is_stage2=True, generator="ddpm", ddpm_n_steps=50, ema_decay=0.0
)
with pytest.raises(NotImplementedError):
FlowDDPMStageTrainer(
name="stage2",
model=torch.nn.Linear(1, 1),
is_stage2=True,
generator="ddpm",
lambda_weight=1.0,
n_sec_lambda=0.1,
lambda_balance=0.0,
lambda_proc=0.0,
lambda_entropy=0.0,
gumbel_tau_start=1.0,
gumbel_tau_end=0.1,
lr=3e-4,
weight_decay=0.01,
ema_decay=0.0,
warmup_epochs=0,
epochs=1,
steps_per_epoch=1,
ddpm_n_steps=50,
device=torch.device("cpu"),
)
FlowDDPMStageTrainer(spec, torch.nn.Linear(1, 1), torch.device("cpu"))
# --- AR trainer wiring (v0.3.0 step 5) --------------------------------------
@@ -649,7 +629,7 @@ def test_build_stage_trainers_ar_scheduled_and_attention_step_runs(
model_config = _model_config(cfg)
models = build_models(model_config)
critics = build_critics(model_config)
trainers = _build_stage_trainers(
trainers = build_stage_trainers(
cfg, models, critics, torch.device("cpu"), total_train_batches=4
)
trainer = trainers["stage2"]
@@ -686,7 +666,7 @@ def test_train_end_to_end_ar_attention_history_scheduled_teacher_forcing(
rows = list(csv.DictReader(f))
assert len(rows) == cfg["train"]["epochs"]
loss_col = (
"stage2_train_g_loss" if stage2_generator == "wgan" else "stage2_train_loss"
"stage2/train/g_loss" if stage2_generator == "wgan" else "stage2/train/loss"
)
assert all(math.isfinite(float(r[loss_col])) for r in rows)
@@ -705,10 +685,10 @@ def test_ar_wgan_onehot_grad_norm_instrumentation_populates_metrics():
_run_train(cfg, out_dir)
with open(out_dir / "metrics.csv", newline="") as f:
rows = list(csv.DictReader(f))
assert "stage2_train_grad_norm_type_slice" in rows[0]
assert "stage2_train_grad_norm_cont_slice" in rows[0]
assert any(float(r["stage2_train_grad_norm_type_slice"]) > 0 for r in rows)
assert any(float(r["stage2_train_grad_norm_cont_slice"]) > 0 for r in rows)
assert "stage2/train/grad_norm_type_slice" in rows[0]
assert "stage2/train/grad_norm_cont_slice" in rows[0]
assert any(float(r["stage2/train/grad_norm_type_slice"]) > 0 for r in rows)
assert any(float(r["stage2/train/grad_norm_cont_slice"]) > 0 for r in rows)
def test_wgan_onehot_one_shot_also_gets_grad_norm_instrumentation():
@@ -721,8 +701,8 @@ def test_wgan_onehot_one_shot_also_gets_grad_norm_instrumentation():
_run_train(cfg, out_dir)
with open(out_dir / "metrics.csv", newline="") as f:
rows = list(csv.DictReader(f))
assert any(float(r["stage2_train_grad_norm_type_slice"]) > 0 for r in rows)
assert any(float(r["stage2_train_grad_norm_cont_slice"]) > 0 for r in rows)
assert any(float(r["stage2/train/grad_norm_type_slice"]) > 0 for r in rows)
assert any(float(r["stage2/train/grad_norm_cont_slice"]) > 0 for r in rows)
def test_wgan_physical_omits_grad_norm_slice_columns():
@@ -731,5 +711,5 @@ def test_wgan_physical_omits_grad_norm_slice_columns():
out_dir = Path(tmp) / "run"
_run_train(cfg, out_dir)
header = (out_dir / "metrics.csv").read_text().splitlines()[0].split(",")
assert "stage2_train_grad_norm_type_slice" not in header
assert "stage2_train_grad_norm_cont_slice" not in header
assert "stage2/train/grad_norm_type_slice" not in header
assert "stage2/train/grad_norm_cont_slice" not in header