diff --git a/miles/dashboard/hooks.py b/miles/dashboard/hooks.py index 505f7fbba8e..424768a083b 100644 --- a/miles/dashboard/hooks.py +++ b/miles/dashboard/hooks.py @@ -387,8 +387,8 @@ def _collect_worker_infos(cells) -> list[list]: return _ray_get( [ manager_handle.get_worker_infos.remote( - pool=compute_engine_pool(model_idx=cell.model_idx, group_index=cell.group_index), - cell_index=cell.cell_index, + pool=compute_engine_pool(model_idx=cell.meta.model_idx, group_index=cell.meta.group_index), + cell_index=cell.meta.cell_index, ) for cell in cells ] @@ -409,7 +409,7 @@ def _compute_engine_infos(cells, worker_infos_per_cell) -> list[EngineInfo]: engines.append( EngineInfo( addr=cell.addr_info.server_url, - worker_type=cell.worker_type, + worker_type=cell.meta.worker_type, engine_rank=engine_rank, gpus=[ [info.self_addrs["primary"].host.strip("[]"), gpu_id] diff --git a/miles/ray/rollout/rollout_server.py b/miles/ray/rollout/rollout_server.py index 50669d5039d..767c2668a0f 100644 --- a/miles/ray/rollout/rollout_server.py +++ b/miles/ray/rollout/rollout_server.py @@ -7,7 +7,7 @@ from miles.backends.sglang_utils.sglang_config import resolve_sglang_config from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient from miles.ray.rollout.router_manager import wait_router_ready -from miles.ray.rollout.server_cell import ServerCell, compute_nodes_per_engine +from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata, compute_nodes_per_engine logger = logging.getLogger(__name__) @@ -56,17 +56,19 @@ async def start_rollout_servers(args) -> dict[str, "RolloutServer"]: cell_id = format_cell_id(server_id=model_cfg.name, index=len(server_cells)) server_cells[cell_id] = ServerCell( args=args, - worker_type=group_cfg.worker_type, - cell_id=cell_id, - num_gpus_per_engine=gpus_per_engine, - gpu_offset=group_cfg.gpu_offset + cell_start * num_gpu_per_engine_local, - sglang_overrides=group_cfg.overrides, - model_idx=model_idx, - group_index=group_index, - cell_index=cell_start // nodes_per_engine, - needs_offload=group_cfg.needs_offload, - model_path=group_cfg.model_path, - update_weights=model_cfg.update_weights, + meta=ServerCellMetadata( + worker_type=group_cfg.worker_type, + cell_id=cell_id, + num_gpus_per_engine=gpus_per_engine, + gpu_offset=group_cfg.gpu_offset + cell_start * num_gpu_per_engine_local, + sglang_overrides=group_cfg.overrides, + model_idx=model_idx, + group_index=group_index, + cell_index=cell_start // nodes_per_engine, + needs_offload=group_cfg.needs_offload, + model_path=group_cfg.model_path, + update_weights=model_cfg.update_weights, + ), ) servers[model_cfg.name] = RolloutServer( @@ -112,11 +114,11 @@ def clear_has_new_engines(self): @property def engine_gpu_counts(self) -> list[int]: """Per-engine GPU count for all node-0 engines, parallel to ``engines``.""" - return [cell.num_gpus_per_engine for cell in self.server_cells.values()] + return [cell.meta.num_gpus_per_engine for cell in self.server_cells.values()] @property def engine_gpu_offsets(self) -> list[int]: - return [cell.gpu_offset for cell in self.server_cells.values()] + return [cell.meta.gpu_offset for cell in self.server_cells.values()] async def probe_and_mark_dead(self): """Mark unreachable cells stopped so ``recover`` restarts them. @@ -138,12 +140,12 @@ async def stop_cells(self, cell_ids: list[str]): async def offload(self, tags: list[str] | None = None): return await asyncio.gather( - *[cell.offload(tags=tags) for cell in self._allocated_cells_of() if cell.needs_offload] + *[cell.offload(tags=tags) for cell in self._allocated_cells_of() if cell.meta.needs_offload] ) async def onload(self, tags: list[str] | None = None): return await asyncio.gather( - *[cell.onload(tags=tags) for cell in self._allocated_cells_of() if cell.needs_offload] + *[cell.onload(tags=tags) for cell in self._allocated_cells_of() if cell.meta.needs_offload] ) async def check_weights( diff --git a/miles/ray/rollout/server_cell.py b/miles/ray/rollout/server_cell.py index e1346dc575c..da8cbf9c714 100644 --- a/miles/ray/rollout/server_cell.py +++ b/miles/ray/rollout/server_cell.py @@ -18,6 +18,7 @@ StateStopped, ) from miles.ray.specs.inference import compute_engine_pool +from miles.utils.pydantic_utils import FrozenStrictBaseModel from miles.utils.workers.naming import compute_worker_name from miles.utils.workers.worker_provider.base import BaseWorkerProvider from miles.utils.workers.worker_provider.ray import RayWorkerProvider @@ -27,20 +28,24 @@ SHUTDOWN_TIMEOUT = 30 +class ServerCellMetadata(FrozenStrictBaseModel): + worker_type: Literal["regular", "prefill", "decode"] + cell_id: str + num_gpus_per_engine: int + gpu_offset: int + sglang_overrides: dict + model_idx: int + group_index: int + cell_index: int + needs_offload: bool + model_path: str | None + update_weights: bool + + @dataclass class ServerCell: args: Any - worker_type: Literal["regular", "prefill", "decode"] - cell_id: str - num_gpus_per_engine: int = 1 - gpu_offset: int = 0 - sglang_overrides: dict = dataclasses.field(default_factory=dict) - model_idx: int = 0 - group_index: int = 0 - cell_index: int = 0 - needs_offload: bool = False - model_path: str | None = None - update_weights: bool = True + meta: ServerCellMetadata _state: CellState = dataclasses.field(default_factory=StateStopped) @property @@ -63,11 +68,11 @@ def api_client(self) -> SGLangApiClient: @property def _pool_id(self) -> str: - return compute_engine_pool(model_idx=self.model_idx, group_index=self.group_index) + return compute_engine_pool(model_idx=self.meta.model_idx, group_index=self.meta.group_index) async def start_engines(self) -> None: - assert not ({"host", "port"} & set(self.sglang_overrides)), ( - f"sglang_overrides must not override host/port ({self.sglang_overrides=}): the rollout process derives " + assert not ({"host", "port"} & set(self.meta.sglang_overrides)), ( + f"sglang_overrides must not override host/port ({self.meta.sglang_overrides=}): the rollout process derives " f"each engine's url from the addr allocator, so an override would make it talk to the wrong endpoint" ) assert not self.is_allocated, "the caller starts only stopped cells" @@ -80,7 +85,7 @@ async def start_engines(self) -> None: self._mark_allocated_uninitialized() provider: BaseWorkerProvider = RayWorkerProvider.create() # TODO inject instance - worker_name = compute_worker_name(pool=self._pool_id, cell_index=self.cell_index) + worker_name = compute_worker_name(pool=self._pool_id, cell_index=self.meta.cell_index) master_addrs = await provider.get_addrs(worker_name=worker_name) primary = master_addrs["primary"] disaggregation_bootstrap = master_addrs.get("disaggregation_bootstrap") @@ -93,15 +98,15 @@ async def start_engines(self) -> None: await wait_server_healthy( server_url=self.addr_info.server_url, - api_key=compute_api_key(self.args, sglang_overrides=self.sglang_overrides), + api_key=compute_api_key(self.args, sglang_overrides=self.meta.sglang_overrides), ) async def start(self, router_api_client: SGLangRouterApiClient, recover: bool = False) -> None: await self.start_engines() - if recover and self.needs_offload: + if recover and self.meta.needs_offload: await self.api_client.release_memory_occupation() - if self.update_weights or self.model_path: + if self.meta.update_weights or self.meta.model_path: await self.api_client.resume_memory_occupation(tags=[GPU_MEMORY_TYPE_WEIGHTS]) self._mark_alive() @@ -113,9 +118,11 @@ async def stop(self, router_api_client: SGLangRouterApiClient) -> None: try: await asyncio.wait_for(self.unregister(router_api_client), timeout=SHUTDOWN_TIMEOUT) except Exception as e: - logger.warning(f"Unregistering cell {self.cell_id} from the router failed, tearing down anyway ({e})") + logger.warning( + f"Unregistering cell {self.meta.cell_id} from the router failed, tearing down anyway ({e})" + ) else: - logger.info(f"Cell {self.cell_id} is already stopped") + logger.info(f"Cell {self.meta.cell_id} is already stopped") self._mark_stopped() def _mark_allocated_uninitialized(self) -> None: @@ -145,10 +152,10 @@ def _change_state( old_state_cls: type[CellState] | tuple[type[CellState], ...], new_state: CellState, ) -> None: - logger.info(f"Cell {self.cell_id} {debug_name} start old={self._state}") + logger.info(f"Cell {self.meta.cell_id} {debug_name} start old={self._state}") assert isinstance(self._state, old_state_cls), f"{self._state=}" self._state = new_state - logger.info(f"Cell {self.cell_id} {debug_name} end new={self._state}") + logger.info(f"Cell {self.meta.cell_id} {debug_name} end new={self._state}") async def probe_and_mark_dead(self) -> None: if not self.is_allocated: @@ -173,7 +180,7 @@ async def check_weights(self, action: str, allow_quant_error: bool, selector: st async def register(self, router_api_client: SGLangRouterApiClient) -> None: await router_api_client.add_worker( worker_url=self.addr_info.server_url, - worker_type=self.worker_type, + worker_type=self.meta.worker_type, use_legacy_api=use_legacy_router_api(self.args), bootstrap_port=self.addr_info.bootstrap_port, ) diff --git a/tests/fast/dashboard/test_hooks.py b/tests/fast/dashboard/test_hooks.py index 9b88b47e380..01686ef925d 100644 --- a/tests/fast/dashboard/test_hooks.py +++ b/tests/fast/dashboard/test_hooks.py @@ -6,6 +6,7 @@ from miles.dashboard import backend, hooks from miles.dashboard.hooks import BATCH_MAX_EVENTS, BATCH_MAX_SECONDS, _Identity from miles.dashboard.store import Role +from miles.ray.rollout.server_cell import ServerCellMetadata from miles.utils.timer import Timer from miles.utils.workers.ray_worker_manager import RayWorkerManager from miles.utils.workers.worker_info import WorkerInfo @@ -161,11 +162,20 @@ class FakeCell: """Duck-typed ServerCell: the hooks read only the driver-side routing facts.""" def __init__(self, url, cell_index=0, alive=True): - self.model_idx = 0 - self.group_index = 0 - self.cell_index = cell_index + self.meta = ServerCellMetadata( + worker_type="regular", + cell_id=f"cell-{cell_index}", + num_gpus_per_engine=1, + gpu_offset=0, + sglang_overrides={}, + model_idx=0, + group_index=0, + cell_index=cell_index, + needs_offload=False, + model_path=None, + update_weights=False, + ) self.addr_info = type("FakeAddrInfo", (), {"server_url": url})() - self.worker_type = "regular" self.is_alive = alive diff --git a/tests/fast/ray/rollout/conftest.py b/tests/fast/ray/rollout/conftest.py index 65b55d0079f..509eda6cdc6 100644 --- a/tests/fast/ray/rollout/conftest.py +++ b/tests/fast/ray/rollout/conftest.py @@ -2,7 +2,7 @@ import textwrap from argparse import ArgumentParser, Namespace -from typing import Any +from typing import TYPE_CHECKING, Any from unittest.mock import MagicMock import pytest @@ -12,6 +12,9 @@ from miles.utils import object_store from miles.utils.types import Sample +if TYPE_CHECKING: + from miles.ray.rollout.server_cell import ServerCell + def fake_actor_handle() -> MagicMock: """MagicMock that passes ``isinstance(x, ray.actor.ActorHandle)``. @@ -299,19 +302,27 @@ def make_dataclass_cells( num_cells: int = 2, num_gpus_per_engine: int = 1, gpu_offset: int = 0, -): +) -> list[ServerCell]: """Build configured ``ServerCell``s. Each cell starts unallocated.""" - from miles.ray.rollout.server_cell import ServerCell + from miles.ray.rollout.server_cell import ServerCell, ServerCellMetadata args = make_args(num_gpus_per_node=8) return [ ServerCell( args=args, - worker_type="regular", - cell_id=f"cell-{cell_index}", - num_gpus_per_engine=num_gpus_per_engine, - gpu_offset=gpu_offset + cell_index * min(num_gpus_per_engine, 8), - cell_index=cell_index, + meta=ServerCellMetadata( + worker_type="regular", + cell_id=f"cell-{cell_index}", + num_gpus_per_engine=num_gpus_per_engine, + gpu_offset=gpu_offset + cell_index * min(num_gpus_per_engine, 8), + sglang_overrides={}, + model_idx=0, + group_index=0, + cell_index=cell_index, + needs_offload=False, + model_path=None, + update_weights=False, + ), ) for cell_index in range(num_cells) ] diff --git a/tests/fast/ray/rollout/test_rollout_server.py b/tests/fast/ray/rollout/test_rollout_server.py index e712288ae26..2f875158ed0 100644 --- a/tests/fast/ray/rollout/test_rollout_server.py +++ b/tests/fast/ray/rollout/test_rollout_server.py @@ -283,12 +283,12 @@ def _cells_for(self, tmp_path, *, num_gpus: int, num_gpus_per_engine: int): def test_a_single_node_engine_becomes_its_own_cell(self, stub_engine_startup, tmp_path): """With one gpu per engine on 8-gpu nodes, every engine is a one-engine cell.""" cells = self._cells_for(tmp_path, num_gpus=8, num_gpus_per_engine=1) - assert [cell.cell_index for cell in cells] == list(range(8)) + assert [cell.meta.cell_index for cell in cells] == list(range(8)) def test_a_multi_node_engine_chunks_its_node_ranks_into_one_cell(self, stub_engine_startup, tmp_path): """With 16 gpus per engine on 8-gpu nodes, the 32 gpus collapse into two cells.""" cells = self._cells_for(tmp_path, num_gpus=32, num_gpus_per_engine=16) - assert [cell.cell_index for cell in cells] == [0, 1] + assert [cell.meta.cell_index for cell in cells] == [0, 1] def test_a_trailing_partial_multi_node_engine_is_rejected(self, stub_engine_startup, tmp_path): """24 gpus do not divide into whole 2-node engines, so startup must fail fast.""" @@ -298,4 +298,4 @@ def test_a_trailing_partial_multi_node_engine_is_rejected(self, stub_engine_star def test_cells_carry_contiguous_gpu_offsets(self, stub_engine_startup, tmp_path): """Each multi-node cell's gpu span starts where the previous one ended.""" cells = self._cells_for(tmp_path, num_gpus=32, num_gpus_per_engine=16) - assert [cell.gpu_offset for cell in cells] == [0, 16] + assert [cell.meta.gpu_offset for cell in cells] == [0, 16]