Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions miles/dashboard/hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
]
Expand All @@ -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]
Expand Down
34 changes: 18 additions & 16 deletions miles/ray/rollout/rollout_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)

Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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.
Expand All @@ -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(
Expand Down
53 changes: 30 additions & 23 deletions miles/ray/rollout/server_cell.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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"
Expand All @@ -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")
Expand All @@ -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()
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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,
)
Expand Down
18 changes: 14 additions & 4 deletions tests/fast/dashboard/test_hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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


Expand Down
27 changes: 19 additions & 8 deletions tests/fast/ray/rollout/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)``.
Expand Down Expand Up @@ -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)
]
Expand Down
6 changes: 3 additions & 3 deletions tests/fast/ray/rollout/test_rollout_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -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]
Loading