From 30c9801b399cee2c58098195961a7221f7a5012f Mon Sep 17 00:00:00 2001 From: Kangyan-Zhou Date: Wed, 1 Jul 2026 10:57:34 -0700 Subject: [PATCH] [disagg] Fix KV-event publisher port collision under pure data parallelism (#29211) Co-authored-by: Claude Opus 4.8 (1M context) --- python/sglang/srt/disaggregation/kv_events.py | 22 +++++ .../kv_events_publisher.py | 6 +- .../unit/disaggregation/test_kv_events.py | 98 +++++++++++++++++++ 3 files changed, 125 insertions(+), 1 deletion(-) create mode 100644 test/registered/unit/disaggregation/test_kv_events.py diff --git a/python/sglang/srt/disaggregation/kv_events.py b/python/sglang/srt/disaggregation/kv_events.py index 91d8b5ee9..0a11009e0 100644 --- a/python/sglang/srt/disaggregation/kv_events.py +++ b/python/sglang/srt/disaggregation/kv_events.py @@ -36,6 +36,28 @@ from pydantic import BaseModel logger = logging.getLogger(__name__) +def select_kv_publisher_dp_rank( + attn_dp_size: int, attn_dp_rank: int, dp_rank: Optional[int] +) -> int: + """Index used to offset this scheduler's KV-event publisher port. + + Each independent KV cache must publish on its own port so a consumer can + subscribe per replica. There are always ``dp_size`` such publishers; which + rank distinguishes them depends on the parallelism mode: + + - DP-attention (``attn_dp_size > 1``): each attention-DP rank owns a KV + cache shard, so distinguish by ``attn_dp_rank``. + - Pure DP (``attn_dp_size == 1``): every worker has ``attn_dp_rank == 0``, + so distinguish by ``dp_rank`` (the data-parallel replica index). + + Both span ``0..dp_size-1``, matching the ``dp_size`` advertised in + ``/server_info`` and the per-rank ports the router subscribes to. + """ + if attn_dp_size > 1: + return attn_dp_rank + return dp_rank or 0 + + class EventBatch( msgspec.Struct, array_like=True, # type: ignore[call-arg] diff --git a/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py b/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py index 9cd95e565..5aee67f7b 100644 --- a/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py +++ b/python/sglang/srt/managers/scheduler_components/kv_events_publisher.py @@ -15,6 +15,7 @@ import zmq from sglang.srt.disaggregation.kv_events import ( EventPublisherFactory, KVEventBatch, + select_kv_publisher_dp_rank, ) from sglang.srt.managers.io_struct import hook_custom_types, sock_send @@ -69,7 +70,10 @@ class SchedulerKvEventsPublisher: if self.enable_kv_cache_events: self.kv_event_publisher = EventPublisherFactory.create( - kv_events_config, self.ps.attn_dp_rank + kv_events_config, + select_kv_publisher_dp_rank( + self.ps.attn_dp_size, self.ps.attn_dp_rank, self.ps.dp_rank + ), ) def emit_kv_metrics(self): diff --git a/test/registered/unit/disaggregation/test_kv_events.py b/test/registered/unit/disaggregation/test_kv_events.py new file mode 100644 index 000000000..d935eb8cf --- /dev/null +++ b/test/registered/unit/disaggregation/test_kv_events.py @@ -0,0 +1,98 @@ +"""Unit tests for srt/disaggregation/kv_events KV-event publisher rank selection. + +Covers the data-parallel rank used to offset each scheduler's KV-event +publisher port, across pure DP, DP-attention, and single-replica modes. The +port offset must make every independent KV cache publish on a distinct port so +the router can subscribe per replica (the `dp_size` it reads from +`/server_info`). +""" + +import unittest + +from sglang.srt.disaggregation.kv_events import ( + ZmqEventPublisher, + select_kv_publisher_dp_rank, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +class TestSelectKvPublisherDpRank(CustomTestCase): + def test_select_rank_across_modes(self): + # (label, attn_dp_size, attn_dp_rank, dp_rank, expected) + cases = [ + # Pure DP (no dp-attention): attn_dp_rank is 0 for every worker, + # so the replica is distinguished by dp_rank. + ("pure_dp_worker0", 1, 0, 0, 0), + ("pure_dp_worker1", 1, 0, 1, 1), + ("pure_dp_worker3", 1, 0, 3, 3), + # DP-attention: each attn-dp rank owns a KV shard; distinguish by + # attn_dp_rank. dp_rank is ignored entirely in this mode. + ("dp_attention_rank0", 2, 0, None, 0), + ("dp_attention_rank1", 2, 1, None, 1), + ("dp_attention_ignores_dp_rank", 2, 1, 99, 1), + # Single replica / no DP. + ("single_dp_rank_none", 1, 0, None, 0), + ("single_dp_rank_zero", 1, 0, 0, 0), + ] + for label, attn_dp_size, attn_dp_rank, dp_rank, expected in cases: + with self.subTest(label): + self.assertEqual( + select_kv_publisher_dp_rank(attn_dp_size, attn_dp_rank, dp_rank), + expected, + ) + + def test_workers_bind_sequential_ports_per_replica(self): + # Each replica r must publish on port_base + r, since the router opens + # one SUB socket per rank at port_base + r. Regression: pre-fix every + # pure-DP worker offset by attn_dp_rank == 0, so all collapsed onto the + # single port tcp://*:5557 -> the 2nd worker crashed binding an + # already-bound port. + endpoint = "tcp://*:5557" + expected = [f"tcp://*:{5557 + r}" for r in range(4)] + + # Pure DP: replica index is dp_rank (attn_dp_rank is 0 for all). + pure_dp = [ + ZmqEventPublisher.offset_endpoint_port( + endpoint, select_kv_publisher_dp_rank(1, 0, r) + ) + for r in range(4) + ] + self.assertEqual(pure_dp, expected) + + # DP-attention: replica index is attn_dp_rank. + dp_attention = [ + ZmqEventPublisher.offset_endpoint_port( + endpoint, select_kv_publisher_dp_rank(4, a, None) + ) + for a in range(4) + ] + self.assertEqual(dp_attention, expected) + + def test_publisher_rank_count_matches_advertised_dp_size(self): + # The router subscribes to `dp_size` per-rank ports (from /server_info). + # The engine must produce exactly `dp_size` distinct publisher ranks in + # both modes, otherwise some subscribed ports get no data. + for dp_size in (1, 2, 4): + with self.subTest(f"pure_dp_{dp_size}"): + ranks = { + select_kv_publisher_dp_rank( + attn_dp_size=1, attn_dp_rank=0, dp_rank=r + ) + for r in range(dp_size) + } + self.assertEqual(len(ranks), dp_size) + with self.subTest(f"dp_attention_{dp_size}"): + ranks = { + select_kv_publisher_dp_rank( + attn_dp_size=dp_size, attn_dp_rank=a, dp_rank=None + ) + for a in range(dp_size) + } + self.assertEqual(len(ranks), dp_size) + + +if __name__ == "__main__": + unittest.main()