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
5 changes: 5 additions & 0 deletions miles/backends/sglang_utils/sglang_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -161,6 +161,11 @@ def resolve(
default_model_path: str,
gpu_offset_cursor: "_MutableBox",
) -> "ServerGroupConfig":
assert not ({"host", "port"} & set(raw.overrides)), (
f"sglang_overrides must not override host/port ({raw.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"
)

rollout_pg_offset = _compute_rollout_offset(args)
megatron_num_gpus = _compute_megatron_num_gpus(args)

Expand Down
4 changes: 0 additions & 4 deletions miles/backends/sglang_utils/sglang_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,10 +82,6 @@ def compute_engine_launch_cmd(
return shlex.join([sys.executable, "-m", "sglang.launch_server", *server_args_to_argv(launch_args)])


def compute_api_key(args, *, sglang_overrides: dict) -> str | None:
return sglang_overrides.get("api_key", args.sglang_api_key)


def _compute_server_args(
args,
*,
Expand Down
16 changes: 6 additions & 10 deletions miles/dashboard/hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
)
from miles.utils.lifecycle import TrajectoryLifecycle
from miles.utils.timer import Timer
from miles.utils.workers.naming import parse_worker_name

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -380,19 +381,14 @@ def _alive_engine_cells(servers) -> list:


def _collect_worker_infos(cells) -> list[list]:
from miles.ray.specs.inference import compute_engine_pool
from miles.utils.workers.ray_worker_manager import RayWorkerManager

manager_handle = RayWorkerManager.get_handle()
return _ray_get(
[
manager_handle.get_worker_infos.remote(
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
]
)
futures = []
for cell in cells:
pool, cell_index, _ = parse_worker_name(cell.meta.worker_name)
futures.append(manager_handle.get_worker_infos.remote(pool=pool, cell_index=cell_index))
return _ray_get(futures)


def _compute_engine_infos(cells, worker_infos_per_cell) -> list[EngineInfo]:
Expand Down
13 changes: 8 additions & 5 deletions miles/ray/rollout/rollout_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,8 @@
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, ServerCellMetadata, compute_nodes_per_engine
from miles.ray.specs.inference import compute_engine_pool
from miles.utils.workers.naming import compute_worker_name

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -63,17 +65,18 @@ async def start_rollout_servers(args) -> dict[str, "RolloutServer"]:
)
for cell_start in range(0, num_engines, nodes_per_engine):
cell_id = format_cell_id(server_id=model_cfg.name, index=cell_count)
cell_index = cell_start // nodes_per_engine
pool = compute_engine_pool(model_idx=model_idx, group_index=group_index)
worker_name = compute_worker_name(pool=pool, cell_index=cell_index)
cell_meta = ServerCellMetadata(
model_id=model_cfg.name,
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,
sglang_api_key=group_cfg.overrides.get("api_key", args.sglang_api_key),
worker_name=worker_name,
needs_offload=group_cfg.needs_offload,
model_path=group_cfg.model_path,
update_weights=model_cfg.update_weights,
)
cell_count += 1
Expand Down
25 changes: 6 additions & 19 deletions miles/ray/rollout/server_cell.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
from sglang.srt.constants import GPU_MEMORY_TYPE_WEIGHTS

from miles.backends.sglang_utils.sglang_api_client import SGLangApiClient, wait_server_healthy
from miles.backends.sglang_utils.sglang_engine import build_server_url, compute_api_key
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.cell_state import (
AddrInfo,
Expand All @@ -16,9 +16,7 @@
StateAllocatedBase,
StateAllocatedUninitialized,
)
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 @@ -28,16 +26,14 @@


class ServerCellMetadata(FrozenStrictBaseModel):
model_id: str
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
sglang_api_key: str | None
worker_name: str
needs_offload: bool
model_path: str | None
update_weights: bool


Expand Down Expand Up @@ -65,23 +61,14 @@ def addr_info(self) -> AddrInfo:
def api_client(self) -> SGLangApiClient:
return SGLangApiClient(server_url=self.addr_info.server_url)

@property
def _pool_id(self) -> str:
return compute_engine_pool(model_idx=self.meta.model_idx, group_index=self.meta.group_index)

async def _add_raw(self) -> None:
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"
)
if self.args.rollout_external:
raise NotImplementedError(
"external rollout address allocation was removed and a new implementation is coming"
)

provider: BaseWorkerProvider = RayWorkerProvider.create() # TODO inject instance
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)
master_addrs = await provider.get_addrs(worker_name=self.meta.worker_name)
primary = master_addrs["primary"]
disaggregation_bootstrap = master_addrs.get("disaggregation_bootstrap")
# TODO simplify (remove) later
Expand All @@ -94,7 +81,7 @@ async def _add_raw(self) -> None:

await wait_server_healthy(
server_url=self.addr_info.server_url,
api_key=compute_api_key(self.args, sglang_overrides=self.meta.sglang_overrides),
api_key=self.meta.sglang_api_key,
)

async def add(self, router_api_client: SGLangRouterApiClient, recover: bool = False) -> None:
Expand Down
7 changes: 7 additions & 0 deletions miles/utils/workers/naming.py
Original file line number Diff line number Diff line change
@@ -1,2 +1,9 @@
# TODO refactor & move later
def compute_worker_name(*, pool_id: str, cell_index: int = 0, worker_in_cell_index: int = 0) -> str:
return f"{pool_id}-{cell_index}-{worker_in_cell_index}"


# TODO refactor & move later
def parse_worker_name(worker_name: str) -> tuple[str, int, int]:
pool_id, cell_index, worker_in_cell_index = worker_name.rsplit("-", maxsplit=2)
return pool_id, int(cell_index), int(worker_in_cell_index)
32 changes: 32 additions & 0 deletions tests/fast/backends/sglang_utils/test_sglang_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,3 +288,35 @@ def test_an_explicit_memory_saver_override_wins(self, tmp_path):
group = cfg.models[0].server_groups[0]
assert group.needs_offload is False
assert group.overrides["enable_memory_saver"] is True


class TestHostPortOverrideRejection:
def test_a_port_override_is_rejected_at_resolve_time(self, tmp_path):
"""Overriding the allocator-owned port must fail fast instead of desyncing engine and controller."""
with pytest.raises(AssertionError, match="must not override host/port"):
_resolve_yaml(
tmp_path,
"sglang:\n"
" - name: actor\n"
" server_groups:\n"
" - worker_type: regular\n"
" num_gpus: 8\n"
" overrides:\n"
" port: 12345\n",
rollout_num_gpus=8,
)

def test_a_host_override_is_rejected_at_resolve_time(self, tmp_path):
"""Overriding the allocator-owned host must fail fast as well."""
with pytest.raises(AssertionError, match="must not override host/port"):
_resolve_yaml(
tmp_path,
"sglang:\n"
" - name: actor\n"
" server_groups:\n"
" - worker_type: regular\n"
" num_gpus: 8\n"
" overrides:\n"
" host: 10.0.0.1\n",
rollout_num_gpus=8,
)
18 changes: 1 addition & 17 deletions tests/fast/backends/sglang_utils/test_sglang_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
pytest.importorskip("sglang")

from miles.backends.sglang_utils.server_args_utils import parse_server_args_argv
from miles.backends.sglang_utils.sglang_engine import compute_api_key, compute_engine_launch_cmd
from miles.backends.sglang_utils.sglang_engine import compute_engine_launch_cmd


def _cmd(*, worker_type: str = "regular", args=None, addr_overrides: dict | None = None, **kwargs) -> str:
Expand Down Expand Up @@ -75,19 +75,3 @@ def test_the_command_carries_the_api_key_from_args(self):
cmd = _cmd(args=make_engine_args(sglang_api_key="secret"))
parsed = parse_server_args_argv(shlex.split(cmd)[3:])
assert parsed.api_key == "secret"


class TestComputeApiKey:
def test_the_args_key_is_used_when_no_override_exists(self):
"""The health wait needs the same key the generic passthrough gave the server."""
args = make_engine_args(sglang_api_key="secret")
assert compute_api_key(args, sglang_overrides={}) == "secret"

def test_an_override_key_wins_over_the_args_key(self):
"""Overrides beat args exactly like they do in the rendered command."""
args = make_engine_args(sglang_api_key="from-args")
assert compute_api_key(args, sglang_overrides={"api_key": "from-override"}) == "from-override"

def test_no_key_anywhere_means_the_health_wait_sends_none(self):
"""No key configured means the health wait sends none."""
assert compute_api_key(make_engine_args(), sglang_overrides={}) is None
Loading