[PD] Pack draft KV head slices for DCP transfers (#40500)
Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user