mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Keep realistic hot restart WandB identity stable (#3070)
This commit is contained in:
+3
-3
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user