"""Tests for `giant train`'s stage-prefixed CLI flags: --stage1-*/--stage2-* must independently override each stage's config block, and must take precedence over the older shared flags (--mode/--hidden-dim/--n-critic/... ) that still apply the same value to both stages for backward compatibility.""" from __future__ import annotations from pathlib import Path from typer.testing import CliRunner import giant.cli as cli runner = CliRunner() def _invoke_and_capture_cfg(monkeypatch, tmp_path: Path, args: list[str]) -> dict: captured: dict = {} def _fake_run_train_job(*, data, cfg, out_dir, **kwargs): captured["cfg"] = cfg monkeypatch.setattr(cli, "run_train_job", _fake_run_train_job) result = runner.invoke( cli.app, ["train", "dummy.parquet", "--out", str(tmp_path / "run")] + args, ) assert result.exit_code == 0, result.output return captured["cfg"] def test_stage_prefixed_generator_overrides_shared_mode(monkeypatch, tmp_path): cfg = _invoke_and_capture_cfg( monkeypatch, tmp_path, ["--mode", "wgan", "--stage1-generator", "flow"], ) assert cfg["stage1_model"]["generator"] == "flow" assert cfg["stage2_model"]["generator"] == "wgan" def test_stage2_only_knobs(monkeypatch, tmp_path): cfg = _invoke_and_capture_cfg( monkeypatch, tmp_path, [ "--stage2-decoder", "one_shot", "--stage2-k-max", "8", "--stage2-hidden-dim", "32", "--stage2-context-dim", "16", "--stage2-stage1-context", "sampled", ], ) assert cfg["stage2_model"]["decoder"] == "one_shot" assert cfg["stage2_model"]["k_max"] == 8 assert cfg["stage2_model"]["hidden_dim"] == 32 assert cfg["stage2_model"]["context_dim"] == 16 assert cfg["stage2_model"]["stage1_context"] == "sampled" # untouched stage1 defaults assert cfg["stage1_model"]["hidden_dim"] == 256 def test_stage1_hidden_dim_flag_overrides_legacy_hidden_dim_flag(monkeypatch, tmp_path): cfg = _invoke_and_capture_cfg( monkeypatch, tmp_path, ["--hidden-dim", "64", "--stage1-hidden-dim", "128"], ) assert cfg["stage1_model"]["hidden_dim"] == 128 def test_wgan_knobs_split_per_stage(monkeypatch, tmp_path): cfg = _invoke_and_capture_cfg( monkeypatch, tmp_path, [ "--mode", "wgan", "--n-critic", "5", "--stage1-n-critic", "3", "--stage2-gp-weight", "2.5", ], ) assert cfg["stage1_model"]["wgan"]["n_critic"] == 3 assert cfg["stage1_model"]["wgan"]["gp_weight"] == 10.0 assert cfg["stage2_model"]["wgan"]["n_critic"] == 5 assert cfg["stage2_model"]["wgan"]["gp_weight"] == 2.5