Honour wgan.critic_hidden_dim/critic_n_res_blocks in build_critics (gitea #28)

build_critics always sized a WGAN critic off the generator's own
hidden_dim/n_res_blocks, silently discarding the documented 0=inherit
sentinel on stage{1,2}_model.wgan.critic_hidden_dim/critic_n_res_blocks
(the same convention critic_lr already honoured). Now both keys are read
with the 0 -> inherit fallback, and stage-scoped-only CLI flags
(--stage{1,2}-critic-hidden-dim/--stage{1,2}-critic-n-res-blocks) are
added -- no shared alias, since critic sizing is an architectural
per-stage knob like --hidden-dim/--n-res-blocks, not a shared training
hyperparameter like --n-critic/--gp-weight/--noise-dim/--critic-lr.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
2026-08-13 15:36:26 +02:00
parent c3fc768b40
commit da717971b6
6 changed files with 101 additions and 10 deletions
+32
View File
@@ -459,6 +459,34 @@ def train(
Optional[float],
typer.Option("--stage2-critic-lr", help="Overrides --critic-lr for stage 2 only"),
] = None,
stage1_critic_hidden_dim: Annotated[
Optional[int],
typer.Option(
"--stage1-critic-hidden-dim",
help="WGAN-GP (--mode wgan only): critic width for stage 1 (default: same as generator's hidden_dim)",
),
] = None,
stage1_critic_n_res_blocks: Annotated[
Optional[int],
typer.Option(
"--stage1-critic-n-res-blocks",
help="WGAN-GP (--mode wgan only): critic depth for stage 1 (default: same as generator's n_res_blocks)",
),
] = None,
stage2_critic_hidden_dim: Annotated[
Optional[int],
typer.Option(
"--stage2-critic-hidden-dim",
help="WGAN-GP (--mode wgan only): critic width for stage 2 (default: same as generator's hidden_dim)",
),
] = None,
stage2_critic_n_res_blocks: Annotated[
Optional[int],
typer.Option(
"--stage2-critic-n-res-blocks",
help="WGAN-GP (--mode wgan only): critic depth for stage 2 (default: same as generator's n_res_blocks)",
),
] = None,
val_fraction: Annotated[Optional[float], typer.Option("--val-fraction", "-f")] = None,
seed: Annotated[
Optional[int],
@@ -619,6 +647,10 @@ def train(
"stage2_gp_weight": stage2_gp_weight,
"stage2_noise_dim": stage2_noise_dim,
"stage2_critic_lr": stage2_critic_lr,
"stage1_critic_hidden_dim": stage1_critic_hidden_dim,
"stage1_critic_n_res_blocks": stage1_critic_n_res_blocks,
"stage2_critic_hidden_dim": stage2_critic_hidden_dim,
"stage2_critic_n_res_blocks": stage2_critic_n_res_blocks,
}
overrides = gconfig.overrides_from_flags(flag_values)
+7
View File
@@ -979,6 +979,13 @@ FLAG_SPECS: tuple[FlagSpec, ...] = (
FlagSpec("critic_lr", ("stage1_model.wgan.critic_lr", "stage2_model.wgan.critic_lr"), precedence=0),
FlagSpec("stage1_critic_lr", ("stage1_model.wgan.critic_lr",), precedence=1),
FlagSpec("stage2_critic_lr", ("stage2_model.wgan.critic_lr",), precedence=1),
# Critic sizing: stage-scoped only, no shared alias — this is an
# architectural per-stage knob like hidden_dim/n_res_blocks above, not a
# shared training hyperparameter like the wgan knobs above it.
FlagSpec("stage1_critic_hidden_dim", ("stage1_model.wgan.critic_hidden_dim",)),
FlagSpec("stage1_critic_n_res_blocks", ("stage1_model.wgan.critic_n_res_blocks",)),
FlagSpec("stage2_critic_hidden_dim", ("stage2_model.wgan.critic_hidden_dim",)),
FlagSpec("stage2_critic_n_res_blocks", ("stage2_model.wgan.critic_n_res_blocks",)),
)
+4 -4
View File
@@ -172,8 +172,8 @@ def build_critics(model_config: dict) -> dict[str, nn.Module | None]:
particle_cfg=particle_cfg,
material_cfg=material_cfg,
in_dim=X_DIM,
hidden_dim=s1_spec.hidden_dim,
n_res_blocks=s1_spec.n_res_blocks,
hidden_dim=s1_spec.wgan.critic_hidden_dim or s1_spec.hidden_dim,
n_res_blocks=s1_spec.wgan.critic_n_res_blocks or s1_spec.n_res_blocks,
cond_out_dim=cond_out_dim,
dropout=s1_spec.dropout,
stage="stage1",
@@ -189,8 +189,8 @@ def build_critics(model_config: dict) -> dict[str, nn.Module | None]:
particle_cfg=particle_cfg,
material_cfg=material_cfg,
in_dim=in_dim,
hidden_dim=s2_spec.hidden_dim,
n_res_blocks=s2_spec.n_res_blocks,
hidden_dim=s2_spec.wgan.critic_hidden_dim or s2_spec.hidden_dim,
n_res_blocks=s2_spec.wgan.critic_n_res_blocks or s2_spec.n_res_blocks,
cond_out_dim=cond_out_dim,
dropout=s2_spec.dropout,
stage="stage2",
+17
View File
@@ -1022,6 +1022,23 @@ def test_overrides_from_flags_wgan_knobs_split_per_stage(shared, stage1_specific
assert overrides["stage2_model"]["wgan"][path_key] == 2.5
@pytest.mark.parametrize(
("stage_flag", "stage_model", "path_key"),
[
("stage1_critic_hidden_dim", "stage1_model", "critic_hidden_dim"),
("stage1_critic_n_res_blocks", "stage1_model", "critic_n_res_blocks"),
("stage2_critic_hidden_dim", "stage2_model", "critic_hidden_dim"),
("stage2_critic_n_res_blocks", "stage2_model", "critic_n_res_blocks"),
],
)
def test_overrides_from_flags_critic_sizing_is_stage_scoped_only(stage_flag, stage_model, path_key):
"""critic_hidden_dim/critic_n_res_blocks are architectural per-stage
knobs (gitea #28) — unlike n_critic/gp_weight/noise_dim/critic_lr above,
there is deliberately no shared alias that fans out to both stages."""
overrides = gconfig.overrides_from_flags({stage_flag: 32})
assert overrides == {stage_model: {"wgan": {path_key: 32}}}
# ---------------------------------------------------------------------------
# checkpoint config-mismatch warnings (unchanged surface, still exercised)
# ---------------------------------------------------------------------------
-6
View File
@@ -60,12 +60,6 @@ _KNOWN_UNUSED = {
"read by any build/train consumer file since only 'truth' can pass "
"validation — see Issue 16 for the real implementation"
),
"stage1_model.wgan.critic_hidden_dim": (
"issues.md Issue 2 — build_critics always sizes the critic off the generator's own hidden_dim, never this key"
),
"stage1_model.wgan.critic_n_res_blocks": ("issues.md Issue 2 — same as critic_hidden_dim"),
"stage2_model.wgan.critic_hidden_dim": ("issues.md Issue 2 — same as stage1_model.wgan.critic_hidden_dim"),
"stage2_model.wgan.critic_n_res_blocks": ("issues.md Issue 2 — same as stage1_model.wgan.critic_hidden_dim"),
"stage2_model.autoregressive.order": (
"issues.md Issue 4 — validate_config checks history/teacher_forcing "
"but never order, and nothing reads it either"
+41
View File
@@ -708,3 +708,44 @@ def test_build_critics_omitted_particle_type_matches_default_config():
# this also confirms the critic was actually built in onehot mode by
# default, not silently falling back to physical.
assert onehot_in_dim != physical_in_dim
# ── build_critics: critic_hidden_dim/critic_n_res_blocks honoured (gitea #28) ─
def test_build_critics_stage1_critic_hidden_dim_and_n_res_blocks_override_generator_size():
cfg = _minimal_model_config(share_stages=False)
cfg["stage1_model"]["generator"] = "wgan"
cfg["stage1_model"]["hidden_dim"] = 8
cfg["stage1_model"]["n_res_blocks"] = 1
inherited = build_critics(cfg)["stage1"]
assert inherited is not None
assert inherited.input_proj.out_features == 8
assert len(inherited.blocks) == 1
cfg["stage1_model"]["wgan"]["critic_hidden_dim"] = 16
cfg["stage1_model"]["wgan"]["critic_n_res_blocks"] = 3
overridden = build_critics(cfg)["stage1"]
assert overridden is not None
assert overridden.input_proj.out_features == 16
assert len(overridden.blocks) == 3
def test_build_critics_stage2_critic_hidden_dim_and_n_res_blocks_override_generator_size():
cfg = _minimal_model_config(share_stages=False)
cfg["stage2_model"]["generator"] = "wgan"
cfg["stage2_model"]["hidden_dim"] = 8
cfg["stage2_model"]["n_res_blocks"] = 1
inherited = build_critics(cfg)["stage2"]
assert inherited is not None
assert inherited.input_proj.out_features == 8
assert len(inherited.blocks) == 1
cfg["stage2_model"]["wgan"]["critic_hidden_dim"] = 16
cfg["stage2_model"]["wgan"]["critic_n_res_blocks"] = 3
overridden = build_critics(cfg)["stage2"]
assert overridden is not None
assert overridden.input_proj.out_features == 16
assert len(overridden.blocks) == 3