mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
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:
@@ -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 ============================
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user