diff --git a/miles/ray/rollout/rollout_server.py b/miles/ray/rollout/rollout_server.py index 73dbcfb27cf..004a93ea870 100644 --- a/miles/ray/rollout/rollout_server.py +++ b/miles/ray/rollout/rollout_server.py @@ -104,11 +104,11 @@ def start_rollout_servers(args, pg) -> dict[str, "RolloutServer"]: engine_offset += num_engines gpu_offset += group_cfg.num_gpus - new_engine_indices_per_group = async_utils.wait_futures(start_futures) + new_cell_indices_per_group = async_utils.wait_futures(start_futures) - for group, new_engine_indices in zip(server_groups, new_engine_indices_per_group, strict=True): - group.mark_alive(engine_indices=new_engine_indices) - async_utils.run(group.register_workers(new_engine_indices)) + for group, new_cell_indices in zip(server_groups, new_cell_indices_per_group, strict=True): + group.mark_alive(cell_indices=new_cell_indices) + async_utils.run(group.register_workers(new_cell_indices)) servers[model_cfg.name] = RolloutServer( server_groups=server_groups, diff --git a/miles/ray/rollout/server_cell.py b/miles/ray/rollout/server_cell.py index ce03170c8aa..cf565a8c2dd 100644 --- a/miles/ray/rollout/server_cell.py +++ b/miles/ray/rollout/server_cell.py @@ -10,6 +10,7 @@ from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy from miles.backends.sglang_utils.sglang_engine import SGLangEngine, build_server_url +from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient, use_legacy_router_api from miles.ray.rollout.addr_allocator import PortAllocator from miles.ray.rollout.server_engine import AddrInfo, ServerEngine from miles.ray.utils import NOSET_VISIBLE_DEVICES_ENV_VARS_LIST @@ -138,6 +139,20 @@ async def check_weights(self, action: str, allow_quant_error: bool, selector: st action=action, allow_quant_error=allow_quant_error, selector=selector, skip_list=skip_list ) + async def register(self, router_api_client: SGLangRouterApiClient) -> None: + await router_api_client.add_worker( + worker_url=self.primary_engine.addr_info.server_url, + worker_type=self.worker_type, + use_legacy_api=use_legacy_router_api(self.args), + bootstrap_port=self.primary_engine.addr_info.bootstrap_port, + ) + + async def unregister(self, router_api_client: SGLangRouterApiClient) -> None: + await router_api_client.remove_worker( + worker_url=self.primary_engine.addr_info.server_url, + use_legacy_api=use_legacy_router_api(self.args), + ) + def flatten_cells(cells: list[ServerCell]) -> list[ServerEngine]: return [engine for cell in cells for engine in cell.engines] diff --git a/miles/ray/rollout/server_group.py b/miles/ray/rollout/server_group.py index 3ae00153a41..3a8ec497a22 100644 --- a/miles/ray/rollout/server_group.py +++ b/miles/ray/rollout/server_group.py @@ -5,9 +5,9 @@ from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS -from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient, use_legacy_router_api +from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient from miles.ray.rollout.addr_allocator import PortAllocator -from miles.ray.rollout.server_cell import SHUTDOWN_TIMEOUT, ServerCell, flatten_cells +from miles.ray.rollout.server_cell import SHUTDOWN_TIMEOUT, ServerCell from miles.ray.rollout.server_engine import ServerEngine from miles.utils import async_utils @@ -55,9 +55,9 @@ async def start_engines( """Create Ray actors, allocate ports, and run ``engine.init()`` on every new engine. Mutates ``port_allocator`` in place to advance past any newly assigned ports. - Returns the list of indices into the group's flat engine list that were just - allocated. Actor creation, port allocation and state marking all happen before - the first await point, so concurrent callers cannot double-start a slot. + Returns the indices of the cells that were just allocated. Actor creation, + port allocation and state marking all happen before the first await point, + so concurrent callers cannot double-start a slot. """ if self.args.debug_train_only or self.worker_type == "placeholder": self.has_new_engines = False @@ -77,52 +77,25 @@ async def start_engines( await asyncio.gather(*cell_starts) - new_engine_indices = [ - cell_index * self.nodes_per_engine + local_index - for cell_index in started_cell_indices - for local_index in range(self.nodes_per_engine) - ] - self.has_new_engines |= bool(new_engine_indices) - return new_engine_indices + self.has_new_engines |= bool(started_cell_indices) + return started_cell_indices - async def register_workers(self, engine_indices: list[int]) -> None: + async def register_workers(self, cell_indices: list[int]) -> None: if self.args.rollout_external or not (self.router_ip and self.router_port): return await asyncio.gather( - *[ - self._router_api_client.add_worker( - worker_url=engine.addr_info.server_url, - worker_type=self.worker_type, - use_legacy_api=use_legacy_router_api(self.args), - bootstrap_port=engine.addr_info.bootstrap_port, - ) - for engine in self._primary_engines_of(engine_indices) - ] + *[cell.register(self._router_api_client) for cell in self._allocated_cells_of(cell_indices)] ) - async def unregister_workers(self, engine_indices: list[int]) -> None: + async def unregister_workers(self, cell_indices: list[int]) -> None: if self.args.rollout_external or not (self.router_ip and self.router_port): return await asyncio.gather( - *[ - self._router_api_client.remove_worker( - worker_url=engine.addr_info.server_url, - use_legacy_api=use_legacy_router_api(self.args), - ) - for engine in self._primary_engines_of(engine_indices) - ] + *[cell.unregister(self._router_api_client) for cell in self._allocated_cells_of(cell_indices)] ) - def _engine_indices_of_cell(self, cell_index: int) -> range: - return range(cell_index * self.nodes_per_engine, (cell_index + 1) * self.nodes_per_engine) - - def _primary_engines_of(self, engine_indices: list[int]) -> list[ServerEngine]: - all_engines = flatten_cells(self.cells) - return [ - all_engines[index] - for index in engine_indices - if index % self.nodes_per_engine == 0 and all_engines[index].is_allocated - ] + def _allocated_cells_of(self, cell_indices: list[int]) -> list[ServerCell]: + return [self.cells[cell_index] for cell_index in cell_indices if self.cells[cell_index].is_allocated] @property def _router_api_client(self) -> SGLangRouterApiClient: @@ -135,11 +108,10 @@ def _router_api_client(self) -> SGLangRouterApiClient: # moving `shutdown` mainly to local code def stop_engines(self, cell_indices: list[int]): logger.info(f"Killing server {cell_indices=}...") - engine_indices = [i for cell_index in cell_indices for i in self._engine_indices_of_cell(cell_index)] try: - async_utils.run(asyncio.wait_for(self.unregister_workers(engine_indices), timeout=SHUTDOWN_TIMEOUT)) + async_utils.run(asyncio.wait_for(self.unregister_workers(cell_indices), timeout=SHUTDOWN_TIMEOUT)) except Exception as e: - logger.warning(f"Unregistering {engine_indices=} from the router failed, tearing down anyway (e: {e})") + logger.warning(f"Unregistering {cell_indices=} from the router failed, tearing down anyway (e: {e})") for cell_index in sorted(set(cell_indices)): self.cells[cell_index].stop() @@ -147,14 +119,13 @@ async def recover(self, port_allocator: PortAllocator, filter_cell_indices: list if filter_cell_indices is None: filter_cell_indices = [cell_index for cell_index, cell in enumerate(self.cells) if not cell.is_allocated] - new_engine_indices = await self.start_engines(port_allocator, start_cell_indices=filter_cell_indices) + started_cell_indices = await self.start_engines(port_allocator, start_cell_indices=filter_cell_indices) - all_engines = flatten_cells(self.cells) release_handles = [] all_resume_engines = [] - logger.info(f"Recovered {len(new_engine_indices)} dead rollout engines (worker_type={self.worker_type})") - if self.needs_offload and new_engine_indices: - new_primary_engines = [all_engines[i] for i in new_engine_indices if i % self.nodes_per_engine == 0] + logger.info(f"Recovered {len(started_cell_indices)} dead rollout cells (worker_type={self.worker_type})") + if self.needs_offload and started_cell_indices: + new_primary_engines = [self.cells[i].primary_engine for i in started_cell_indices] release_handles.extend(engine.api_client.release_memory_occupation() for engine in new_primary_engines) if self.update_weights or self.model_path: all_resume_engines.extend(new_primary_engines) @@ -169,13 +140,13 @@ async def recover(self, port_allocator: PortAllocator, filter_cell_indices: list ] ) - self.mark_alive(engine_indices=new_engine_indices) - await self.register_workers(new_engine_indices) + self.mark_alive(cell_indices=started_cell_indices) + await self.register_workers(started_cell_indices) - def mark_alive(self, engine_indices: list[int]): - all_engines = flatten_cells(self.cells) - for engine_index in engine_indices: - all_engines[engine_index].mark_alive() + def mark_alive(self, cell_indices: list[int]): + for cell_index in cell_indices: + for engine in self.cells[cell_index].engines: + engine.mark_alive() async def offload(self, tags: list[str] | None = None): if not self.needs_offload: diff --git a/tests/fast/ray/rollout/test_server_cell.py b/tests/fast/ray/rollout/test_server_cell.py index cb0602f74a3..bcd7689f99b 100644 --- a/tests/fast/ray/rollout/test_server_cell.py +++ b/tests/fast/ray/rollout/test_server_cell.py @@ -47,6 +47,59 @@ async def test_check_weights_forwards_all_arguments_to_the_primary_engine(self): ) +def _addressed_cell( + *, worker_type: str = "regular", bootstrap_port: int | None = None, **args_overrides +) -> ServerCell: + engines = [ServerEngine(), ServerEngine()] + for index, engine in enumerate(engines): + engine.mark_allocated_uninitialized(fake_actor_handle()) + engine.set_addressing( + AddrInfo(server_url=f"http://10.0.0.{index + 1}:3000{index}", bootstrap_port=bootstrap_port) + ) + engine.mark_alive() + return ServerCell(args=make_args(num_gpus_per_node=8, **args_overrides), worker_type=worker_type, engines=engines) + + +class TestServerCellRouterMembership: + async def test_register_publishes_the_primary_engine_url_and_worker_type(self): + """The router routes to the cell through its node-0 engine only.""" + client = MagicMock() + client.add_worker = AsyncMock() + await _addressed_cell().register(client) + client.add_worker.assert_awaited_once_with( + worker_url="http://10.0.0.1:30000", + worker_type="regular", + use_legacy_api=False, + bootstrap_port=None, + ) + + async def test_register_passes_the_bootstrap_port_of_a_prefill_worker(self): + """PD disaggregation needs the decode side to dial this port.""" + client = MagicMock() + client.add_worker = AsyncMock() + await _addressed_cell(worker_type="prefill", bootstrap_port=8998).register(client) + assert client.add_worker.await_args.kwargs["worker_type"] == "prefill" + assert client.add_worker.await_args.kwargs["bootstrap_port"] == 8998 + + async def test_unregister_removes_the_same_url_register_published(self): + """A mismatch would leave the router routing to a dead worker.""" + client = MagicMock() + client.remove_worker = AsyncMock() + await _addressed_cell().unregister(client) + client.remove_worker.assert_awaited_once_with(worker_url="http://10.0.0.1:30000", use_legacy_api=False) + + async def test_use_miles_router_pins_the_legacy_api_on_both_calls(self): + """--use-miles-router selects the query-string API for register and unregister alike.""" + client = MagicMock() + client.add_worker = AsyncMock() + client.remove_worker = AsyncMock() + cell = _addressed_cell(use_miles_router=True) + await cell.register(client) + await cell.unregister(client) + assert client.add_worker.await_args.kwargs["use_legacy_api"] is True + assert client.remove_worker.await_args.kwargs["use_legacy_api"] is True + + def _build_servers( *, num_servers: int = 1, groups_per_server: int = 1, engines_per_group: int = 2, num_gpus_per_engine: int = 1 ) -> dict[str, RolloutServer]: diff --git a/tests/fast/ray/rollout/test_server_group_router_registration.py b/tests/fast/ray/rollout/test_server_group_router_registration.py index b34e0a1165f..b23eded4443 100644 --- a/tests/fast/ray/rollout/test_server_group_router_registration.py +++ b/tests/fast/ray/rollout/test_server_group_router_registration.py @@ -115,11 +115,24 @@ async def test_registration_addresses_only_node0_of_a_multi_node_engine(): group = _build_group(events=events, num_engines=2, num_gpus_per_engine=16) with _with_recording_client(group): - await group.register_workers([0, 1]) + await group.register_workers([0]) assert [kwargs["worker_url"] for _name, kwargs in events] == ["http://10.0.0.1:30000"] +async def test_registration_skips_a_cell_that_is_not_allocated(): + """A stopped cell has no url to publish, so it must be filtered out.""" + events: list[tuple[str, dict]] = [] + group = _build_group(events=events, num_engines=2) + for engine in group.cells[0].engines: + engine.mark_stopped() + + with _with_recording_client(group): + await group.register_workers([0, 1]) + + assert [kwargs["worker_url"] for _name, kwargs in events] == ["http://10.0.0.2:30001"] + + @pytest.mark.parametrize("missing", [dict(router_ip=None), dict(router_port=None)]) async def test_registration_is_skipped_without_a_router(missing): events: list[tuple[str, dict]] = []