[rust-server] Use node-local HTTP ports for DP attention (#34430)
Signed-off-by: Sage Ahrac <sagiahrak@gmail.com>
This commit is contained in:
@@ -1169,11 +1169,16 @@ class Engine(EngineScoreMixin, EngineBase):
|
|||||||
weight_cache_daemon_procs,
|
weight_cache_daemon_procs,
|
||||||
)
|
)
|
||||||
|
|
||||||
launch_dummy_health_check_server(
|
# A node-local Rust listener owns the health endpoints when present.
|
||||||
get_serving().host,
|
rust_server_owns_base_port = (
|
||||||
get_serving().port,
|
envs.SGLANG_RUST_SERVER.get() and node_hosts_rust_server()
|
||||||
get_observability().enable_metrics,
|
|
||||||
)
|
)
|
||||||
|
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()
|
scheduler_init_result.block_until_scheduler_exits()
|
||||||
return (
|
return (
|
||||||
@@ -1870,6 +1875,31 @@ def _calculate_rank_ranges(
|
|||||||
return pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node
|
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]:
|
def _compute_parallelism_ranks(tp_rank: int) -> Tuple[int, int, int]:
|
||||||
"""Compute attention-CP, MoE-DP, and MoE-EP ranks for a TP rank.
|
"""Compute attention-CP, MoE-DP, and MoE-EP ranks for a TP rank.
|
||||||
|
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from sglang.srt.managers.utils import (
|
|||||||
MsgpackDecodeError,
|
MsgpackDecodeError,
|
||||||
msgpack_decode_explained,
|
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.config import _build_server_args, _partition_cores
|
||||||
from sglang.srt.rust_server.multimodal import (
|
from sglang.srt.rust_server.multimodal import (
|
||||||
RUST_MM_FAMILIES,
|
RUST_MM_FAMILIES,
|
||||||
@@ -79,10 +79,19 @@ class RustServer:
|
|||||||
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
||||||
|
|
||||||
server_args = scheduler.server_args
|
server_args = scheduler.server_args
|
||||||
# Per-DP-rank HTTP port with client load balancing. `None` when DP is off,
|
# Preserve the DP startup log; ports use node-local offsets.
|
||||||
# so the rank is not conflated with rank 0 of a one-rank group.
|
|
||||||
dp_rank = scheduler.ps.attn_dp_rank if scheduler.ps.dp_size > 1 else None
|
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()
|
listen_addr = NetworkAddress(get_serving().host, listen_port).to_host_port_str()
|
||||||
|
|
||||||
launch_cores, server_cores = _partition_cores(
|
launch_cores, server_cores = _partition_cores(
|
||||||
@@ -97,7 +106,7 @@ class RustServer:
|
|||||||
_build_server_args(scheduler),
|
_build_server_args(scheduler),
|
||||||
# None -> run unpinned; the list carries the pinning decision.
|
# None -> run unpinned; the list carries the pinning decision.
|
||||||
cores=server_cores,
|
cores=server_cores,
|
||||||
port_offset=dp_rank,
|
port_offset=local_dp_rank,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Multimodal models must have a Rust pipeline — there is no Python
|
# Multimodal models must have a Rust pipeline — there is no Python
|
||||||
|
|||||||
@@ -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"]))
|
||||||
Reference in New Issue
Block a user