From da717971b6b7ed46429936adffd6e6d043630898 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Thu, 13 Aug 2026 15:36:26 +0200 Subject: [PATCH] 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 --- giant/cli.py | 32 +++++++++++++++++++++++ giant/config.py | 7 +++++ giant/model/builders.py | 8 +++--- tests/test_config.py | 17 +++++++++++++ tests/test_config_consumed_keys.py | 6 ----- tests/test_network.py | 41 ++++++++++++++++++++++++++++++ 6 files changed, 101 insertions(+), 10 deletions(-) diff --git a/giant/cli.py b/giant/cli.py index 764399b..853c778 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -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) diff --git a/giant/config.py b/giant/config.py index b022796..24b2850 100644 --- a/giant/config.py +++ b/giant/config.py @@ -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",)), ) diff --git a/giant/model/builders.py b/giant/model/builders.py index 887a573..ea169db 100644 --- a/giant/model/builders.py +++ b/giant/model/builders.py @@ -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", diff --git a/tests/test_config.py b/tests/test_config.py index 8e7751a..400410e 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -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) # --------------------------------------------------------------------------- diff --git a/tests/test_config_consumed_keys.py b/tests/test_config_consumed_keys.py index 05527bf..e21a216 100644 --- a/tests/test_config_consumed_keys.py +++ b/tests/test_config_consumed_keys.py @@ -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" diff --git a/tests/test_network.py b/tests/test_network.py index 7857d39..bf68efa 100644 --- a/tests/test_network.py +++ b/tests/test_network.py @@ -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