Make config dataclasses the single source of truth for DEFAULT_CONFIG
CI / Format (ruff format) (push) Successful in 31s
CI / Lint (ruff check) (push) Successful in 31s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 44s
CI / Type check (ty) (push) Successful in 46s
CI / Format (ruff format) (pull_request) Successful in 39s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 39s
CI / Tests (pull_request) Successful in 3m38s
CI / Tests (push) Successful in 3m45s
CI / Format (ruff format) (push) Successful in 31s
CI / Lint (ruff check) (push) Successful in 31s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 44s
CI / Type check (ty) (push) Successful in 46s
CI / Format (ruff format) (pull_request) Successful in 39s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 39s
CI / Tests (pull_request) Successful in 3m38s
CI / Tests (push) Successful in 3m45s
DEFAULT_CONFIG and build_models/build_critics/StageSpec.from_config's inline .get(key, default) fallbacks had already drifted: two keys (stage2_model.decoder, stage2_model.particle_type.target) resolved differently depending on whether a config dict came from merge_cli_overrides (fully populated, correct) or was hand-built and partial (fell back to stale v0.2-shaped literals). Introduce frozen dataclasses (GiantConfig and its nested blocks) in giant/config.py as the actual single declaration of every default; DEFAULT_CONFIG is now generated from them instead of hand-maintained, and build_models, build_critics, and StageSpec.from_config consume the dataclasses instead of duplicating literal fallbacks, so this class of drift can't recur. Router/n_sec sub-blocks keep an `extra` catch-all for their genuinely dynamic keys (composed-router axes, runtime-seeded centers_init, legacy_owner). Fixing the fallback surfaced the same latent bug in two existing partial-config callers that had been silently depending on it: a test fixture in test_train.py and scripts/warm_setup_cache.py's minimal cfg (now merged against DEFAULT_CONFIG instead of hand-rolled, closing the gap for good). See issues.md Issue 1. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
+21
-1
@@ -265,6 +265,12 @@ def _base_cfg():
|
||||
"k_max": K_MAX,
|
||||
"context_dim": 16,
|
||||
"n_sec": {"mode": "head", "lambda": 0.1},
|
||||
# Explicit, not relying on the fallback default (which is
|
||||
# "onehot", matching DEFAULT_CONFIG — see issues.md Issue 1):
|
||||
# the "physical"-labelled cases below (and this fixture's own
|
||||
# comment history) intend this as the base "physical" case,
|
||||
# with "*_onehot"/"*_embedding" cases opting in explicitly.
|
||||
"particle_type": {"target": "physical", "lambda": 1.0},
|
||||
"flow": {"time_dim": 16},
|
||||
"ddpm": {"time_dim": 16, "n_steps": 50},
|
||||
"wgan": {
|
||||
@@ -566,6 +572,20 @@ def test_flow_stage_trainer_ddpm_not_implemented_for_stage2():
|
||||
FlowDDPMStageTrainer(spec, torch.nn.Linear(1, 1), torch.device("cpu"))
|
||||
|
||||
|
||||
def test_stage_spec_from_config_omitted_decoder_and_particle_type_match_default_config():
|
||||
"""Regression for issues.md Issue 1: StageSpec.from_config's own fallback
|
||||
defaults for stage2_model.decoder/particle_type must equal
|
||||
DEFAULT_CONFIG's ("autoregressive" / "onehot"), not the old, now-wrong
|
||||
("one_shot" / "physical") literals a .get(key, default) call used to
|
||||
supply when a hand-built cfg omitted these keys."""
|
||||
cfg = _base_cfg()
|
||||
del cfg["stage2_model"]["decoder"]
|
||||
del cfg["stage2_model"]["particle_type"]
|
||||
spec = StageSpec.from_config(cfg, "stage2", is_stage2=True, steps_per_epoch=1)
|
||||
assert spec.decoder == "autoregressive"
|
||||
assert spec.particle_type.target == "onehot"
|
||||
|
||||
|
||||
# --- AR trainer wiring (v0.3.0 step 5) --------------------------------------
|
||||
|
||||
|
||||
@@ -664,7 +684,7 @@ def test_wgan_onehot_one_shot_also_gets_grad_norm_instrumentation():
|
||||
|
||||
|
||||
def test_wgan_physical_omits_grad_norm_slice_columns():
|
||||
cfg = _base_cfg() # default stage2_model has no particle_type -> "physical"
|
||||
cfg = _base_cfg() # _base_cfg's stage2_model.particle_type.target is "physical"
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
out_dir = Path(tmp) / "run"
|
||||
_run_train(cfg, out_dir)
|
||||
|
||||
Reference in New Issue
Block a user