mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
Gate Ray fault injection between weight updates (#3066)
This commit is contained in:
@@ -29,7 +29,6 @@ from miles.utils.init_once import InitOnce, init_once
|
||||
from miles.utils.logging_utils import configure_logger
|
||||
from miles.utils.misc import SimpleTicker
|
||||
from miles.utils.test_utils.fault_injector import FailureMode
|
||||
from miles.utils.workers.ray_worker_manager import RayWorkerManager
|
||||
from miles.utils.workers.registration.hub import RegistrationHub
|
||||
from miles.utils.workers.registration.models import RegistrationSnapshot
|
||||
from miles.utils.workers.worker_provider.base import BaseWorkerProvider, CellInfo, StopWatchFn
|
||||
@@ -94,6 +93,22 @@ class InferenceController:
|
||||
async def stop_cell_between_weight_updates(self, cell_id: str) -> None:
|
||||
await self._engine_provider.stop_cells(cell_ids=[cell_id])
|
||||
|
||||
# TEMPORARY: exists only so fault injection can take this lock, reverted with the weight-update fault tolerance work
|
||||
@with_lock
|
||||
async def inject_fault_between_weight_updates(self, cell_id: str, *, mode: FailureMode, sub_index: int) -> None:
|
||||
# TEMPORARY: colocate cannot kill rollout workers while trainer ranks own the shared GPUs
|
||||
server = next((srv for srv in self.servers.values() if cell_id in srv.server_cells), None)
|
||||
if server is None:
|
||||
raise KeyError(f"Unknown rollout cell {cell_id!r}")
|
||||
if not server.health_checker_activeness.get().active:
|
||||
raise RuntimeError(f"Rollout cell {cell_id!r} is offloaded; refusing fault injection")
|
||||
|
||||
await self._engine_provider._worker_manager_handle.inject_fault.remote(
|
||||
cell_id,
|
||||
mode=mode.value,
|
||||
worker_in_cell_index=sub_index,
|
||||
)
|
||||
|
||||
# -------------------------- take over -----------------------------
|
||||
|
||||
@lock_exempt
|
||||
@@ -285,18 +300,6 @@ class InferenceController:
|
||||
f"Pass model_id to update exactly one of them."
|
||||
)
|
||||
|
||||
# -------------------------- cell operations -----------------------------
|
||||
|
||||
# TEMPORARY: exists only so fault injection can take this lock, reverted with the weight-update fault tolerance work
|
||||
@with_lock
|
||||
async def inject_fault_between_weight_updates(self, cell_id: str, *, mode: FailureMode, sub_index: int) -> None:
|
||||
# TEMPORARY: colocate cannot kill rollout workers while trainer ranks own the shared GPUs
|
||||
if not self._health_checker_activeness.get().active:
|
||||
raise RuntimeError(f"Rollout cell {cell_id!r} is offloaded; refusing fault injection")
|
||||
await RayWorkerManager.get_handle().inject_fault.remote(
|
||||
cell_id, mode=mode.value, worker_in_cell_index=sub_index
|
||||
)
|
||||
|
||||
# -------------------------- eval fleet -----------------------------
|
||||
|
||||
@lock_exempt
|
||||
|
||||
@@ -27,14 +27,36 @@ class RayCellOperations(BaseCellOperations):
|
||||
return await self._worker_manager_handle.get_cell_infos.remote(pool_ids=pool_ids)
|
||||
|
||||
async def suspend(self, *, cell_id: str) -> None:
|
||||
if _is_trainer_cell_id(cell_id):
|
||||
await self._worker_manager_handle.stop_cells.remote([cell_id])
|
||||
return
|
||||
|
||||
# TEMPORARY: taking the lock the weight update holds, reverted with that fault tolerance work
|
||||
# await self._worker_manager_handle.stop_cells.remote([cell_id]) # use this later
|
||||
if self._inference_controller is None:
|
||||
self._inference_controller = self._resolve_inference_controller()
|
||||
await self._inference_controller.stop_cell_between_weight_updates(cell_id=cell_id)
|
||||
await self._controller().stop_cell_between_weight_updates(cell_id=cell_id)
|
||||
|
||||
async def resume(self, *, cell_id: str) -> None:
|
||||
await self._worker_manager_handle.start_cells.remote([cell_id])
|
||||
|
||||
async def inject_fault(self, *, cell_id: str, mode: FailureMode, sub_index: int) -> None:
|
||||
await self._worker_manager_handle.inject_fault.remote(cell_id, mode=mode.value, worker_in_cell_index=sub_index)
|
||||
if _is_trainer_cell_id(cell_id):
|
||||
await self._worker_manager_handle.inject_fault.remote(
|
||||
cell_id, mode=mode.value, worker_in_cell_index=sub_index
|
||||
)
|
||||
return
|
||||
|
||||
# TEMPORARY: taking the lock the weight update holds, reverted with that fault tolerance work
|
||||
await self._controller().inject_fault_between_weight_updates(
|
||||
cell_id=cell_id,
|
||||
mode=mode,
|
||||
sub_index=sub_index,
|
||||
)
|
||||
|
||||
def _controller(self) -> BaseWorkerHandle:
|
||||
if self._inference_controller is None:
|
||||
self._inference_controller = self._resolve_inference_controller()
|
||||
return self._inference_controller
|
||||
|
||||
|
||||
# TEMPORARY: matched by pool prefix only, until the fault tolerance rework routes trainer cells properly
|
||||
def _is_trainer_cell_id(cell_id: str) -> bool:
|
||||
return cell_id.startswith("trainer-engine-")
|
||||
|
||||
+12
-5
@@ -307,7 +307,7 @@ Healing witness: one heal per target phase, at P+2 (healed = last cell, ckpt src
|
||||
```
|
||||
Type: comparison; both sides run the identical command, only the target is wrapped in the
|
||||
fault injector, through the pipeline's target_side_context hook
|
||||
Entry: test_rollout_deterministic__kill_rollout__dp2_cp2__colocate.py, ft-long
|
||||
Entry: test_rollout_deterministic__kill_rollout__dp4__colocate.py, ft-long
|
||||
Steps: 8 rollouts (NUM_ROLLOUTS)
|
||||
Requires: mode.has_real_rollout, and ft_components == ("rollout",) exactly
|
||||
Compare: dumps rel <= 0 (bitwise); metrics rtol=0 / atol=0 over train/* and rollout/*,
|
||||
@@ -318,12 +318,14 @@ Regime (both sides):
|
||||
--sglang-attention-backend flashinfer and --deterministic-mode
|
||||
- --debug-deterministic-collective and scenario_trainer_deterministic's deterministic env vars
|
||||
- --sglang-disable-radix-cache
|
||||
- --rollout-health-check-interval 5
|
||||
- --rollout-health-check-interval 1
|
||||
|
||||
Injection (target side only):
|
||||
1. Rollout cells, seed 42, exponential mean CRASH_INTERVAL_SECONDS (120s)
|
||||
1. Rollout cells, seed 42, exponential mean CRASH_INTERVAL_SECONDS (30s)
|
||||
2. Forms drawn per (cluster backend, cell type), as in the soaks
|
||||
3. Stop the injector with a 5s timeout, then re-use the soak's rollout witnesses: >= 2
|
||||
3. Stop accepting faults after six completed rollouts, leaving the final two rollouts for recovery
|
||||
4. Stop the injector, waiting out a mid-flight injection for at most
|
||||
STOP_AND_JOIN_TIMEOUT_SECONDS (180s), then re-use the soak's rollout witnesses: >= 2
|
||||
accepted rollout injections, each paired with one completed recovery cycle
|
||||
|
||||
Assertions:
|
||||
@@ -337,7 +339,12 @@ Assertions:
|
||||
- **Why it exists**: an engine dying and being replaced mid-generation is supposed to be invisible to training, and "invisible" is a claim about bits; the rollout soak only ever asserted survival.
|
||||
- **Why the shared deterministic recipe**: the assertion is deterministic replay across fresh inference engines, not true-on-policy training. Reusing the same FlashInfer recipe as the main deterministic trainer-FT test avoids a second, incompatible attention-backend contract.
|
||||
- **Why `--sglang-disable-radix-cache`**: a replacement engine serves with a cold prefix cache where the baseline's was warm, and deterministic inference is nowhere documented as prefix-cache-length invariant.
|
||||
- **Why `--rollout-health-check-interval 5`**: the generation retry loop gives up after ~60s while the default health check needs 90-120s to evict a dead worker, so a request could exhaust its retries against a corpse.
|
||||
- **Why this recipe disables batch-variant MM fallback**: a rollout worker loss changes co-batching while the pool is healing; permitting an `einsum` fallback would make the same seeded request depend on that temporary batch shape. The scenario injects the environment override without changing the production default.
|
||||
- **Why `--rollout-health-check-interval 1`**: healthy generation can finish between two five-second polls; the short scenario needs at least one fresh Serving observation before its lock-protected injection attempt.
|
||||
- **Why this scenario polls the fault window every 0.2 seconds**: colocated generation windows are only a few seconds long, so the generic two-second scheduler cadence can miss every Serving observation in an eight-rollout run.
|
||||
- **Why one quiescent poll on Ray only**: colocate exposes Serving only between train phases, and on Ray the injection endpoint takes the inference-controller lock and atomically rejects inactive servers, so the generic stable-serving gate is redundant. The Kubernetes forms (`exec_sigkill`, `delete_pod`) act on the pod without that lock, so a stale Serving poll or a slow `kubectl` would land the fault after the colocated engines are already offloaded; Kubernetes therefore keeps the generic 60-poll gate.
|
||||
- **Why the final two rollouts accept no new fault**: the scheduler keeps observing recovery but closes admission after rollout 5, so teardown cannot race a newly accepted replacement.
|
||||
- **Why Ray checks Serving again inside the injection lock**: the lock excludes weight-update and offload transitions, while the Serving check also rejects the subsequent colocated trainer phase after `offload()` has released the lock.
|
||||
- **Why every namespace, not just `train/`**: an engine crash shows up first in `rollout/raw_reward` or `rollout/log_probs`. `perf/` is left out by name, being wall-clock and throughput that a relaunch moves by definition, and a metric in neither namespace fails the run rather than being dropped quietly.
|
||||
- **Why the weights-moved gate**: bitwise equality is also satisfied by two runs that trained on nothing.
|
||||
- **Why not a loss or reward curve**: neither is a progress signal here — the reward is `deterministic_random`, a hash of the response, and GRPO's surrogate loss is not monotone even while a run learns. Over eight rollouts neither moves for a reason worth asserting, and the weights either changed or they did not.
|
||||
|
||||
@@ -133,29 +133,12 @@ class _RecordingServer:
|
||||
self.cells_timeouts.append(timeout)
|
||||
|
||||
|
||||
class _AbortRecordingServer:
|
||||
def __init__(
|
||||
self,
|
||||
name: str,
|
||||
*,
|
||||
started: list[str],
|
||||
completed: list[str],
|
||||
error: Exception | None = None,
|
||||
scheduling_turns: int = 1,
|
||||
) -> None:
|
||||
self.name = name
|
||||
self.started = started
|
||||
self.completed = completed
|
||||
self.error = error
|
||||
self.scheduling_turns = scheduling_turns
|
||||
class _RecordingRemoteMethod:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[tuple[tuple[Any, ...], dict[str, Any]]] = []
|
||||
|
||||
async def abort_all(self) -> None:
|
||||
self.started.append(self.name)
|
||||
for _ in range(self.scheduling_turns):
|
||||
await asyncio.sleep(0)
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
self.completed.append(self.name)
|
||||
async def remote(self, *args: Any, **kwargs: Any) -> None:
|
||||
self.calls.append((args, kwargs))
|
||||
|
||||
|
||||
class _FakeUpdatableCell:
|
||||
@@ -191,6 +174,7 @@ class _FakeWorkerProvider(BaseWorkerProvider):
|
||||
def __init__(self, cell_infos: list[CellInfo], *, pool_ids: list[str] | None = None) -> None:
|
||||
self._cell_infos = cell_infos
|
||||
self._pools = pool_ids or []
|
||||
self._worker_manager_handle = SimpleNamespace(inject_fault=_RecordingRemoteMethod())
|
||||
self.watched_pool_ids: list[str] | None = None
|
||||
self.initialized = False
|
||||
self.stop_watch_calls = 0
|
||||
@@ -252,37 +236,6 @@ def _make_controller(
|
||||
return controller
|
||||
|
||||
|
||||
class TestAbortAll:
|
||||
async def test_abort_all_reaches_every_server_before_propagating_the_first_failure(self) -> None:
|
||||
"""Every server is aborted to completion before the first fleet failure is propagated."""
|
||||
started: list[str] = []
|
||||
completed: list[str] = []
|
||||
first_failure = RuntimeError("first server refused abort")
|
||||
controller = _make_controller(
|
||||
{
|
||||
"first": _AbortRecordingServer(
|
||||
"first",
|
||||
started=started,
|
||||
completed=completed,
|
||||
error=first_failure,
|
||||
),
|
||||
"second": _AbortRecordingServer(
|
||||
"second",
|
||||
started=started,
|
||||
completed=completed,
|
||||
scheduling_turns=2,
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError) as exc_info:
|
||||
await controller.abort_all()
|
||||
|
||||
assert exc_info.value is first_failure
|
||||
assert set(started) == {"first", "second"}
|
||||
assert completed == ["second"]
|
||||
|
||||
|
||||
class TestHealthCheckerActiveness:
|
||||
@pytest.mark.asyncio
|
||||
async def test_offload_pauses_probing_before_putting_engines_to_sleep(self):
|
||||
@@ -346,6 +299,63 @@ class TestHealthCheckerActiveness:
|
||||
assert srv.health_checker_activeness.get().active
|
||||
|
||||
|
||||
class TestRolloutFaultInjectionWindow:
|
||||
async def test_fault_injection_refuses_an_unknown_rollout_cell(self) -> None:
|
||||
"""An unknown rollout cell raises its identifier without reaching the worker manager."""
|
||||
cell_id = "unknown-rollout-cell"
|
||||
provider = _FakeWorkerProvider([])
|
||||
controller = _make_controller(
|
||||
{"default": _RecordingServer(server_cells={"inference-engine-0-0-0": object()})},
|
||||
engine_provider=provider,
|
||||
)
|
||||
|
||||
with pytest.raises(KeyError) as exc_info:
|
||||
await controller.inject_fault_between_weight_updates(
|
||||
cell_id=cell_id,
|
||||
mode=FailureMode.SIGKILL,
|
||||
sub_index=0,
|
||||
)
|
||||
|
||||
assert cell_id in str(exc_info.value)
|
||||
assert provider._worker_manager_handle.inject_fault.calls == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fault_injection_reaches_a_serving_rollout_cell(self) -> None:
|
||||
"""A serving rollout cell accepts a fault through the Ray worker manager."""
|
||||
cell_id = "inference-engine-0-0-0"
|
||||
server = _RecordingServer(server_cells={cell_id: object()})
|
||||
provider = _FakeWorkerProvider([])
|
||||
controller = _make_controller({"default": server}, engine_provider=provider)
|
||||
|
||||
await controller.inject_fault_between_weight_updates(
|
||||
cell_id=cell_id,
|
||||
mode=FailureMode.SIGKILL,
|
||||
sub_index=0,
|
||||
)
|
||||
|
||||
assert provider._worker_manager_handle.inject_fault.calls == [
|
||||
((cell_id,), {"mode": "sigkill", "worker_in_cell_index": 0})
|
||||
]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fault_injection_refuses_an_offloaded_rollout_cell(self) -> None:
|
||||
"""Colocate must not kill rollout processes while trainer ranks own the shared GPUs."""
|
||||
cell_id = "inference-engine-0-0-0"
|
||||
server = _RecordingServer(server_cells={cell_id: object()})
|
||||
server.health_checker_activeness.bump_active(False)
|
||||
provider = _FakeWorkerProvider([])
|
||||
controller = _make_controller({"default": server}, engine_provider=provider)
|
||||
|
||||
with pytest.raises(RuntimeError, match="is offloaded; refusing fault injection"):
|
||||
await controller.inject_fault_between_weight_updates(
|
||||
cell_id=cell_id,
|
||||
mode=FailureMode.SIGKILL,
|
||||
sub_index=0,
|
||||
)
|
||||
|
||||
assert provider._worker_manager_handle.inject_fault.calls == []
|
||||
|
||||
|
||||
class TestReconcile:
|
||||
@pytest.fixture
|
||||
def servers(self) -> dict[str, _RecordingServer]:
|
||||
@@ -642,9 +652,12 @@ class TestInitSubscription:
|
||||
assert "session-server" not in provider.watched_pool_ids
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_init_survives_a_router_cell_offered_by_the_provider(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""A router cell carries no engine meta, so a too-wide subscription kills startup in the initial sync."""
|
||||
async def test_init_subscribes_narrowly_enough_to_never_see_a_router_cell(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""A router cell carries no engine meta, so reading one as engine meta would kill startup; the
|
||||
controller is safe because it subscribes to the engine pools alone."""
|
||||
args = make_args()
|
||||
assert compute_router_pool_id(0) not in compute_engine_pool_ids(args)
|
||||
|
||||
router_info = CellInfo(
|
||||
cell_id="inference-router-0-0",
|
||||
pool_id=compute_router_pool_id(0),
|
||||
@@ -1158,194 +1171,124 @@ async def _raise_async(cell: ServerCell) -> None:
|
||||
raise RuntimeError("injected init failure")
|
||||
|
||||
|
||||
class _RecordingWorkerManager:
|
||||
def __init__(self) -> None:
|
||||
_ROLLOUT_CELL_ID = "inference-engine-0-0-0"
|
||||
|
||||
|
||||
class _StoppingWorkerProvider(_FakeWorkerProvider):
|
||||
def __init__(self, *, completion: asyncio.Future | None = None) -> None:
|
||||
super().__init__([])
|
||||
self.stopped_cells: list[list[str]] = []
|
||||
self.injected: list[tuple[str, dict[str, Any]]] = []
|
||||
self.stop_requested = asyncio.Event()
|
||||
self._completion = completion
|
||||
|
||||
@property
|
||||
def stop_cells(self) -> Any:
|
||||
return _RecordingRemoteCall(lambda cell_ids: self.stopped_cells.append(list(cell_ids)))
|
||||
|
||||
@property
|
||||
def inject_fault(self) -> Any:
|
||||
return _RecordingRemoteCall(lambda cell_id, **kwargs: self.injected.append((cell_id, kwargs)))
|
||||
async def stop_cells(self, *, cell_ids: list[str]) -> None:
|
||||
self.stopped_cells.append(list(cell_ids))
|
||||
self.stop_requested.set()
|
||||
if self._completion is not None:
|
||||
await self._completion
|
||||
|
||||
|
||||
class _RecordingRemoteCall:
|
||||
def __init__(self, record: Any) -> None:
|
||||
self._record = record
|
||||
|
||||
def remote(self, *args: Any, **kwargs: Any) -> asyncio.Future:
|
||||
self._record(*args, **kwargs)
|
||||
future: asyncio.Future = asyncio.get_event_loop().create_future()
|
||||
future.set_result(None)
|
||||
return future
|
||||
def _make_cell_operations_controller(
|
||||
provider: _StoppingWorkerProvider, *, probing_paused: bool = False
|
||||
) -> InferenceController:
|
||||
server = _RecordingServer(server_cells={_ROLLOUT_CELL_ID: object()})
|
||||
if probing_paused:
|
||||
server.health_checker_activeness.bump_active(False)
|
||||
return _make_controller({"default": server}, engine_provider=provider)
|
||||
|
||||
|
||||
def _patch_worker_manager(monkeypatch: pytest.MonkeyPatch) -> _RecordingWorkerManager:
|
||||
manager = _RecordingWorkerManager()
|
||||
monkeypatch.setattr(inference_controller_module, "RayWorkerManager", SimpleNamespace(get_handle=lambda: manager))
|
||||
return manager
|
||||
async def _hold_context_lock(controller: InferenceController) -> tuple[asyncio.Task, asyncio.Event]:
|
||||
entered, may_finish = asyncio.Event(), asyncio.Event()
|
||||
|
||||
async def _hold() -> None:
|
||||
async with controller.context_lock:
|
||||
entered.set()
|
||||
await may_finish.wait()
|
||||
|
||||
class _BlockingWorkerManager:
|
||||
def __init__(self, *, completion: asyncio.Future) -> None:
|
||||
self.requested = asyncio.Event()
|
||||
self.stopped_cells: list[list[str]] = []
|
||||
self.completion = completion
|
||||
|
||||
@property
|
||||
def stop_cells(self) -> Any:
|
||||
return _BlockingRemoteCall(manager=self)
|
||||
|
||||
|
||||
class _BlockingRemoteCall:
|
||||
def __init__(self, *, manager: _BlockingWorkerManager) -> None:
|
||||
self._manager = manager
|
||||
|
||||
def remote(self, cell_ids: list[str]) -> asyncio.Future:
|
||||
self._manager.stopped_cells.append(list(cell_ids))
|
||||
self._manager.requested.set()
|
||||
return self._manager.completion
|
||||
|
||||
|
||||
def _patch_blocking_worker_manager(
|
||||
monkeypatch: pytest.MonkeyPatch, *, completion: asyncio.Future
|
||||
) -> _BlockingWorkerManager:
|
||||
manager = _BlockingWorkerManager(completion=completion)
|
||||
monkeypatch.setattr(inference_controller_module, "RayWorkerManager", SimpleNamespace(get_handle=lambda: manager))
|
||||
return manager
|
||||
holder = asyncio.create_task(_hold())
|
||||
await entered.wait()
|
||||
return holder, may_finish
|
||||
|
||||
|
||||
class TestCellOperations:
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_cell_between_weight_updates_is_forwarded_to_the_worker_manager(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""The manager owns the processes, so the controller only serializes the suspension."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
async def test_stop_cell_between_weight_updates_is_forwarded_to_the_engine_provider(self):
|
||||
"""The provider owns the processes, so the controller only serializes the suspension."""
|
||||
provider = _StoppingWorkerProvider()
|
||||
controller = _make_cell_operations_controller(provider)
|
||||
|
||||
await controller.stop_cell_between_weight_updates("engine-0")
|
||||
await controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID)
|
||||
|
||||
assert manager.stopped_cells == [["engine-0"]]
|
||||
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inject_fault_between_weight_updates_is_forwarded_to_the_worker_manager(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""Injection targets one worker of one cell, and only the manager can reach it."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
|
||||
await controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=1)
|
||||
|
||||
assert manager.injected == [("engine-0", {"mode": "sigkill", "worker_in_cell_index": 1})]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inject_fault_between_weight_updates_is_refused_while_probing_is_paused(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""An offloaded or updating cell would report a crash that the trainer cannot distinguish from its own pause."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
controller._health_checker_activeness.bump_active(False)
|
||||
|
||||
with pytest.raises(RuntimeError, match="refusing fault injection"):
|
||||
await controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
|
||||
|
||||
assert manager.injected == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_cell_between_weight_updates_waits_until_the_weight_update_window_closes(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
async def test_stop_cell_between_weight_updates_waits_until_the_weight_update_window_closes(self):
|
||||
"""Suspending a cell mid-broadcast leaves the trainer waiting on an engine that is being torn down."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
entered = asyncio.Event()
|
||||
may_finish = asyncio.Event()
|
||||
provider = _StoppingWorkerProvider()
|
||||
controller = _make_cell_operations_controller(provider)
|
||||
holder, may_finish = await _hold_context_lock(controller)
|
||||
|
||||
async def _hold_lock() -> None:
|
||||
async with controller.context_lock:
|
||||
entered.set()
|
||||
await may_finish.wait()
|
||||
|
||||
holder = asyncio.create_task(_hold_lock())
|
||||
await entered.wait()
|
||||
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates("engine-0"))
|
||||
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID))
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert manager.stopped_cells == []
|
||||
assert provider.stopped_cells == []
|
||||
|
||||
may_finish.set()
|
||||
await holder
|
||||
await stopping
|
||||
|
||||
assert manager.stopped_cells == [["engine-0"]]
|
||||
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_inject_fault_between_weight_updates_waits_until_the_weight_update_window_closes(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
async def test_inject_fault_between_weight_updates_waits_until_the_weight_update_window_closes(self):
|
||||
"""Injection racing a broadcast is the same hazard as suspension, so it takes the same turn."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
entered = asyncio.Event()
|
||||
may_finish = asyncio.Event()
|
||||
provider = _StoppingWorkerProvider()
|
||||
controller = _make_cell_operations_controller(provider)
|
||||
holder, may_finish = await _hold_context_lock(controller)
|
||||
|
||||
async def _hold_lock() -> None:
|
||||
async with controller.context_lock:
|
||||
entered.set()
|
||||
await may_finish.wait()
|
||||
|
||||
holder = asyncio.create_task(_hold_lock())
|
||||
await entered.wait()
|
||||
injecting = asyncio.create_task(
|
||||
controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
|
||||
controller.inject_fault_between_weight_updates(_ROLLOUT_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0)
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert manager.injected == []
|
||||
assert provider._worker_manager_handle.inject_fault.calls == []
|
||||
|
||||
may_finish.set()
|
||||
await holder
|
||||
await injecting
|
||||
|
||||
assert manager.injected == [("engine-0", {"mode": "sigkill", "worker_in_cell_index": 0})]
|
||||
assert provider._worker_manager_handle.inject_fault.calls == [
|
||||
((_ROLLOUT_CELL_ID,), {"mode": "sigkill", "worker_in_cell_index": 0})
|
||||
]
|
||||
|
||||
async def test_stop_cell_between_weight_updates_is_allowed_while_probing_is_paused(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_cell_between_weight_updates_is_allowed_while_probing_is_paused(self):
|
||||
"""An offloaded cell is the one a heal loop most needs to suspend, so only injection is refused."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
controller._health_checker_activeness.bump_active(False)
|
||||
provider = _StoppingWorkerProvider()
|
||||
controller = _make_cell_operations_controller(provider, probing_paused=True)
|
||||
|
||||
await controller.stop_cell_between_weight_updates("engine-0")
|
||||
await controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID)
|
||||
|
||||
assert manager.stopped_cells == [["engine-0"]]
|
||||
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
|
||||
|
||||
async def test_inject_fault_between_weight_updates_refuses_a_pause_that_began_while_it_waited(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
@pytest.mark.asyncio
|
||||
async def test_inject_fault_between_weight_updates_refuses_a_pause_that_began_while_it_waited(self):
|
||||
"""Reading the pause before taking the lock would kill a cell the offload has since put to sleep."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
entered = asyncio.Event()
|
||||
may_finish = asyncio.Event()
|
||||
provider = _StoppingWorkerProvider()
|
||||
controller = _make_cell_operations_controller(provider)
|
||||
server = controller.servers["default"]
|
||||
entered, may_finish = asyncio.Event(), asyncio.Event()
|
||||
|
||||
async def _pause_probing_under_the_lock() -> None:
|
||||
async with controller.context_lock:
|
||||
entered.set()
|
||||
await may_finish.wait()
|
||||
controller._health_checker_activeness.bump_active(False)
|
||||
server.health_checker_activeness.bump_active(False)
|
||||
|
||||
holder = asyncio.create_task(_pause_probing_under_the_lock())
|
||||
await entered.wait()
|
||||
injecting = asyncio.create_task(
|
||||
controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
|
||||
controller.inject_fault_between_weight_updates(_ROLLOUT_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0)
|
||||
)
|
||||
await asyncio.sleep(0)
|
||||
may_finish.set()
|
||||
@@ -1354,38 +1297,39 @@ class TestCellOperations:
|
||||
with pytest.raises(RuntimeError, match="refusing fault injection"):
|
||||
await injecting
|
||||
|
||||
assert manager.injected == []
|
||||
assert provider._worker_manager_handle.inject_fault.calls == []
|
||||
|
||||
async def test_a_refused_injection_leaves_the_weight_update_lock_free(self, monkeypatch: pytest.MonkeyPatch):
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_refused_injection_leaves_the_weight_update_lock_free(self):
|
||||
"""A refusal that kept the lock would hang the next weight update instead of only skipping the injection."""
|
||||
manager = _patch_worker_manager(monkeypatch)
|
||||
controller = _make_controller({})
|
||||
controller._health_checker_activeness.bump_active(False)
|
||||
provider = _StoppingWorkerProvider()
|
||||
controller = _make_cell_operations_controller(provider, probing_paused=True)
|
||||
|
||||
with pytest.raises(RuntimeError, match="refusing fault injection"):
|
||||
await controller.inject_fault_between_weight_updates("engine-0", mode=FailureMode.SIGKILL, sub_index=0)
|
||||
await controller.inject_fault_between_weight_updates(
|
||||
_ROLLOUT_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0
|
||||
)
|
||||
|
||||
assert not controller.context_lock.locked
|
||||
|
||||
await controller.stop_cell_between_weight_updates("engine-0")
|
||||
await controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID)
|
||||
|
||||
assert manager.stopped_cells == [["engine-0"]]
|
||||
assert provider.stopped_cells == [[_ROLLOUT_CELL_ID]]
|
||||
|
||||
async def test_a_weight_update_cannot_start_while_a_suspension_is_still_running(
|
||||
self, monkeypatch: pytest.MonkeyPatch
|
||||
):
|
||||
"""Releasing the lock before the manager has torn the cell down reopens the very race this serializes."""
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_weight_update_cannot_start_while_a_suspension_is_still_running(self):
|
||||
"""Releasing the lock before the provider has torn the cell down reopens the very race this serializes."""
|
||||
completion: asyncio.Future = asyncio.get_running_loop().create_future()
|
||||
manager = _patch_blocking_worker_manager(monkeypatch, completion=completion)
|
||||
controller = _make_controller({})
|
||||
provider = _StoppingWorkerProvider(completion=completion)
|
||||
controller = _make_cell_operations_controller(provider)
|
||||
weight_update_started = asyncio.Event()
|
||||
|
||||
async def _start_weight_update() -> None:
|
||||
async with controller.context_lock:
|
||||
weight_update_started.set()
|
||||
|
||||
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates("engine-0"))
|
||||
await manager.requested.wait()
|
||||
stopping = asyncio.create_task(controller.stop_cell_between_weight_updates(_ROLLOUT_CELL_ID))
|
||||
await provider.stop_requested.wait()
|
||||
weight_update = asyncio.create_task(_start_weight_update())
|
||||
await asyncio.sleep(0)
|
||||
|
||||
|
||||
@@ -5,14 +5,20 @@ from dataclasses import dataclass
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from miles.ray.rollout.inference_controller import InferenceController
|
||||
from miles.utils.context_lock import ContextLock
|
||||
from miles.utils.ft_utils.health_checker import ActivenessTracker
|
||||
from miles.utils.test_utils.fault_injector import FailureMode
|
||||
from miles.utils.workers.cell_operations.ray import RayCellOperations
|
||||
|
||||
_TRAINER_CELL_ID = "trainer-engine-actor-00001"
|
||||
|
||||
|
||||
class _RecordingEngineProvider:
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, *, worker_manager: _RecordingWorkerManagerHandle) -> None:
|
||||
self._worker_manager_handle = worker_manager
|
||||
self.stopped: list[str] = []
|
||||
|
||||
async def stop_cells(self, *, cell_ids: list[str]) -> None:
|
||||
@@ -35,6 +41,7 @@ class _RecordingWorkerManagerHandle:
|
||||
self.calls: list[tuple[str, tuple[Any, ...], dict[str, Any]]] = []
|
||||
self.get_cell_infos = _RecordingRemoteMethod(name="get_cell_infos", calls=self.calls)
|
||||
self.start_cells = _RecordingRemoteMethod(name="start_cells", calls=self.calls)
|
||||
self.stop_cells = _RecordingRemoteMethod(name="stop_cells", calls=self.calls)
|
||||
self.inject_fault = _RecordingRemoteMethod(name="inject_fault", calls=self.calls)
|
||||
|
||||
|
||||
@@ -47,9 +54,15 @@ class _Fixture:
|
||||
|
||||
|
||||
def _make_fixture() -> _Fixture:
|
||||
provider = _RecordingEngineProvider()
|
||||
controller = InferenceController(SimpleNamespace(), engine_provider=provider, router_providers=[])
|
||||
worker_manager = _RecordingWorkerManagerHandle()
|
||||
provider = _RecordingEngineProvider(worker_manager=worker_manager)
|
||||
controller = InferenceController(SimpleNamespace(), engine_provider=provider, router_providers=[])
|
||||
controller.servers = {
|
||||
"actor": SimpleNamespace(
|
||||
server_cells={"engine-0-2": SimpleNamespace()},
|
||||
health_checker_activeness=ActivenessTracker(active=True),
|
||||
)
|
||||
}
|
||||
return _Fixture(
|
||||
provider=provider,
|
||||
controller=controller,
|
||||
@@ -71,30 +84,6 @@ async def _settle() -> None:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
|
||||
class TestRayCellOperationsProtocol:
|
||||
async def test_cell_infos_forwards_pool_ids_and_returns_the_actor_result(self) -> None:
|
||||
"""Cell info reads forward every pool ID by keyword and preserve the actor result."""
|
||||
worker_manager = _RecordingWorkerManagerHandle()
|
||||
operations = RayCellOperations(worker_manager_handle=worker_manager)
|
||||
pool_ids = ["engine-0", "rollout-1"]
|
||||
actor_result = {"engine-0-2": object(), "rollout-1-3": object()}
|
||||
worker_manager.get_cell_infos.result = actor_result
|
||||
|
||||
result = await operations.cell_infos(pool_ids=pool_ids)
|
||||
|
||||
assert worker_manager.calls == [("get_cell_infos", (), {"pool_ids": pool_ids})]
|
||||
assert result is actor_result
|
||||
|
||||
async def test_resume_starts_exactly_the_requested_cell(self) -> None:
|
||||
"""Resuming a cell sends exactly that cell ID in a one-element list."""
|
||||
worker_manager = _RecordingWorkerManagerHandle()
|
||||
operations = RayCellOperations(worker_manager_handle=worker_manager)
|
||||
|
||||
await operations.resume(cell_id="engine-0-2")
|
||||
|
||||
assert worker_manager.calls == [("start_cells", (["engine-0-2"],), {})]
|
||||
|
||||
|
||||
async def test_a_suspend_waits_for_the_controller_lock_instead_of_reaching_the_worker_manager() -> None:
|
||||
"""A suspend arriving mid weight update must not reach the worker manager until the update ends."""
|
||||
fixture = _make_fixture()
|
||||
@@ -115,8 +104,78 @@ async def test_a_suspend_waits_for_the_controller_lock_instead_of_reaching_the_w
|
||||
assert fixture.worker_manager.calls == []
|
||||
|
||||
|
||||
async def test_every_other_operation_goes_straight_through() -> None:
|
||||
"""Only a suspend takes a rank out of a live collective, so nothing else pays for the lock."""
|
||||
async def test_inject_fault_waits_for_the_controller_lock() -> None:
|
||||
"""A fault arriving mid weight update must wait before killing an engine worker."""
|
||||
fixture = _make_fixture()
|
||||
acquired, release = asyncio.Event(), asyncio.Event()
|
||||
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
|
||||
await acquired.wait()
|
||||
|
||||
injecting = asyncio.create_task(
|
||||
fixture.operations.inject_fault(cell_id="engine-0-2", mode=FailureMode.SIGKILL, sub_index=0)
|
||||
)
|
||||
await _settle()
|
||||
assert not injecting.done()
|
||||
assert fixture.worker_manager.calls == []
|
||||
|
||||
release.set()
|
||||
await holding
|
||||
await injecting
|
||||
assert fixture.worker_manager.calls == [
|
||||
("inject_fault", ("engine-0-2",), {"mode": "sigkill", "worker_in_cell_index": 0})
|
||||
]
|
||||
|
||||
|
||||
async def test_a_trainer_cells_fault_reaches_the_worker_manager() -> None:
|
||||
"""Regression: routing a trainer cell through the rollout controller raised, so the actor never died."""
|
||||
fixture = _make_fixture()
|
||||
acquired, release = asyncio.Event(), asyncio.Event()
|
||||
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
|
||||
await acquired.wait()
|
||||
|
||||
await asyncio.wait_for(
|
||||
fixture.operations.inject_fault(cell_id=_TRAINER_CELL_ID, mode=FailureMode.SIGKILL, sub_index=0),
|
||||
timeout=5.0,
|
||||
)
|
||||
|
||||
assert fixture.worker_manager.calls == [
|
||||
("inject_fault", (_TRAINER_CELL_ID,), {"mode": "sigkill", "worker_in_cell_index": 0})
|
||||
]
|
||||
|
||||
release.set()
|
||||
await holding
|
||||
|
||||
|
||||
async def test_a_trainer_cells_suspend_reaches_the_worker_manager() -> None:
|
||||
"""A trainer cell is none of the controller's business, and its lock would only stall the stop."""
|
||||
fixture = _make_fixture()
|
||||
acquired, release = asyncio.Event(), asyncio.Event()
|
||||
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
|
||||
await acquired.wait()
|
||||
|
||||
await asyncio.wait_for(fixture.operations.suspend(cell_id=_TRAINER_CELL_ID), timeout=5.0)
|
||||
|
||||
assert fixture.worker_manager.calls == [("stop_cells", ([_TRAINER_CELL_ID],), {})]
|
||||
assert fixture.provider.stopped == []
|
||||
|
||||
release.set()
|
||||
await holding
|
||||
|
||||
|
||||
async def test_a_rollout_cell_the_controller_does_not_list_yet_still_goes_through_the_controller() -> None:
|
||||
"""Routing on live membership would kill an engine being replaced without the weight-update lock."""
|
||||
fixture = _make_fixture()
|
||||
|
||||
with pytest.raises(KeyError):
|
||||
await asyncio.wait_for(
|
||||
fixture.operations.inject_fault(cell_id="engine-0-7", mode=FailureMode.SIGKILL, sub_index=0), timeout=5.0
|
||||
)
|
||||
|
||||
assert fixture.worker_manager.calls == []
|
||||
|
||||
|
||||
async def test_non_disruptive_operations_go_straight_through() -> None:
|
||||
"""Cell reads and resumes do not wait for the weight-update lock."""
|
||||
fixture = _make_fixture()
|
||||
acquired, release = asyncio.Event(), asyncio.Event()
|
||||
holding = asyncio.create_task(_hold_lock(lock=fixture.controller.context_lock, acquired=acquired, release=release))
|
||||
@@ -124,18 +183,36 @@ async def test_every_other_operation_goes_straight_through() -> None:
|
||||
|
||||
await asyncio.wait_for(fixture.operations.cell_infos(pool_ids=["engine-0"]), timeout=5.0)
|
||||
await asyncio.wait_for(fixture.operations.resume(cell_id="engine-0-2"), timeout=5.0)
|
||||
await asyncio.wait_for(
|
||||
fixture.operations.inject_fault(cell_id="engine-0-2", mode=FailureMode.SIGKILL, sub_index=0),
|
||||
timeout=5.0,
|
||||
)
|
||||
|
||||
assert [name for name, _, _ in fixture.worker_manager.calls] == ["get_cell_infos", "start_cells", "inject_fault"]
|
||||
assert [name for name, _, _ in fixture.worker_manager.calls] == ["get_cell_infos", "start_cells"]
|
||||
assert fixture.provider.stopped == []
|
||||
|
||||
release.set()
|
||||
await holding
|
||||
|
||||
|
||||
class TestRayCellOperationsProtocol:
|
||||
async def test_cell_infos_forwards_pool_ids_and_returns_the_actor_result(self) -> None:
|
||||
"""Cell info reads forward every pool ID by keyword and preserve the actor result."""
|
||||
fixture = _make_fixture()
|
||||
pool_ids = ["engine-0", "rollout-1"]
|
||||
actor_result = {"engine-0-2": SimpleNamespace(), "rollout-1-3": SimpleNamespace()}
|
||||
fixture.worker_manager.get_cell_infos.result = actor_result
|
||||
|
||||
result = await fixture.operations.cell_infos(pool_ids=pool_ids)
|
||||
|
||||
assert fixture.worker_manager.calls == [("get_cell_infos", (), {"pool_ids": pool_ids})]
|
||||
assert result is actor_result
|
||||
|
||||
async def test_resume_starts_exactly_the_requested_cell(self) -> None:
|
||||
"""Resuming a cell sends exactly that cell ID in a one-element list."""
|
||||
fixture = _make_fixture()
|
||||
|
||||
await fixture.operations.resume(cell_id="engine-0-2")
|
||||
|
||||
assert fixture.worker_manager.calls == [("start_cells", (["engine-0-2"],), {})]
|
||||
|
||||
|
||||
class TestRayCellOperationsInferenceControllerResolution:
|
||||
async def test_the_inference_controller_is_resolved_only_when_a_disruptive_operation_needs_it(self) -> None:
|
||||
"""Construction, reads, and resumes do not resolve the controller before a suspend needs it."""
|
||||
|
||||
Reference in New Issue
Block a user