[PD] Pack draft KV head slices for DCP transfers (#40500)

Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
This commit is contained in:
Khoa Pham
2026-09-21 21:11:27 -07:00
committed by GitHub
co-authored by Qiaolin Yu
parent b44e248682
commit 018b73c7a0
6 changed files with 273 additions and 23 deletions
@@ -1,12 +1,14 @@
import concurrent.futures
import unittest
from threading import Event
from types import SimpleNamespace
from unittest.mock import MagicMock, call
from unittest.mock import MagicMock, call, patch
import numpy as np
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
@@ -286,5 +288,67 @@ class TestDcpDraftHeadTransfer(unittest.TestCase):
self.assertFalse(dst_buffers[1000000].any())
class TestDcpPackLifetime(CustomTestCase):
def test_failed_transfer_drains_before_pack_buffer_reuse(self):
"""A failed layer must not release the pack buffer while another transfer reads it."""
manager = TestMooncakeTransferBatching._make_manager(
enable_custom_mem_pool=True
)
manager.kv_args = SimpleNamespace(
page_size=1, kv_layer_ids=[], kv_data_ptrs=[1000, 2000], num_draft_entries=0
)
source = np.array([11], dtype=np.uint8)
observed = []
running, release = Event(), Event()
def transfer(session, blocks):
if blocks[0][0] == 1000:
self.assertTrue(running.wait(10))
return 17
running.set()
self.assertTrue(release.wait(10))
observed.append(int(source[0]))
return 0
def send(executor):
result = MooncakeKVManager.send_kvcache_dcp(
manager,
"session",
np.array([0, 1], dtype=np.int32),
[5000, 6000],
np.array([0], dtype=np.int32),
dcp_token_item_lens=[1, 1],
dst_dcp_size=2,
dst_dcp_rank=0,
src_page_offset=0,
decode_prefix_len=0,
num_kv_tokens=2,
executor=executor,
dst_layer_ids=[],
pack_buffer=object(),
)
source[0] = 22
return result
manager._transfer_data = transfer
with (
patch(
"sglang.srt.disaggregation.common.dcp_pack.try_pack_dcp_src",
return_value=([1000, 2000], np.array([0], dtype=np.int64)),
),
concurrent.futures.ThreadPoolExecutor(max_workers=2) as transfers,
concurrent.futures.ThreadPoolExecutor(max_workers=1) as worker,
):
future = worker.submit(send, transfers)
try:
self.assertTrue(running.wait(10))
with self.assertRaises(concurrent.futures.TimeoutError):
future.result(timeout=1)
finally:
release.set()
self.assertEqual(future.result(timeout=10), 17)
self.assertEqual(observed, [11])
if __name__ == "__main__":
unittest.main()