diff --git a/giant/config.py b/giant/config.py index 37d91b1..d021693 100644 --- a/giant/config.py +++ b/giant/config.py @@ -339,6 +339,29 @@ def _router_candidate(train, model): return f"r-{router['type']}{router['n_experts']}" +def _router_flag_candidate(field, token_map): + """Candidate factory for a boolean `model.router` sub-field. + + Gated on `router.enabled` like `_router_candidate` (a disabled router's + sub-fields are meaningless), then omitted unless `field` differs from + its DEFAULT_CONFIG value — same "only show non-default" rule as every + other candidate. `token_map` need only cover the non-default value(s), + since the default value always yields None. + """ + + def _candidate(train, model): + router = model["router"] + default_router = DEFAULT_CONFIG["model"]["router"] + if router["enabled"] == default_router["enabled"]: + return None + value = router[field] + if value == default_router[field]: + return None + return token_map[value] + + return _candidate + + def _conditioning_candidate(train, model): if model["conditioning"] == DEFAULT_CONFIG["model"]["conditioning"]: return None @@ -360,6 +383,10 @@ def _default_field_candidate(section_key, field, prefix): _OUT_DIR_NAME_CANDIDATES = [ ("mode", _mode_candidate), ("router", _router_candidate), + ("gumbel", _router_flag_candidate("gumbel", {True: "gum"})), + ("learn_centers", _router_flag_candidate("learn_centers", {False: "nolc"})), + ("learn_width", _router_flag_candidate("learn_width", {True: "lw"})), + ("learn_temperature", _router_flag_candidate("learn_temperature", {True: "lt"})), ("conditioning", _conditioning_candidate), ("hidden_dim", _default_field_candidate("model", "hidden_dim", "h")), ("n_blocks", _default_field_candidate("model", "n_blocks", "b")), diff --git a/tests/test_config.py b/tests/test_config.py index 3795367..30f6e41 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -207,6 +207,64 @@ def test_default_out_dir_name_router_disabled_omitted_even_if_subfields_nondefau assert gconfig.default_out_dir_name(cfg, now=_NOW) == "20260729_1430" +def test_default_out_dir_name_router_gumbel_shown_when_enabled(): + cfg = _default_cfg( + router={"enabled": True, "type": "energy", "n_experts": 8, "gumbel": True} + ) + assert gconfig.default_out_dir_name(cfg, now=_NOW) == "20260729_1430_r-energy8_gum" + + +def test_default_out_dir_name_router_gumbel_omitted_when_router_disabled(): + cfg = _default_cfg( + router={"enabled": False, "type": "energy", "n_experts": 8, "gumbel": True} + ) + assert gconfig.default_out_dir_name(cfg, now=_NOW) == "20260729_1430" + + +def test_default_out_dir_name_router_learn_centers_shown_only_when_disabled(): + cfg_default = _default_cfg( + router={"enabled": True, "type": "energy", "n_experts": 8} + ) + assert ( + gconfig.default_out_dir_name(cfg_default, now=_NOW) == "20260729_1430_r-energy8" + ) + + cfg_off = _default_cfg( + router={ + "enabled": True, + "type": "energy", + "n_experts": 8, + "learn_centers": False, + } + ) + assert ( + gconfig.default_out_dir_name(cfg_off, now=_NOW) + == "20260729_1430_r-energy8_nolc" + ) + + +def test_default_out_dir_name_router_learn_width_and_temperature_shown(): + cfg = _default_cfg( + router={ + "enabled": True, + "type": "energy", + "n_experts": 8, + "learn_width": True, + } + ) + assert gconfig.default_out_dir_name(cfg, now=_NOW) == "20260729_1430_r-energy8_lw" + + cfg2 = _default_cfg( + router={ + "enabled": True, + "type": "energy", + "n_experts": 8, + "learn_temperature": True, + } + ) + assert gconfig.default_out_dir_name(cfg2, now=_NOW) == "20260729_1430_r-energy8_lt" + + def test_default_out_dir_name_mode_shown_bare_no_prefix(): cfg = _default_cfg(mode="wgan") assert gconfig.default_out_dir_name(cfg, now=_NOW) == "20260729_1430_wgan"