[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:
co-authored by
Claude Opus 4.8
parent
03b9278da0
commit
30c9801b39
@@ -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()
|
||||||
Reference in New Issue
Block a user