Keep realistic hot restart WandB identity stable (#3070)

This commit is contained in:
fzyzcjy
2026-09-26 20:43:22 +08:00
committed by GitHub
parent 87cbb8ac80
commit 2b42fdb24c
2 changed files with 23 additions and 7 deletions
@@ -83,7 +83,7 @@ def run_ci(
mean_interval_seconds_of_cell_type={_HOT_RESTART_CELL_TYPE: hot_restart_interval_seconds},
create_forms=create_forms,
get_virtual_cells=lambda: _create_virtual_cells_before(hot_restart_form.value),
build_extra_train_args=_build_train_args,
build_extra_train_args=lambda dump_dir: _build_train_args(dump_dir, wandb_run_id=config.run_id),
enable_fault_tolerance=False,
)
@@ -115,8 +115,8 @@ def run_ci(
print(f"Hot restart realistic gsm8k test PASSED (seed={seed}, rollouts={num_rollout})")
def _build_train_args(dump_dir: str) -> str:
return build_checkpoint_args(dump_dir) + "--ci-disable-weight-update-checker "
def _build_train_args(dump_dir: str, *, wandb_run_id: str) -> str:
return build_checkpoint_args(dump_dir) + f"--wandb-run-id {wandb_run_id} " + "--ci-disable-weight-update-checker "
def assert_no_take_over_attempt_failed(events: list[Event]) -> None:
@@ -40,11 +40,27 @@ class TestTheRecipeIsTheOneFtConverges:
assert scenario.DEFAULT_METRIC_THRESHOLD is scenario_realistic_gsm8k.DEFAULT_METRIC_THRESHOLD
assert scenario.DEFAULT_NUM_ROLLOUT is scenario_realistic_gsm8k.DEFAULT_NUM_ROLLOUT
def test_this_scenario_spells_no_training_arguments_of_its_own_beyond_its_checkpoints(self):
"""Hot restart keeps the FT recipe except for checkpointing and its incompatible tensor checker."""
declared = [one for one in shlex.split(scenario._build_train_args("/dumps")) if one.startswith("--")]
def test_this_scenario_only_adds_arguments_its_take_over_requires(self):
"""Hot restart keeps the FT recipe except for checkpointing, tracking identity, and its tensor checker."""
declared = [
one
for one in shlex.split(scenario._build_train_args("/dumps", wandb_run_id="run-one"))
if one.startswith("--")
]
assert sorted(declared) == ["--ci-disable-weight-update-checker", "--load", "--save", "--save-interval"]
assert sorted(declared) == [
"--ci-disable-weight-update-checker",
"--load",
"--save",
"--save-interval",
"--wandb-run-id",
]
def test_the_run_id_is_the_identity_of_the_deployment_this_test_restarts(self):
"""Every process receives the outer config ID instead of an identity derived from another subsystem."""
argv = shlex.split(scenario._build_train_args("/dumps/another-id", wandb_run_id="run-one"))
assert ArgvManipulator.get(argv, "--wandb-run-id") == ["run-one"]
def test_the_shared_recipe_can_keep_its_api_without_enabling_training_ft(self):
"""Hot restart uses the cell API for injection without combining with automatic FT recovery."""