Skip to content
Open
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
8 changes: 4 additions & 4 deletions miles/ray/rollout/rollout_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
15 changes: 15 additions & 0 deletions miles/ray/rollout/server_cell.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]
Expand Down
79 changes: 25 additions & 54 deletions miles/ray/rollout/server_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -135,26 +108,24 @@ 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()

async def recover(self, port_allocator: PortAllocator, filter_cell_indices: list[int] | None = None):
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)
Expand All @@ -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:
Expand Down
53 changes: 53 additions & 0 deletions tests/fast/ray/rollout/test_server_cell.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down
15 changes: 14 additions & 1 deletion tests/fast/ray/rollout/test_server_group_router_registration.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]] = []
Expand Down
Loading