Assert a hot restart's changed argument reached only the restarted components

Two assertions on the recorded pod commands and the rollout data on disk: the successive orchestrator pods carry the changed flag in their argv and the successive rollout-executor pods in their served --config payload, in generation order, with no other pod carrying any of the values; and each generation's directory holds exactly the rollouts that generation generated per the freeze schedule.
This commit is contained in:
Tom
2026-10-01 15:43:35 +08:00
parent 3bd66e8a62
commit 31a3479884
4 changed files with 133 additions and 3 deletions
+12 -2
View File
@@ -87,7 +87,16 @@ Entries: test_hot_restart_checkpointed.py, test_hot_restart_no_checkpoint.py
an unchanged custom resource
3. Assert process: one trainer rpc boot uuid throughout, answering the take-over's fresh client, and
read once off a whole-release snapshot taken before the first take-over stamped anything
4. Assert redo, measured off the logs, per mode:
4. Assert the changed argument reached the two restarted components and nothing else, off the
commands of every pod the observer saw (recorded once per pod uid):
- the successive rollout-executor pods carry --save-debug-rollout-data generation_0, _1, ... in
their served --config payload
- no pod of any other workload carries any of those values, the orchestrator included: its
--orchestrator-config payload holds the OrchestratorConfig, which has no rollout-executor-only field
- generation k's directory holds exactly the rollouts that generation generated: 0..frozen_0
for generation 0, (saved_{k-1}, frozen_k] for the k-th take-over, (saved_last, 5] for the
last - so the relaunched executor demonstrably ran with its new arguments
5. Assert redo, measured off the logs, per mode:
- checkpointed: one .trash_* per restart; resume point == the pinned save (the snapshot
beside that checkpoint), so the run resumed there, not at step 0; the redone steps are
exactly the pinned (save, frozen step] windows; per-step attempts all 1 or 2
@@ -95,7 +104,7 @@ Entries: test_hot_restart_checkpointed.py, test_hot_restart_no_checkpoint.py
thrown away (steps 0..1, each once) and sharing no step with the log that replaced it; the
surviving log describes each of the 6 steps exactly once; the run still saves after the
restart, past the step it was frozen at
5. Compare: bitwise as in scenario_split_deterministic, engine checksums included, with one
6. Compare: bitwise as in scenario_split_deterministic, engine checksums included, with one
exemption - rollout/weight_version mean/median/max/min. The trainer outlives a take-over, so
its weight update counter keeps counting through the steps the target redoes and stands
ahead of the baseline's at the same step
@@ -106,6 +115,7 @@ over at rollout 0 with the run.
```
- **Why `--save-debug-rollout-data`**: it is read by the rollout executor alone (a `DebugRolloutOnlyConfig` field), so a relaunch that changes it must leave the trainer and inference-controller payloads byte-identical - which the launcher enforces by refusing any other diff - and what it does is visible on disk without touching a single training bit, so the bitwise comparison against the baseline still holds.
- **Why the pod commands and not only the files**: the files prove the restarted executor ran the new value; the commands prove that no other component, the restarted orchestrator included, was handed it.
### `scenario_hot_restart_realistic_gsm8k`
@@ -0,0 +1,37 @@
from collections.abc import Sequence
from pathlib import Path
from tests.e2e.deploy.conftest_deploy.hot_restart.driver import ScheduledFreeze
def assert_generations_recorded_their_steps(
templates: Sequence[str], *, schedule: Sequence[ScheduledFreeze], num_rollouts: int
) -> None:
directories = [Path(template).parent for template in templates]
stray = sorted(set(directories[0].parent.iterdir()) - set(directories))
assert not stray, (
f"{directories[0].parent} holds {[one.name for one in stray]} beside the {len(templates)} generation "
f"directories the take-overs relaunched with, so some executor wrote where nothing relaunched it"
)
for generation, (directory, rollout_ids) in enumerate(
zip(directories, _compute_rollout_ids_of_generation(schedule, num_rollouts=num_rollouts), strict=True)
):
recorded = sorted(int(one.stem) for one in directory.glob("*.pt"))
assert recorded == rollout_ids, (
f"generation {generation} of the rollout executor recorded the rollouts {recorded} under {directory}, "
f"and the freeze schedule has that generation generate exactly {rollout_ids}: the relaunched executor "
f"did not run with the arguments it was relaunched with, or generated steps it should not have"
)
print("every generation of the rollout executor recorded the steps it generated")
def _compute_rollout_ids_of_generation(schedule: Sequence[ScheduledFreeze], *, num_rollouts: int) -> list[list[int]]:
windows: list[list[int]] = []
start = 0
for scheduled in schedule:
windows.append(list(range(start, scheduled.frozen_rollout_id + 1)))
start = 0 if scheduled.saved_iteration is None else scheduled.saved_iteration + 1
windows.append(list(range(start, num_rollouts)))
return windows
@@ -26,6 +26,7 @@ from tests.e2e.deploy.conftest_deploy.hot_restart.assert_redone_from_checkpoint
from tests.e2e.deploy.conftest_deploy.hot_restart.assert_redone_from_scratch import (
assert_unsaved_run_redone_from_scratch,
)
from tests.e2e.deploy.conftest_deploy.hot_restart.assert_rollout_data import assert_generations_recorded_their_steps
from tests.e2e.deploy.conftest_deploy.hot_restart.driver import (
HotRestartDriver,
ScheduledFreeze,
@@ -42,7 +43,10 @@ from tests.e2e.ft.conftest_ft.execution import DATA_DIR, MODEL_DIR
from tests.e2e.ft.conftest_ft.modes import DENSE_MODEL_HF_REPO, DENSE_MODEL_NAME, DENSE_MODEL_TYPE, FTTestMode
from tests.utils.deploy.hot_restart.evidence import TRAIN_STEP_METRIC_KEY, HotRestartEvidence
from tests.utils.soak.core.utils import compute_release_of_config
from tests.utils.soak.deploy.checkers.takeover_scope import assert_take_overs_replaced_only_script
from tests.utils.soak.deploy.checkers.takeover_scope import (
assert_take_overs_carried_rollout_only_args,
assert_take_overs_replaced_only_script,
)
from tests.utils.soak.deploy.utils import compute_checkpoint_dir
from miles.utils.audit_utils.event_logger.logger import EVENTS_DIRNAME
@@ -314,6 +318,8 @@ def _driving_take_overs_of(
yield
driver.assert_all_restarts_happened()
assert_generations_recorded_their_steps(templates, schedule=restart_mode.schedule, num_rollouts=NUM_ROLLOUTS)
assert_take_overs_carried_rollout_only_args(driver.evidence, flag=SAVE_DEBUG_ROLLOUT_DATA_FLAG, values=templates)
# ========================= comparison and assertions ==========================
@@ -8,6 +8,14 @@ from tests.utils.deploy.hot_restart.assert_process import (
from tests.utils.deploy.hot_restart.cluster_observer import ClusterSnapshot, compute_hot_restart_workloads
from tests.utils.deploy.hot_restart.evidence import HotRestartEvidence
from miles.ray.specs.rollout import ROLLOUT_EXECUTOR_POOL_ID
from miles.utils.external_utils.command_utils.common import ArgvManipulator
from miles.utils.workers.serving.utils import parse_serve_worker_config
from miles.utils.workers.worker_provider.kubernetes.helm.naming import component_name
SERVE_CONFIG_FLAG: str = "--config"
# ============================ what a take-over rolls ==========================
@@ -60,6 +68,75 @@ def assert_only_orchestration_restarted(evidence: HotRestartEvidence, *, num_res
), f"a hot restart stamps exactly the two pod templates it replaces, and these carry a stamp too: {unexpected}"
# ======================== what a take-over's pods carry =======================
def assert_take_overs_carried_rollout_only_args(
evidence: HotRestartEvidence, *, flag: str, values: Sequence[str]
) -> None:
rollout_executor = component_name(evidence.release, ROLLOUT_EXECUTOR_POOL_ID)
uids_of_workload = _compute_pod_uids_of_workload(evidence.snapshots)
carried = [
_read_flag_of_serve_config(command, flag=flag)
for command in _commands_of(evidence, workload=rollout_executor, uids_of_workload=uids_of_workload)
]
assert len(carried) == len(values), (
f"{rollout_executor} ran as {len(carried)} pod(s) while {len(values)} generation(s) of arguments were "
f"installed, so the pods and the arguments cannot be paired up"
)
assert carried == list(values), (
f"the successive pods of {rollout_executor} carried {flag} as {carried}, and the launches installed "
f"{list(values)}: a take-over ran the component with arguments other than the ones it was relaunched with"
)
leaked = {
uid: workload
for workload, uids in uids_of_workload.items()
if workload != rollout_executor
for uid in uids
if any(value in part for part in evidence.commands_of_pod_uid[uid] for value in values)
}
assert not leaked, (
f"{flag} is read by the rollout executor alone, and the pods {leaked} of other workloads carry one of its "
f"values: the argument leaked into a payload a hot restart must leave untouched"
)
print(f"every take-over's rollout executor carried {flag} as relaunched, and no other pod did: {list(values)}")
def _compute_pod_uids_of_workload(snapshots: Sequence[ClusterSnapshot]) -> dict[str, list[str]]:
uids_of_workload: dict[str, list[str]] = {}
for snapshot in snapshots:
for pod in snapshot.pods:
if (workload := _compute_workload_of_pod(pod.name, workloads=snapshot.workload_names)) is None:
continue
uids = uids_of_workload.setdefault(workload, [])
if pod.uid not in uids:
uids.append(pod.uid)
return uids_of_workload
def _commands_of(
evidence: HotRestartEvidence, *, workload: str, uids_of_workload: dict[str, list[str]]
) -> list[list[str]]:
commands = []
for uid in uids_of_workload.get(workload, []):
assert (command := evidence.commands_of_pod_uid.get(uid)) is not None, (
f"pod {uid} of {workload} was observed, but no read of the release recorded its command, so what it "
f"ran is unknown"
)
commands.append(list(command))
return commands
def _read_flag_of_serve_config(command: list[str], *, flag: str) -> str | None:
assert (
rendered := ArgvManipulator.get_effective(command, SERVE_CONFIG_FLAG)
) is not None, f"a served worker is started with {SERVE_CONFIG_FLAG}, and this command carries none: {command}"
return parse_serve_worker_config(rendered).args.get(flag.removeprefix("--").replace("-", "_"))
# ========================== what the snapshots say ============================