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
32 changes: 0 additions & 32 deletions miles/backends/sglang_utils/sglang_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,14 +5,11 @@
import os

import httpx
import sglang_router
from packaging.version import parse
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import kill_process_tree

from miles.backends.megatron_utils.lora_utils import convert_target_modules_to_hf, sglang_lora_target_all_sentinel
from miles.backends.sglang_utils.sglang_api_client import wait_server_healthy
from miles.backends.sglang_utils.sglang_router_api_client import SGLangRouterApiClient
from miles.ray.ray_actor import RayActor
from miles.utils import async_utils
from miles.utils.env_report import collect_and_print_node_env_report
Expand Down Expand Up @@ -99,10 +96,6 @@ def build_server_url(host: str, port: int) -> str:
return f"http://{format_v6_uri(host)}:{port}"


def _use_legacy_router_api(args) -> bool:
return parse(sglang_router.__version__) <= parse("0.2.1") or args.use_miles_router


class SGLangEngine(RayActor):
def __init__(
self,
Expand Down Expand Up @@ -146,8 +139,6 @@ def init(
nccl_port,
host=None,
disaggregation_bootstrap_port=None,
router_ip=None,
router_port=None,
engine_info_bootstrap_port=None,
):
if env_report := self.args.env_report:
Expand All @@ -157,9 +148,6 @@ def init(
partial_env_report=env_report,
)

self.router_ip = router_ip if router_ip is not None else self.args.sglang_router_ip
self.router_port = router_port if router_port is not None else self.args.sglang_router_port

host = host or get_host_info()[1]

host = format_v6_uri(host)
Expand All @@ -186,7 +174,6 @@ def init(
self.server_port = server_args_dict["port"]

self.server_url = build_server_url(self.server_host, self.server_port)
self.router_api_client = SGLangRouterApiClient(router_url=f"http://{self.router_ip}:{self.router_port}")

if self.args.rollout_external:
self._init_external(server_args_dict, external_engine_need_check_fields=external_engine_need_check_fields)
Expand Down Expand Up @@ -221,30 +208,11 @@ def _init_normal(self, server_args_dict):
server_args = ServerArgs(**{**server_args_dict, "host": server_args_dict["host"].strip("[]")})
self.process = launch_server_process(server_args)

if self.node_rank == 0 and self.router_ip and self.router_port:
async_utils.run(
self.router_api_client.add_worker(
worker_url=self.server_url,
worker_type=self.worker_type,
use_legacy_api=_use_legacy_router_api(self.args),
bootstrap_port=(
server_args_dict["disaggregation_bootstrap_port"] if self.worker_type == "prefill" else None
),
)
)

def shutdown(self):
if self.args.rollout_external:
return

logger.info(f"Shutdown engine {self.server_host}:{self.server_port}...")
if self.node_rank == 0:
async_utils.run(
self.router_api_client.remove_worker(
worker_url=self.server_url,
use_legacy_api=_use_legacy_router_api(self.args),
)
)
kill_process_tree(self.process.pid)

def simulate_crash(self):
Expand Down
4 changes: 4 additions & 0 deletions miles/backends/sglang_utils/sglang_router_api_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,10 @@
logger = logging.getLogger(__name__)


def use_legacy_router_api(args) -> bool:
return parse(sglang_router.__version__) <= parse("0.2.1") or args.use_miles_router


@dataclasses.dataclass(frozen=True)
class SGLangRouterApiClient:
router_url: str
Expand Down
2 changes: 2 additions & 0 deletions miles/ray/rollout/rollout_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
from miles.ray.rollout.router_manager import start_router
from miles.ray.rollout.server_engine import ServerEngine
from miles.ray.rollout.server_group import ServerGroup
from miles.utils import async_utils

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -89,6 +90,7 @@ def start_rollout_servers(args, pg) -> dict[str, "RolloutServer"]:

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))

servers[model_cfg.name] = RolloutServer(
server_groups=server_groups,
Expand Down
25 changes: 16 additions & 9 deletions miles/ray/rollout/server_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,18 +27,18 @@ def __init__(self):
def mark_allocated_uninitialized(self, actor_handle: ray.actor.ActorHandle):
self._change_state("mark_allocated", _StateStopped, _StateAllocatedUninitialized(actor_handle=actor_handle))

def set_server_url(self, server_url: str) -> None:
def set_addressing(self, addr_info: AddrInfo) -> None:
self._change_state(
"set_server_url",
"set_addressing",
_StateAllocatedUninitialized,
_StateAllocatedUninitialized(actor_handle=self.actor_handle, server_url=server_url),
_StateAllocatedUninitialized(actor_handle=self.actor_handle, addr_info=addr_info),
)

def mark_alive(self):
self._change_state(
"mark_alive",
_StateAllocatedUninitialized,
_StateAllocatedAlive(actor_handle=self.actor_handle, server_url=self.server_url),
_StateAllocatedAlive(actor_handle=self.actor_handle, addr_info=self.addr_info),
)

def mark_stopped(self):
Expand All @@ -50,14 +50,14 @@ def actor_handle(self) -> ray.actor.ActorHandle:
return self._state.actor_handle

@property
def server_url(self) -> str:
def addr_info(self) -> AddrInfo:
assert isinstance(self._state, _StateAllocatedBase)
assert self._state.server_url is not None, f"{self._state=}"
return self._state.server_url
assert self._state.addr_info is not None, f"{self._state=}"
return self._state.addr_info

@property
def api_client(self) -> SGLangApiClient:
return SGLangApiClient(server_url=self.server_url)
return SGLangApiClient(server_url=self.addr_info.server_url)

@property
def is_allocated(self) -> bool:
Expand All @@ -80,6 +80,13 @@ def _change_state(
logger.info(f"{debug_name} end new={self._state}")


class AddrInfo(BaseModel):
model_config = ConfigDict(extra="forbid", frozen=True)

server_url: str
bootstrap_port: int | None = None


# ------------------------- states -----------------------------


Expand All @@ -93,7 +100,7 @@ class _StateStopped(_StateBase):

class _StateAllocatedBase(_StateBase):
actor_handle: ray.actor.ActorHandle
server_url: str | None = None
addr_info: AddrInfo | None = None


class _StateAllocatedUninitialized(_StateAllocatedBase):
Expand Down
67 changes: 55 additions & 12 deletions miles/ray/rollout/server_group.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,14 +9,15 @@
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS

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 (
PortCursors,
allocate_rollout_engine_addr_and_ports_external,
allocate_rollout_engine_addr_and_ports_normal,
)
from miles.ray.rollout.server_engine import ServerEngine
from miles.ray.rollout.server_engine import AddrInfo, ServerEngine
from miles.ray.utils import NOSET_VISIBLE_DEVICES_ENV_VARS_LIST
from miles.utils import dumper_utils
from miles.utils import async_utils, dumper_utils

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -173,27 +174,68 @@ def start_engines(

for index, _ in new_engines:
engine_addr_and_ports = addr_and_ports[index]
self.all_engines[index - self.rank_offset].set_server_url(
build_server_url(host=engine_addr_and_ports["host"], port=engine_addr_and_ports["port"])
self.all_engines[index - self.rank_offset].set_addressing(
AddrInfo(
server_url=build_server_url(
host=engine_addr_and_ports["host"], port=engine_addr_and_ports["port"]
),
bootstrap_port=engine_addr_and_ports.get("disaggregation_bootstrap_port"),
)
)

init_handles = [
engine.init.remote(
**addr_and_ports[index],
router_ip=self.router_ip,
router_port=self.router_port,
)
for index, engine in new_engines
]
init_handles = [engine.init.remote(**addr_and_ports[index]) for index, engine in new_engines]
return init_handles, new_engine_indices

async def register_workers(self, engine_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)
]
)

async def unregister_workers(self, engine_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)
]
)

def _primary_engines_of(self, engine_indices: list[int]) -> list[ServerEngine]:
return [
self.all_engines[index]
for index in engine_indices
if index % self.nodes_per_engine == 0 and self.all_engines[index].is_allocated
]

@property
def _router_api_client(self) -> SGLangRouterApiClient:
return SGLangRouterApiClient(router_url=f"http://{self.router_ip}:{self.router_port}")

# Called from InferenceController.stop_cell (main thread, async): deliberately non-async here
# to avoid introducing two states like "stopping (but not stopped)" vs "stopped", since
# single-thread async code will not yield without an await point
# it has the drawback of freezing the whole async thread, which may be avoided later by
# moving `shutdown` mainly to local code
def stop_engines(self, engine_indices: list[int]):
logger.info(f"Killing server {engine_indices=}...")
try:
async_utils.run(asyncio.wait_for(self.unregister_workers(engine_indices), timeout=_SHUTDOWN_TIMEOUT))
except Exception as e:
logger.warning(f"Unregistering {engine_indices=} from the router failed, tearing down anyway (e: {e})")
for i in engine_indices:
engine = self.all_engines[i]
if engine.is_allocated:
Expand Down Expand Up @@ -242,6 +284,7 @@ async def recover(self, port_cursors: PortCursors, filter_indices: list[int] | N
)

self.mark_alive(engine_indices=new_engine_indices)
await self.register_workers(new_engine_indices)

def mark_alive(self, engine_indices: list[int]):
for engine_index in engine_indices:
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
from argparse import Namespace
from unittest.mock import MagicMock, patch

import pytest

_MODULE = "miles.backends.sglang_utils.sglang_engine"


def test_init_brackets_bare_ipv6_inputs_and_keeps_the_allocated_port():
"""A bare v6 host makes host:port ambiguous, so init must bracket what it forwards and what
it serves on -- at the port the allocator handed it, not a fixed one."""
pytest.importorskip("sglang")
from miles.backends.sglang_utils.sglang_engine import SGLangEngine

engine = SGLangEngine.__new__(SGLangEngine)
engine.args = Namespace(env_report=None, rollout_external=False)
engine.rank = 0
engine.worker_type = "regular"
engine.base_gpu_id = 0
engine.sglang_overrides = {}
engine.num_gpus_per_engine = 1

forwarded: dict = {}

def fake_compute_server_args(*args, **kwargs):
forwarded["dist_init_addr"] = args[2]
forwarded["host"] = args[4]
return {"node_rank": 0, "host": "fd00::2", "port": 31007, "disaggregation_bootstrap_port": None}, []

with (
patch(f"{_MODULE}._compute_server_args", side_effect=fake_compute_server_args),
patch(f"{_MODULE}.ServerArgs"),
patch(f"{_MODULE}.launch_server_process", return_value=MagicMock(pid=4242)),
):
engine.init(dist_init_addr="fd00::1:15003", port=31007, nccl_port=6000, host="fd00::2")

assert forwarded == {"dist_init_addr": "[fd00::1]:15003", "host": "[fd00::2]"}
assert engine.server_url == "http://[fd00::2]:31007"
Loading
Loading