From 8733da8cf46089a1fc7d86b36eb5f6ba62ea89c5 Mon Sep 17 00:00:00 2001 From: Sage <80211083+sagearc@users.noreply.github.com> Date: Wed, 9 Sep 2026 21:22:44 +0300 Subject: [PATCH] [rust-server] Use node-local HTTP ports for DP attention (#34430) Signed-off-by: Sage Ahrac --- python/sglang/srt/entrypoints/engine.py | 38 +++++++- python/sglang/srt/rust_server/server.py | 19 +++- .../entrypoints/test_rust_server_dp_ports.py | 96 +++++++++++++++++++ 3 files changed, 144 insertions(+), 9 deletions(-) create mode 100644 test/registered/unit/entrypoints/test_rust_server_dp_ports.py diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index 7903d9331..29575196c 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -1169,11 +1169,16 @@ class Engine(EngineScoreMixin, EngineBase): weight_cache_daemon_procs, ) - launch_dummy_health_check_server( - get_serving().host, - get_serving().port, - get_observability().enable_metrics, + # A node-local Rust listener owns the health endpoints when present. + rust_server_owns_base_port = ( + envs.SGLANG_RUST_SERVER.get() and node_hosts_rust_server() ) + if not rust_server_owns_base_port: + launch_dummy_health_check_server( + get_serving().host, + get_serving().port, + get_observability().enable_metrics, + ) scheduler_init_result.block_until_scheduler_exits() return ( @@ -1870,6 +1875,31 @@ def _calculate_rank_ranges( return pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node +def node_hosts_rust_server() -> bool: + """Whether this node contains a Rust listener rank, assuming Rust mode.""" + parallel = get_parallel() + pp_rank_range, tp_rank_range, _, _ = _calculate_rank_ranges( + parallel.nnodes, + parallel.pp_size, + parallel.tp_size, + parallel.node_rank, + ) + if 0 not in pp_rank_range: + return False + + if get_exec().moe.is_ep_scale_joiner: + # Scale joiners launch the full local TP group, including its first rank. + return True + + # Each attention DP group hosts a listener on its first rank (CP=TP=0). + ranks_per_dp_group = parallel.attn_tp_size * parallel.attn_cp_size + for tp_rank in tp_rank_range: + rank_within_dp_group = tp_rank % ranks_per_dp_group + if rank_within_dp_group == 0: + return True + return False + + def _compute_parallelism_ranks(tp_rank: int) -> Tuple[int, int, int]: """Compute attention-CP, MoE-DP, and MoE-EP ranks for a TP rank. diff --git a/python/sglang/srt/rust_server/server.py b/python/sglang/srt/rust_server/server.py index 8dd564f7b..557a933b5 100644 --- a/python/sglang/srt/rust_server/server.py +++ b/python/sglang/srt/rust_server/server.py @@ -23,7 +23,7 @@ from sglang.srt.managers.utils import ( MsgpackDecodeError, msgpack_decode_explained, ) -from sglang.srt.runtime_context import get_mm, get_serving +from sglang.srt.runtime_context import get_exec, get_mm, get_parallel, get_serving from sglang.srt.rust_server.config import _build_server_args, _partition_cores from sglang.srt.rust_server.multimodal import ( RUST_MM_FAMILIES, @@ -79,10 +79,19 @@ class RustServer: os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") server_args = scheduler.server_args - # Per-DP-rank HTTP port with client load balancing. `None` when DP is off, - # so the rank is not conflated with rank 0 of a one-rank group. + # Preserve the DP startup log; ports use node-local offsets. dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None - listen_port = get_serving().port + (dp_rank or 0) + if get_exec().moe.is_ep_scale_joiner: + # The joining TP group is entirely local to this node. + tp_size_per_node = scheduler.ps.tp_size + else: + nnodes_per_pp_rank = max(get_parallel().nnodes // scheduler.ps.pp_size, 1) + tp_size_per_node = scheduler.ps.tp_size // nnodes_per_pp_rank + dp_group_width = scheduler.ps.attn_tp_size * scheduler.ps.attn_cp_size + # Count DP leaders within this node's TP range. The first leader must + # use the base port even when a DP group spans multiple nodes. + local_dp_rank = (scheduler.ps.tp_rank % tp_size_per_node) // dp_group_width + listen_port = get_serving().port + local_dp_rank listen_addr = NetworkAddress(get_serving().host, listen_port).to_host_port_str() launch_cores, server_cores = _partition_cores( @@ -97,7 +106,7 @@ class RustServer: _build_server_args(scheduler), # None -> run unpinned; the list carries the pinning decision. cores=server_cores, - port_offset=dp_rank, + port_offset=local_dp_rank, ) # Multimodal models must have a Rust pipeline — there is no Python diff --git a/test/registered/unit/entrypoints/test_rust_server_dp_ports.py b/test/registered/unit/entrypoints/test_rust_server_dp_ports.py new file mode 100644 index 000000000..1740b0aec --- /dev/null +++ b/test/registered/unit/entrypoints/test_rust_server_dp_ports.py @@ -0,0 +1,96 @@ +from types import SimpleNamespace +from unittest.mock import patch + +import pytest + +from sglang.srt import rust_extensions +from sglang.srt.entrypoints.engine import node_hosts_rust_server +from sglang.srt.runtime_context import get_context, get_parallel +from sglang.srt.rust_server import server as rust_server +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=3, suite="base-a-test-cpu") + + +@pytest.mark.parametrize( + "nnodes,tp_size,dp_size,ep_join_mode,ranks,expected", + [ + (2, 4, 4, None, (0, 1, 2, 3), [0, 1, 0, 1]), + (4, 4, 2, None, (0, 2), [0, 0]), + (2, 2, 2, "scale", (0, 1), [0, 1]), + ], + ids=["multiple-listeners-per-node", "dp-spans-nodes", "scale-joiner"], +) +def test_dp_leaders_reuse_node_local_ports( + nnodes, tp_size, dp_size, ep_join_mode, ranks, expected +): + with ( + get_context().override_server_args( + nnodes=nnodes, + tp_size=tp_size, + dp_size=dp_size, + enable_dp_attention=True, + ep_join_mode=ep_join_mode, + host="0.0.0.0", + port=30000, + ), + patch.object(rust_extensions, "load_rust_extension") as extension, + patch.object(rust_server, "_partition_cores", return_value=(None, None)), + patch.object(rust_server, "_build_server_args"), + ): + parallel = get_parallel() + ports = [] + for dp_rank, tp_rank in enumerate(ranks): + scheduler = SimpleNamespace( + server_args=SimpleNamespace(), + ps=SimpleNamespace( + tp_rank=tp_rank, + tp_size=parallel.tp_size, + pp_size=parallel.pp_size, + attn_tp_size=parallel.attn_tp_size, + attn_cp_size=parallel.attn_cp_size, + attn_dp_rank=dp_rank, + dp_size=dp_size, + ), + model_config=SimpleNamespace(is_multimodal=False), + ) + ports.append(rust_server.RustServer.launch(scheduler).http_port) + + calls = extension.return_value.Server.call_args_list + assert [c.kwargs["port_offset"] for c in calls] == expected + # P/D bootstrap must register against the same ports Rust binds. + assert ports == [30000 + offset for offset in expected] + + +@pytest.mark.parametrize( + "pp_size,expected", + [(1, [True, False, True, False]), (2, [True, True, False, False])], +) +@pytest.mark.parametrize("node_rank", range(4)) +def test_node_listener_placement(pp_size, expected, node_rank): + with get_context().override_server_args( + nnodes=4, + node_rank=node_rank, + tp_size=4, + pp_size=pp_size, + dp_size=2, + enable_dp_attention=True, + attn_cp_size=2, + ): + assert node_hosts_rust_server() == expected[node_rank] + + +def test_scale_joiner_hosts_listener(): + with get_context().override_server_args( + nnodes=2, + node_rank=1, + tp_size=1, + dp_size=1, + enable_dp_attention=True, + ep_join_mode="scale", + ): + assert node_hosts_rust_server() + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__, "-v"]))