[disagg] Fix KV-event publisher port collision under pure data parallelism (#29211)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Kangyan-Zhou
2026-07-01 10:57:34 -07:00
committed by GitHub
co-authored by Claude Opus 4.8
parent 03b9278da0
commit 30c9801b39
3 changed files with 125 additions and 1 deletions
@@ -36,6 +36,28 @@ from pydantic import BaseModel
logger = logging.getLogger(__name__) 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( class EventBatch(
msgspec.Struct, msgspec.Struct,
array_like=True, # type: ignore[call-arg] array_like=True, # type: ignore[call-arg]
@@ -15,6 +15,7 @@ import zmq
from sglang.srt.disaggregation.kv_events import ( from sglang.srt.disaggregation.kv_events import (
EventPublisherFactory, EventPublisherFactory,
KVEventBatch, KVEventBatch,
select_kv_publisher_dp_rank,
) )
from sglang.srt.managers.io_struct import hook_custom_types, sock_send from sglang.srt.managers.io_struct import hook_custom_types, sock_send
@@ -69,7 +70,10 @@ class SchedulerKvEventsPublisher:
if self.enable_kv_cache_events: if self.enable_kv_cache_events:
self.kv_event_publisher = EventPublisherFactory.create( 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): def emit_kv_metrics(self):
@@ -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()