Add gumbel/learn_centers/learn_width/learn_temperature to out-dir naming
CI / Format (ruff format) (push) Successful in 25s
CI / Lint (ruff check) (push) Successful in 27s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 29s
CI / Type check (ty) (push) Successful in 31s
CI / Format (ruff format) (pull_request) Successful in 37s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 38s
CI / Tests (push) Successful in 1m34s
CI / Tests (pull_request) Successful in 1m30s
CI / Format (ruff format) (push) Successful in 25s
CI / Lint (ruff check) (push) Successful in 27s
CI / Sync project version with tag (push) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 29s
CI / Type check (ty) (push) Successful in 31s
CI / Format (ruff format) (pull_request) Successful in 37s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 38s
CI / Tests (push) Successful in 1m34s
CI / Tests (pull_request) Successful in 1m30s
Extends default_out_dir_name's non-default-field convention to the router's new gumbel combine-weight flag and its learnable-knob toggles, so gumbel sweep configs (learn_centers on/off, learn_width, learn_temperature) resolve to distinguishable checkpoint directory names instead of colliding on the same r-<type><n> token. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -339,6 +339,29 @@ def _router_candidate(train, model):
|
|||||||
return f"r-{router['type']}{router['n_experts']}"
|
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):
|
def _conditioning_candidate(train, model):
|
||||||
if model["conditioning"] == DEFAULT_CONFIG["model"]["conditioning"]:
|
if model["conditioning"] == DEFAULT_CONFIG["model"]["conditioning"]:
|
||||||
return None
|
return None
|
||||||
@@ -360,6 +383,10 @@ def _default_field_candidate(section_key, field, prefix):
|
|||||||
_OUT_DIR_NAME_CANDIDATES = [
|
_OUT_DIR_NAME_CANDIDATES = [
|
||||||
("mode", _mode_candidate),
|
("mode", _mode_candidate),
|
||||||
("router", _router_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),
|
("conditioning", _conditioning_candidate),
|
||||||
("hidden_dim", _default_field_candidate("model", "hidden_dim", "h")),
|
("hidden_dim", _default_field_candidate("model", "hidden_dim", "h")),
|
||||||
("n_blocks", _default_field_candidate("model", "n_blocks", "b")),
|
("n_blocks", _default_field_candidate("model", "n_blocks", "b")),
|
||||||
|
|||||||
@@ -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"
|
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():
|
def test_default_out_dir_name_mode_shown_bare_no_prefix():
|
||||||
cfg = _default_cfg(mode="wgan")
|
cfg = _default_cfg(mode="wgan")
|
||||||
assert gconfig.default_out_dir_name(cfg, now=_NOW) == "20260729_1430_wgan"
|
assert gconfig.default_out_dir_name(cfg, now=_NOW) == "20260729_1430_wgan"
|
||||||
|
|||||||
Reference in New Issue
Block a user