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
69 changes: 0 additions & 69 deletions miles/ray/rollout/addr_allocator.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,3 @@
import functools
import logging
from dataclasses import dataclass

Expand Down Expand Up @@ -27,71 +26,3 @@ def alloc(self, *, engine, node_ip: str, consecutive: int = 1) -> int:
)
self._values[node_ip] = port + consecutive
return port


# NOTE: May re-implement this in a potentially easier way if needed
def allocate_rollout_engine_addr_and_ports_normal(
*,
args,
port_allocator: PortAllocator,
rollout_engines,
worker_type="regular",
num_gpus_per_engine=None,
rank_offset=0,
):
# get ports
# there are 4 ports we need to allocate
# 1. server port
# 2. nccl port
# 3. dist_init_addr port
# 4. other ports for dp_attention, which is of size 4 + dp_size
_gpus_per_engine = num_gpus_per_engine or args.rollout_num_gpus_per_engine
num_engines_per_node = max(1, args.num_gpus_per_node // _gpus_per_engine)
addr_and_ports: dict[int, dict] = {}

visited_nodes = set()
for rank, engine in rollout_engines:
local_rank = rank - rank_offset
node_index = local_rank // num_engines_per_node
if node_index in visited_nodes:
continue
visited_nodes.add(node_index)
# TODO: currently when restarting engines, we will set port for all engines on this node starting with this rank.
# e.g. for 8 gpus, if we are restarting engine on gpu 3, we will set port for engine 3,4,5,6,7 on this node.
num_engines_on_this_node = num_engines_per_node - (local_rank % num_engines_per_node)

node_ip, _ = ray.get(engine._get_current_node_ip_and_free_port.remote())

get_port = functools.partial(port_allocator.alloc, engine=engine, node_ip=node_ip)

for i in range(num_engines_on_this_node):
current_rank = rank + i
addr_and_ports.setdefault(current_rank, {})
addr_and_ports[current_rank]["host"] = node_ip
addr_and_ports[current_rank]["port"] = get_port()
addr_and_ports[current_rank]["nccl_port"] = get_port()
# Always allocate a unique engine_info_bootstrap_port per engine
addr_and_ports[current_rank]["engine_info_bootstrap_port"] = get_port()

if worker_type == "prefill":
addr_and_ports[current_rank]["disaggregation_bootstrap_port"] = get_port()

if _gpus_per_engine > args.num_gpus_per_node:
num_node_per_engine = _gpus_per_engine // args.num_gpus_per_node
if local_rank % num_node_per_engine == 0:
dist_init_addr = f"{node_ip}:{get_port(consecutive=30 + args.sglang_dp_size)}"
for i in range(num_node_per_engine):
addr_and_ports.setdefault(rank + i, {})
addr_and_ports[rank + i]["dist_init_addr"] = dist_init_addr
else:
for i in range(num_engines_on_this_node):
addr_and_ports[rank + i][
"dist_init_addr"
] = f"{node_ip}:{get_port(consecutive=30 + args.sglang_dp_size)}"

for i, _ in rollout_engines:
for key in ["port", "nccl_port", "dist_init_addr"]:
assert key in addr_and_ports[i], f"Engine {i} {key} is not set."
logger.info(f"Ports for engine {i}: {addr_and_ports[i]}")

return addr_and_ports
33 changes: 24 additions & 9 deletions miles/ray/rollout/server_group.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,15 @@
import asyncio
import dataclasses
import functools
import logging
from typing import Any

import ray
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS

from miles.backends.sglang_utils.sglang_engine import 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, allocate_rollout_engine_addr_and_ports_normal
from miles.ray.rollout.addr_allocator import PortAllocator
from miles.ray.rollout.server_cell import SHUTDOWN_TIMEOUT, ServerCell, flatten_cells, launch_sglang_ray_actor
from miles.ray.rollout.server_engine import AddrInfo, ServerEngine
from miles.utils import async_utils
Expand Down Expand Up @@ -117,14 +119,27 @@ def start_engines(
if curr_num_new_engines == 0:
return [], []

addr_and_ports = allocate_rollout_engine_addr_and_ports_normal(
args=self.args,
port_allocator=port_allocator,
rollout_engines=new_engines,
worker_type=self.worker_type,
num_gpus_per_engine=self.num_gpus_per_engine,
rank_offset=self.rank_offset,
)
addr_and_ports: dict[int, dict[str, Any]] = {}
for cell_index in sorted({index // self.nodes_per_engine for index in new_engine_indices}):
dist_init_addr = None
for engine_in_cell_index in range(self.nodes_per_engine):
actor = self.cells[cell_index].engines[engine_in_cell_index].actor_handle
node_ip, _ = ray.get(actor._get_current_node_ip_and_free_port.remote())
alloc = functools.partial(port_allocator.alloc, engine=actor, node_ip=node_ip)

if engine_in_cell_index == 0:
dist_init_addr = f"{node_ip}:{alloc(consecutive=30 + self.args.sglang_dp_size)}"

rank = self.rank_offset + cell_index * self.nodes_per_engine + engine_in_cell_index
addr_and_ports[rank] = dict(
host=node_ip,
port=alloc(),
nccl_port=alloc(),
engine_info_bootstrap_port=alloc(),
dist_init_addr=dist_init_addr,
)
if self.worker_type == "prefill":
addr_and_ports[rank]["disaggregation_bootstrap_port"] = alloc()

for index, _ in new_engines:
engine_addr_and_ports = addr_and_ports[index]
Expand Down
5 changes: 4 additions & 1 deletion tests/fast/ray/rollout/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -296,8 +296,11 @@ def fake_engine(host: str = "10.0.0.1", port_seed: int = 30000) -> MagicMock:

Mocks ``_get_current_node_ip_and_free_port.remote(start_port, consecutive)``
with a deterministic ``max(seq, start_port)`` counter so allocator tests
can predict and assert on port assignment."""
can predict and assert on port assignment. It also passes
``isinstance(x, ray.actor.ActorHandle)`` so it can be handed to
``ServerEngine.mark_allocated_uninitialized`` (see ``fake_actor_handle``)."""
e = MagicMock()
e._spec_class = ray.actor.ActorHandle
e._port_cursor = port_seed

def _alloc(start_port: int = 15000, consecutive: int = 1):
Expand Down
6 changes: 2 additions & 4 deletions tests/fast/ray/rollout/real_ray/test_server_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -173,10 +173,8 @@ def test_stop_handles_shutdown_failure_gracefully(self, patched_sglang_engine, p


class TestStartEnginesRealAllocator:
"""Drive ``start_engines`` with the real
``allocate_rollout_engine_addr_and_ports_normal`` (no stub) so that the
actor → driver port round-trip via
``_get_current_node_ip_and_free_port.remote`` actually runs."""
"""Drive ``start_engines`` with real actors so that the actor → driver port
round-trip via ``_get_current_node_ip_and_free_port.remote`` actually runs."""

def test_real_allocator_assigns_distinct_ports_via_remote_calls(
self,
Expand Down
Loading
Loading