From b652669a259c7baa7f149c3ad876b7633af98f5d Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Thu, 30 Jul 2026 18:03:53 +0800 Subject: [PATCH] Pass node_rank instead of rank when computing the engine launch command --- miles/backends/sglang_utils/sglang_engine.py | 7 ++- miles/ray/rollout/server_cell.py | 2 +- .../sglang_utils/test_server_args_utils.py | 6 +-- .../sglang_utils/test_sglang_engine.py | 2 +- tests/fast/ray/rollout/test_server_cell.py | 47 ++++++++++++++++++- 5 files changed, 54 insertions(+), 10 deletions(-) diff --git a/miles/backends/sglang_utils/sglang_engine.py b/miles/backends/sglang_utils/sglang_engine.py index 627d6f77d9b..348cce888c5 100644 --- a/miles/backends/sglang_utils/sglang_engine.py +++ b/miles/backends/sglang_utils/sglang_engine.py @@ -61,7 +61,7 @@ class EngineLaunchPlan: def compute_engine_launch_plan( args, *, - rank: int, + node_rank: int, worker_type: str, base_gpu_id: int, sglang_overrides: dict, @@ -70,7 +70,7 @@ def compute_engine_launch_plan( ) -> EngineLaunchPlan: server_args_dict = _compute_server_args( args, - rank=rank, + node_rank=node_rank, dist_init_addr=addr_and_ports["dist_init_addr"], nccl_port=addr_and_ports["nccl_port"], host=addr_and_ports["host"], @@ -91,7 +91,7 @@ def compute_engine_launch_plan( def _compute_server_args( args, *, - rank, + node_rank: int, dist_init_addr, nccl_port, host, @@ -105,7 +105,6 @@ def _compute_server_args( ): _gpus_per_engine = num_gpus_per_engine or args.rollout_num_gpus_per_engine nnodes = max(1, _gpus_per_engine // args.num_gpus_per_node) - node_rank = rank % nnodes base = _to_local_gpu_id(base_gpu_id) kwargs = { "model_path": args.hf_checkpoint, diff --git a/miles/ray/rollout/server_cell.py b/miles/ray/rollout/server_cell.py index ec64f938a4d..29dfab043a5 100644 --- a/miles/ray/rollout/server_cell.py +++ b/miles/ray/rollout/server_cell.py @@ -159,7 +159,7 @@ async def start_engines(self, port_allocator: PortAllocator) -> None: plans = { rank: compute_engine_launch_plan( self.args, - rank=rank, + node_rank=local_index, worker_type=self.worker_type, base_gpu_id=self.engine_gpu_ids[local_index][0], sglang_overrides=self.sglang_overrides, diff --git a/tests/fast/backends/sglang_utils/test_server_args_utils.py b/tests/fast/backends/sglang_utils/test_server_args_utils.py index de9ad681850..7b7f9eed3d7 100644 --- a/tests/fast/backends/sglang_utils/test_server_args_utils.py +++ b/tests/fast/backends/sglang_utils/test_server_args_utils.py @@ -40,7 +40,7 @@ def _server_args( *, worker_type: str = "regular", - rank: int = 0, + node_rank: int = 0, dist_init_addr: str = "10.0.0.1:20000", args: Namespace | None = None, sglang_overrides: dict | None = None, @@ -52,7 +52,7 @@ def _server_args( overrides = {"device": "cuda", **(sglang_overrides or {})} server_args_dict = _compute_server_args( args or _args(), - rank=rank, + node_rank=node_rank, dist_init_addr=dist_init_addr, nccl_port=20031, host="10.0.0.1", @@ -118,7 +118,7 @@ def test_a_decode_worker_roundtrips(self): def test_a_multi_node_rank_roundtrips(self): """nnodes, node_rank and tp_size of a multi-node engine survive the boundary.""" server_args = _server_args( - rank=1, + node_rank=1, num_gpus_per_engine=16, args=_args(rollout_num_gpus_per_engine=16), ) diff --git a/tests/fast/backends/sglang_utils/test_sglang_engine.py b/tests/fast/backends/sglang_utils/test_sglang_engine.py index 011ce984556..5d91ca73e8b 100644 --- a/tests/fast/backends/sglang_utils/test_sglang_engine.py +++ b/tests/fast/backends/sglang_utils/test_sglang_engine.py @@ -23,7 +23,7 @@ def _plan(*, worker_type: str = "regular", args=None, addr_overrides: dict | Non addr_and_ports.update(addr_overrides or {}) return compute_engine_launch_plan( args or make_engine_args(), - rank=0, + node_rank=0, worker_type=worker_type, base_gpu_id=0, sglang_overrides={}, diff --git a/tests/fast/ray/rollout/test_server_cell.py b/tests/fast/ray/rollout/test_server_cell.py index 7cf9daca164..8124dcf495a 100644 --- a/tests/fast/ray/rollout/test_server_cell.py +++ b/tests/fast/ray/rollout/test_server_cell.py @@ -1,13 +1,18 @@ from __future__ import annotations +import asyncio +from typing import Any from unittest.mock import AsyncMock, MagicMock, patch import pytest -from tests.fast.ray.rollout.conftest import fake_actor_handle, make_args +from tests.fast.ray.rollout.conftest import fake_actor_handle, fake_engine, make_args +import miles.ray.rollout.server_cell as server_cell_module from miles.ray.rollout.cell_state import AddrInfo from miles.ray.rollout.rollout_server import RolloutServer, format_cell_id, list_cell_ids from miles.ray.rollout.server_cell import ServerCell, compute_nodes_per_engine +from miles.utils.test_utils.mock_sglang_engine import parse_cmd_flags +from miles.utils.workers.addr_allocator import PortAllocator def _allocated_cell(num_nodes: int = 1, *, alive: bool = True, addressed: bool = True) -> ServerCell: @@ -136,6 +141,46 @@ async def test_check_weights_forwards_all_arguments_to_the_primary_engine(self): ) +def _launch_command_flags(*, rank_offset: int, num_nodes: int) -> list[dict[str, Any]]: + num_gpus_per_engine: int = 8 * num_nodes + actors: list[MagicMock] = [fake_engine(host=f"10.0.0.{index + 1}", port_seed=30000) for index in range(num_nodes)] + for actor in actors: + actor.run.remote.side_effect = lambda **kwargs: asyncio.sleep(0) + + cell = ServerCell( + args=make_args( + num_gpus_per_node=8, + sglang_pp_size=1, + sglang_ep_size=1, + multi_lora=False, + rollout_num_gpus_per_engine=num_gpus_per_engine, + ), + worker_type="regular", + cell_id="cell-1", + num_nodes=num_nodes, + num_gpus_per_engine=num_gpus_per_engine, + rank_offset=rank_offset, + pg=(None, [], list(range(8)) * num_nodes), + ) + + pending: list[MagicMock] = list(actors) + with ( + patch.object(server_cell_module, "launch_sglang_ray_actor", side_effect=lambda **kwargs: pending.pop(0)), + patch.object(server_cell_module, "wait_server_healthy", new=AsyncMock()), + ): + asyncio.run(cell.start_engines(PortAllocator())) + + return [parse_cmd_flags(actor.run.remote.call_args.kwargs["cmd"]) for actor in actors] + + +class TestMultiNodeEngineNodeRank: + def test_the_second_two_node_cell_numbers_its_own_nodes_from_zero(self, patch_ray_get): + """--node-rank is cell-local, so the cell at rank_offset=2 launches node-ranks 0 and 1, not 2 and 3.""" + flags = _launch_command_flags(rank_offset=2, num_nodes=2) + assert [entry["nnodes"] for entry in flags] == [2, 2] + assert [entry.get("node_rank", 0) for entry in flags] == [0, 1] + + def _addressed_cell( *, worker_type: str = "regular", bootstrap_port: int | None = None, **args_overrides ) -> ServerCell: