"""Tests for giant/train.py helpers.""" from giant.train import _gumbel_tau, _wandb_run_config def test_gumbel_tau_at_step_zero_is_start(): assert _gumbel_tau(0, 1000, 1.0, 0.1) == 1.0 def test_gumbel_tau_at_total_steps_is_end(): assert abs(_gumbel_tau(1000, 1000, 1.0, 0.1) - 0.1) < 1e-9 def test_gumbel_tau_interpolates_linearly_midway(): assert abs(_gumbel_tau(500, 1000, 1.0, 0.1) - 0.55) < 1e-9 def test_gumbel_tau_clamps_beyond_total_steps(): assert _gumbel_tau(5000, 1000, 1.0, 0.1) == _gumbel_tau(1000, 1000, 1.0, 0.1) def test_gumbel_tau_handles_zero_total_steps(): # total_steps=0 is guarded to 1 internally: step=0 gives zero progress # (still tau_start), any step>=1 immediately clamps to full progress. assert _gumbel_tau(0, 0, 1.0, 0.1) == 1.0 assert abs(_gumbel_tau(1, 0, 1.0, 0.1) - 0.1) < 1e-9 def _base_wandb_kwargs(**overrides): kwargs = dict( mode="flow", epochs=30, lr=3e-4, warmup_epochs=3, weight_decay=0.01, ema_decay=0.9999, lambda_nsec=0.1, lambda_s2=1.0, lambda_balance=0.035, lambda_proc=0.0, lambda_entropy=0.0, gumbel_tau_start=1.0, gumbel_tau_end=0.1, n_critic=5, gp_weight=10.0, model_config={"router": {"enabled": False}}, stage1_params=100, sec_decoder_params=50, critic_params=0, sec_critic_params=0, total_params=150, ) kwargs.update(overrides) return kwargs def test_wandb_run_config_omits_router_knobs_when_router_disabled(): cfg = _wandb_run_config(**_base_wandb_kwargs()) for key in ( "lambda_balance", "lambda_proc", "lambda_entropy", "gumbel_tau_start", "gumbel_tau_end", ): assert key not in cfg # still present, nested, regardless of router state assert cfg["model"] == {"router": {"enabled": False}} def test_wandb_run_config_includes_router_knobs_when_router_enabled(): cfg = _wandb_run_config( **_base_wandb_kwargs(model_config={"router": {"enabled": True}}) ) assert cfg["lambda_balance"] == 0.035 assert cfg["lambda_proc"] == 0.0 assert cfg["lambda_entropy"] == 0.0 assert cfg["gumbel_tau_start"] == 1.0 assert cfg["gumbel_tau_end"] == 0.1 def test_wandb_run_config_omits_wgan_knobs_when_mode_is_not_wgan(): cfg = _wandb_run_config(**_base_wandb_kwargs(mode="flow")) assert "n_critic" not in cfg assert "gp_weight" not in cfg def test_wandb_run_config_includes_wgan_knobs_when_mode_is_wgan(): cfg = _wandb_run_config(**_base_wandb_kwargs(mode="wgan")) assert cfg["n_critic"] == 5 assert cfg["gp_weight"] == 10.0 def test_wandb_run_config_handles_missing_model_config(): cfg = _wandb_run_config(**_base_wandb_kwargs(model_config=None)) assert cfg["model"] == {} assert "lambda_balance" not in cfg