[PD] Transfer the DCP-replicated DSPARK draft KV in DCP1->DCP-N relayouts (#37709)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Khoa Pham
2026-09-11 17:15:18 -07:00
committed by GitHub
co-authored by Claude Fable 5 Cursor
parent 7d9c57da6e
commit a207786205
11 changed files with 549 additions and 146 deletions
@@ -1,10 +1,12 @@
import unittest
from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import Mock, patch
import numpy as np
import torch
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.common.dcp_pack import (
dcp_pack_buffer_bytes,
try_pack_dcp_src,
@@ -19,32 +21,170 @@ from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=12, suite="base-a-test-cpu")
class TestPackedDcpGrouping(CustomTestCase):
def test_packed_groups_collapse_cyclic_src(self):
page_size = 64
dcp_size = 4
src_pages = np.arange(4, dtype=np.int32)
dst_pages = np.array([7], dtype=np.int32)
plan = build_dcp_token_transfer_plan(
src_pages,
dst_pages,
physical_page_size=page_size,
dcp_size=dcp_size,
dcp_rank=0,
num_kv_tokens=256,
)
raw_src, _ = group_concurrent_contiguous(
plan.src_token_indices, plan.dst_token_indices
)
self.assertEqual(len(raw_src), 64)
self.assertTrue(all(len(group) == 1 for group in raw_src))
def _plan(*, src, dst, page_size, dcp_size, dcp_rank, **kwargs):
return build_dcp_token_transfer_plan(
np.asarray(src, dtype=np.int32),
np.asarray(dst, dtype=np.int32),
physical_page_size=page_size,
dcp_size=dcp_size,
dcp_rank=dcp_rank,
**kwargs,
)
packed_src = np.arange(plan.dst_token_indices.size, dtype=np.int64)
packed_groups, _ = group_concurrent_contiguous(
packed_src, plan.dst_token_indices
class TestDcpTokenTransferPlan(CustomTestCase):
def test_one_virtual_page_explicit_rows(self):
# P=2, N=4. Prefill pages 5,2,11,4; decode virtual page 7.
# pos 0..7 src rows: 10,11, 4,5, 22,23, 8,9
# draft dest page is P*N=8 → 56..63
# each rank stores local rows 14,15 (page P=2)
expected_draft_src = [10, 11, 4, 5, 22, 23, 8, 9]
expected_draft_dst = list(range(56, 64))
expected_target_src = {
0: [10, 22],
1: [11, 23],
2: [4, 8],
3: [5, 9],
}
seen_src = []
for rank, src in expected_target_src.items():
plan = _plan(
src=[5, 2, 11, 4],
dst=[7],
page_size=2,
dcp_size=4,
dcp_rank=rank,
num_kv_tokens=8,
)
np.testing.assert_array_equal(
plan.draft_src_token_indices, expected_draft_src
)
np.testing.assert_array_equal(
plan.draft_dst_token_indices, expected_draft_dst
)
np.testing.assert_array_equal(plan.target_src_token_indices, src)
np.testing.assert_array_equal(plan.target_dst_token_indices, [14, 15])
seen_src.extend(plan.target_src_token_indices.tolist())
self.assertEqual(sorted(seen_src), sorted(expected_draft_src))
def test_second_chunk_crosses_dest_pages(self):
# P=2, N=2 (virtual page = 4). Decode already holds a 4-token prefix;
# dst=[4, 6] is the full send-range page list. This chunk is the second
# prefill page of the send range (src_page_offset=1), so its 4 tokens
# sit at send-range pos 2..5 (absolute 6..9) and straddle virtual page
# 4 (rows 16..19) and virtual page 6 (rows 24..27).
plan = _plan(
src=[9, 3],
dst=[4, 6],
page_size=2,
dcp_size=2,
dcp_rank=0,
src_page_offset=1,
decode_prefix_len=4,
num_kv_tokens=4,
)
self.assertEqual(len(packed_groups), 1)
self.assertEqual(len(packed_groups[0]), 64)
np.testing.assert_array_equal(plan.draft_src_token_indices, [18, 19, 6, 7])
np.testing.assert_array_equal(plan.draft_dst_token_indices, [18, 19, 24, 25])
# rank 0 owns absolute pos 6, 8 -> per-rank slots 1, 2 -> pages 4, 6.
np.testing.assert_array_equal(plan.target_src_token_indices, [18, 6])
np.testing.assert_array_equal(plan.target_dst_token_indices, [9, 12])
plan_r1 = _plan(
src=[9, 3],
dst=[4, 6],
page_size=2,
dcp_size=2,
dcp_rank=1,
src_page_offset=1,
decode_prefix_len=4,
num_kv_tokens=4,
)
np.testing.assert_array_equal(plan_r1.draft_src_token_indices, [18, 19, 6, 7])
np.testing.assert_array_equal(plan_r1.draft_dst_token_indices, [18, 19, 24, 25])
np.testing.assert_array_equal(plan_r1.target_src_token_indices, [19, 7])
np.testing.assert_array_equal(plan_r1.target_dst_token_indices, [9, 12])
def test_rejects_unaligned_prefix(self):
with self.assertRaisesRegex(ValueError, "align"):
_plan(
src=[0],
dst=[0],
page_size=2,
dcp_size=4,
dcp_rank=0,
decode_prefix_len=1,
num_kv_tokens=2,
)
def test_empty_tokens(self):
plan = _plan(
src=[0], dst=[0], page_size=2, dcp_size=4, dcp_rank=0, num_kv_tokens=0
)
self.assertTrue(plan.empty())
class TestPackedDcpGrouping(CustomTestCase):
def test_target_needs_pack_draft_does_not(self):
plan = _plan(
src=[0, 1, 2, 3],
dst=[0],
page_size=2,
dcp_size=4,
dcp_rank=0,
num_kv_tokens=8,
)
np.testing.assert_array_equal(plan.target_src_token_indices, [0, 4])
np.testing.assert_array_equal(plan.target_dst_token_indices, [0, 1])
target_src, _ = group_concurrent_contiguous(
plan.target_src_token_indices, plan.target_dst_token_indices
)
self.assertEqual(target_src, [[0], [4]])
packed_src, packed_dst = group_concurrent_contiguous(
np.arange(2, dtype=np.int64), plan.target_dst_token_indices
)
self.assertEqual(packed_src, [[0, 1]])
self.assertEqual(packed_dst, [[0, 1]])
draft_src, draft_dst = group_concurrent_contiguous(
plan.draft_src_token_indices, plan.draft_dst_token_indices
)
self.assertEqual(draft_src, [[0, 1, 2, 3, 4, 5, 6, 7]])
self.assertEqual(draft_dst, [[0, 1, 2, 3, 4, 5, 6, 7]])
def _dcp_kv_manager_stub(*, page_size, kv_item_lens, num_draft_entries):
return SimpleNamespace(
kv_args=SimpleNamespace(
page_size=page_size,
kv_item_lens=kv_item_lens,
num_draft_entries=num_draft_entries,
)
)
class TestPrepareDcpTokenItemLens(CustomTestCase):
def test_draft_tail_scales_by_dst_dcp_size(self):
mgr = _dcp_kv_manager_stub(
page_size=64,
kv_item_lens=[64 * 32, 64 * 32, 64 * 16],
num_draft_entries=1,
)
token_lens = CommonKVManager.prepare_dcp_token_item_lens(
mgr, [64 * 32, 64 * 32, 4 * 64 * 16], dst_dcp_size=4
)
self.assertEqual(token_lens, [32, 32, 16])
def test_rejects_unscaled_draft_item_len(self):
mgr = _dcp_kv_manager_stub(
page_size=64,
kv_item_lens=[64 * 32, 64 * 16],
num_draft_entries=1,
)
with self.assertRaisesRegex(RuntimeError, "geometry differs at entry 1"):
CommonKVManager.prepare_dcp_token_item_lens(
mgr, [64 * 32, 64 * 16], dst_dcp_size=4
)
class TestDcpPackBufferBytes(CustomTestCase):
@@ -575,7 +575,9 @@ class TestNixlTransferWorker(CustomTestCase):
mgr.is_hybrid_mla_backend = False
mgr.attn_tp_size = 1
mgr.transfer_source_rank = 0
mgr.kv_args = SimpleNamespace(engine_rank=0, kv_data_ptrs=[0])
mgr.kv_args = SimpleNamespace(
engine_rank=0, kv_data_ptrs=[0], num_draft_entries=0
)
mgr.exceptions = {}
mgr.failure_lock = threading.Lock()
mgr.failure_records = {}
@@ -674,6 +676,7 @@ class TestNixlTransferWorker(CustomTestCase):
engine_rank=0,
kv_data_ptrs=[0x1000],
page_size=4,
num_draft_entries=0,
)
mgr._dcp_pack_buffers = [SimpleNamespace(get_size=lambda: 16)]
@@ -686,7 +689,8 @@ class TestNixlTransferWorker(CustomTestCase):
def send_kvcache_dcp(*args, **kwargs):
submitted.append((args[0], args[-1]))
return f"handle-{args[0]}"
# One handle per transfer part; the worker extends its handle list.
return [f"handle-{args[0]}"]
mgr.send_kvcache_dcp = MagicMock(side_effect=send_kvcache_dcp)
submitted_counts_at_poll = []