[rust-server] fix p/d bootstrap across dp listeners (#36234)
This commit is contained in:
@@ -74,6 +74,8 @@ class KVArgs:
|
||||
page_size: int
|
||||
# for system dp
|
||||
system_dp_rank: int
|
||||
# Local Rust /route registry port; None on scheduler ranks without a listener.
|
||||
rust_http_port: Optional[int]
|
||||
# for pp prefill
|
||||
pp_rank: int
|
||||
prefill_start_layer: int
|
||||
|
||||
@@ -828,6 +828,28 @@ class CommonKVManager(BaseKVManager):
|
||||
"prefill_http_port": get_serving().port,
|
||||
}
|
||||
|
||||
if envs.SGLANG_RUST_SERVER.get() and self.attn_dp_size > 1:
|
||||
topology_rows = get_world_group().all_gather_object(payload)
|
||||
# Every scheduler contributes a topology row. Only the scheduler
|
||||
# ranks that own a Rust listener populate their local registry.
|
||||
if self.kv_args.rust_http_port is None:
|
||||
return
|
||||
registry_host = {"0.0.0.0": "127.0.0.1", "::": "::1"}.get(
|
||||
self.bootstrap_host, self.bootstrap_host
|
||||
)
|
||||
registry = NetworkAddress(registry_host, self.kv_args.rust_http_port)
|
||||
url = f"{registry.to_url()}/route"
|
||||
for topology_row in topology_rows:
|
||||
registry_row = {
|
||||
**topology_row,
|
||||
"prefill_http_port": self.kv_args.rust_http_port,
|
||||
}
|
||||
self._register_topology_row(url, registry_row)
|
||||
return
|
||||
|
||||
self._register_topology_row(url, payload)
|
||||
|
||||
def _register_topology_row(self, url: str, payload: Dict) -> None:
|
||||
max_retries, initial_delay, max_delay = 5, 1.0, 30.0
|
||||
for attempt in range(max_retries):
|
||||
try:
|
||||
|
||||
@@ -208,6 +208,11 @@ class PrefillBootstrapQueue:
|
||||
kv_args.engine_rank = self.tp_rank
|
||||
kv_args.pp_rank = self.pp_rank
|
||||
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
||||
kv_args.rust_http_port = (
|
||||
self.scheduler.rust_server.http_port
|
||||
if self.scheduler.rust_server is not None
|
||||
else None
|
||||
)
|
||||
kv_args.kv_cache_dtype_str = (
|
||||
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
|
||||
)
|
||||
|
||||
@@ -55,10 +55,12 @@ class RustServer:
|
||||
def __init__(
|
||||
self,
|
||||
server: Server,
|
||||
http_port: int,
|
||||
mm_spec: Optional[RustMmSpec] = None,
|
||||
max_per_poll: int = 256,
|
||||
):
|
||||
self.server = server
|
||||
self.http_port = http_port
|
||||
self.mm_spec = mm_spec
|
||||
self._max_per_poll = max_per_poll
|
||||
|
||||
@@ -151,7 +153,7 @@ class RustServer:
|
||||
dp_note,
|
||||
)
|
||||
|
||||
return cls(server, mm_spec=mm_spec)
|
||||
return cls(server, http_port=listen_port, mm_spec=mm_spec)
|
||||
|
||||
def wait_request(self, timeout_ms: int) -> None:
|
||||
"""Block until a request is pushed into the in-process ring or the timeout
|
||||
|
||||
@@ -8,6 +8,7 @@ register_cpu_ci(est_time=11, suite="base-a-test-cpu")
|
||||
import unittest
|
||||
from unittest.mock import MagicMock, call, patch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_context
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -199,6 +200,86 @@ class TestRegisterToBootstrap(CustomTestCase):
|
||||
url_used = mock_put.call_args[0][0]
|
||||
self.assertIn("10.0.0.1", url_used)
|
||||
|
||||
@patch("sglang.srt.disaggregation.common.conn.requests.put")
|
||||
@patch("sglang.srt.disaggregation.common.conn.get_world_group")
|
||||
def test_rust_attention_dp_replicates_complete_topology_across_hosts(
|
||||
self, mock_world_group, mock_put
|
||||
):
|
||||
success_resp = MagicMock()
|
||||
success_resp.status_code = 200
|
||||
mock_put.return_value = success_resp
|
||||
|
||||
schedulers = (
|
||||
(0, 0, "10.0.0.1", 17000, 8765),
|
||||
(0, 1, "10.0.0.1", 17001, None),
|
||||
(1, 0, "10.0.0.2", 17002, 8766),
|
||||
(1, 1, "10.0.0.2", 17003, None),
|
||||
)
|
||||
|
||||
def gather_topology(payload):
|
||||
return [
|
||||
{
|
||||
**payload,
|
||||
"attn_dp_rank": dp_rank,
|
||||
"attn_tp_rank": tp_rank,
|
||||
"rank_ip": host,
|
||||
"rank_port": rank_port,
|
||||
}
|
||||
for dp_rank, tp_rank, host, rank_port, _ in schedulers
|
||||
]
|
||||
|
||||
mock_world_group.return_value.all_gather_object.side_effect = gather_topology
|
||||
|
||||
with envs.SGLANG_RUST_SERVER.override(True):
|
||||
for dp_rank, tp_rank, local_ip, _, rust_http_port in schedulers:
|
||||
manager = self._make_manager()
|
||||
manager.attn_dp_size = 2
|
||||
manager.attn_dp_rank = dp_rank
|
||||
manager.attn_tp_size = 2
|
||||
manager.attn_tp_rank = tp_rank
|
||||
manager.local_ip = local_ip
|
||||
manager.bootstrap_host = local_ip
|
||||
manager.kv_args.rust_http_port = rust_http_port
|
||||
manager.register_to_bootstrap()
|
||||
|
||||
topology_by_registry = {}
|
||||
for put_call in mock_put.call_args_list:
|
||||
payload = put_call.kwargs["json"]
|
||||
topology_by_registry.setdefault(put_call.args[0], {})[
|
||||
(payload["attn_dp_rank"], payload["attn_tp_rank"])
|
||||
] = (payload["rank_ip"], payload["rank_port"])
|
||||
complete_topology = {
|
||||
(dp, tp): (host, rank_port) for dp, tp, host, rank_port, _ in schedulers
|
||||
}
|
||||
self.assertEqual(
|
||||
topology_by_registry,
|
||||
{
|
||||
"http://10.0.0.1:8765/route": complete_topology,
|
||||
"http://10.0.0.2:8766/route": complete_topology,
|
||||
},
|
||||
)
|
||||
self.assertEqual(mock_put.call_count, 8)
|
||||
self.assertEqual(
|
||||
{
|
||||
(put_call.args[0], put_call.kwargs["json"]["prefill_http_port"])
|
||||
for put_call in mock_put.call_args_list
|
||||
},
|
||||
{
|
||||
("http://10.0.0.1:8765/route", 8765),
|
||||
("http://10.0.0.2:8766/route", 8766),
|
||||
},
|
||||
)
|
||||
self.assertEqual(
|
||||
[
|
||||
(
|
||||
gather_call.args[0]["attn_dp_rank"],
|
||||
gather_call.args[0]["attn_tp_rank"],
|
||||
)
|
||||
for gather_call in mock_world_group.return_value.all_gather_object.call_args_list
|
||||
],
|
||||
[(dp, tp) for dp, tp, _, _, _ in schedulers],
|
||||
)
|
||||
|
||||
@patch("sglang.srt.disaggregation.common.conn.time")
|
||||
@patch("sglang.srt.disaggregation.common.conn.requests.put")
|
||||
def test_wildcard_host_0000_uses_ipv4_loopback(self, mock_put, mock_time):
|
||||
@@ -258,6 +339,9 @@ class TestRegisterToBootstrap(CustomTestCase):
|
||||
mgr.register_to_bootstrap = CommonKVManager.register_to_bootstrap.__get__(
|
||||
mgr, CommonKVManager
|
||||
)
|
||||
mgr._register_topology_row = CommonKVManager._register_topology_row.__get__(
|
||||
mgr, CommonKVManager
|
||||
)
|
||||
|
||||
# Set attributes that register_to_bootstrap reads
|
||||
mgr.dist_init_addr = dist_init_addr
|
||||
@@ -278,6 +362,7 @@ class TestRegisterToBootstrap(CustomTestCase):
|
||||
|
||||
mgr.kv_args = MagicMock()
|
||||
mgr.kv_args.page_size = 16
|
||||
mgr.kv_args.rust_http_port = None
|
||||
# Resolved per-runner value threaded through KVArgs (the payload field).
|
||||
mgr.kv_cache_dtype_str = "auto"
|
||||
|
||||
|
||||
Reference in New Issue
Block a user