[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,4 +1,4 @@
from typing import Sequence
from typing import Optional, Sequence
import torch
import triton
@@ -15,10 +15,11 @@ def _copy_mla_rows_into_pack_kernel(
):
layer_id = tl.program_id(0)
block_id = tl.program_id(1)
metadata_offset = layer_id * 3
metadata_offset = layer_id * 4
src = tl.load(src_metadata + metadata_offset).to(pack.dtype)
row_nbytes = tl.load(src_metadata + metadata_offset + 1)
pack_offset = tl.load(src_metadata + metadata_offset + 2)
src_row_stride = tl.load(src_metadata + metadata_offset + 3)
offsets = block_id * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
layer_nbytes = num_rows * row_nbytes
@@ -26,7 +27,7 @@ def _copy_mla_rows_into_pack_kernel(
row = offsets // row_nbytes
byte = offsets % row_nbytes
src_row = tl.load(row_indices + row, mask=mask, other=0)
values = tl.load(src + src_row * row_nbytes + byte, mask=mask)
values = tl.load(src + src_row * src_row_stride + byte, mask=mask)
tl.store(pack + pack_offset + offsets, values, mask=mask)
@@ -35,11 +36,14 @@ def copy_mla_rows_into_pack(
row_indices: torch.Tensor,
pack: torch.Tensor,
token_item_lens: Sequence[int],
src_token_item_lens: Optional[Sequence[int]] = None,
) -> None:
if len(kv_data_ptrs) != len(token_item_lens):
if src_token_item_lens is None:
src_token_item_lens = token_item_lens
if not (len(kv_data_ptrs) == len(token_item_lens) == len(src_token_item_lens)):
raise ValueError(
"kv_data_ptrs and token_item_lens length mismatch: "
f"{len(kv_data_ptrs)} vs {len(token_item_lens)}"
"KV pointers, copy widths, and source strides length mismatch: "
f"{len(kv_data_ptrs)}, {len(token_item_lens)}, {len(src_token_item_lens)}"
)
if not kv_data_ptrs:
return
@@ -47,11 +51,13 @@ def copy_mla_rows_into_pack(
n = int(row_indices.numel())
metadata = []
offset = 0
for ptr, item_len in zip(kv_data_ptrs, token_item_lens):
for ptr, item_len, src_item_len in zip(
kv_data_ptrs, token_item_lens, src_token_item_lens
):
item_len = int(item_len)
if item_len <= 0:
raise ValueError(f"MLA token item length must be positive, got {item_len}")
metadata.extend((int(ptr), item_len, offset))
metadata.extend((int(ptr), item_len, offset, int(src_item_len)))
offset += n * item_len
src_metadata = torch.tensor(metadata, dtype=torch.int64, device=pack.device)
@@ -382,7 +382,9 @@ class CommonKVManager(BaseKVManager):
f"{type(self).__name__} does not support staging memory registration"
)
def _init_dcp_pack_buffers_once(self, dcp_size: int) -> None:
def _init_dcp_pack_buffers_once(
self, dcp_size: int, *, include_draft: bool = False
) -> None:
if self._dcp_pack_buffers is not None:
return
if not self.kv_args.kv_item_lens:
@@ -397,6 +399,7 @@ class CommonKVManager(BaseKVManager):
len(self.transfer_queues),
dcp_size,
max_tokens,
include_draft=include_draft,
)
self._dcp_pack_max_tokens = max_tokens
@@ -41,6 +41,7 @@ def try_pack_dcp_src(
kv_data_ptrs: Sequence[int],
src_token_indices: npt.NDArray[np.integer],
token_item_lens: Sequence[int],
src_token_item_lens: Optional[Sequence[int]] = None,
pack_offset_bytes: int = 0,
pack_capacity_bytes: Optional[int] = None,
) -> Optional[Tuple[List[int], npt.NDArray[np.int64]]]:
@@ -75,7 +76,9 @@ def try_pack_dcp_src(
gather_stream = pack_buffer.get_gather_stream()
gather_stream.wait_stream(torch.cuda.default_stream(pack.device))
with torch.cuda.stream(gather_stream):
copy_mla_rows_into_pack(kv_data_ptrs, row_indices, pack, token_item_lens)
copy_mla_rows_into_pack(
kv_data_ptrs, row_indices, pack, token_item_lens, src_token_item_lens
)
gather_stream.synchronize()
packed_ptrs: List[int] = []
@@ -93,14 +96,16 @@ def init_dcp_pack_buffers(
count: int,
dcp_size: int,
max_tokens: int,
*,
include_draft: bool = False,
) -> List[StagingBuffer]:
from sglang.srt.disaggregation.common.staging_handler import (
_get_custom_mem_pool,
)
kv_item_lens = kv_args.kv_item_lens
if kv_args.num_draft_entries > 0:
kv_item_lens = kv_item_lens[: len(kv_item_lens) - kv_args.num_draft_entries]
if not include_draft and kv_args.num_draft_entries:
kv_item_lens = kv_item_lens[: -kv_args.num_draft_entries]
# Note(kpham-sgl): size = dcp_size x ceil(max_tokens / dcp_size)
# x sum(per-layer token bytes). At 32,768 tokens and 61 MLA layers
# x 576 bf16 dims x 2 B: 2.14 GiB/buffer, 8.58 GiB for 4 queues.
@@ -1173,6 +1173,7 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
for entry in range(num_target)
]
sliced_draft_params = []
draft_to_pack = []
if num_draft > 0 and plan.draft_src_token_indices.size:
if not dst_kv_item_lens and dst_attn_tp_size not in (
None,
@@ -1224,14 +1225,45 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
dst_rank = dst_tp_rank // max(1, dst_span // src_span)
src_offset = (dst_rank * dst_width) % src_width
dst_offset = (src_rank * src_width) % dst_width
sliced_draft_params.append(
(
src_kv_ptrs[entry] + src_offset,
dst_kv_ptrs[entry] + dst_offset,
src_width,
dst_width,
copy_width,
)
params = (
src_kv_ptrs[entry] + src_offset,
dst_kv_ptrs[entry] + dst_offset,
src_width,
dst_width,
copy_width,
)
if pack_buffer is not None and src_width > dst_width:
draft_to_pack.append(params)
else:
sliced_draft_params.append(params)
if draft_to_pack:
from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
draft_src_ptrs, draft_dst_ptrs, src_strides, _, copy_widths = zip(
*draft_to_pack
)
target_pack_bytes = plan.target_src_token_indices.size * sum(
dcp_token_item_lens[:num_target]
)
packed = try_pack_dcp_src(
pack_buffer=pack_buffer,
kv_data_ptrs=draft_src_ptrs,
src_token_indices=plan.draft_src_token_indices,
token_item_lens=copy_widths,
src_token_item_lens=src_strides,
pack_offset_bytes=target_pack_bytes,
)
if packed is None:
sliced_draft_params.extend(draft_to_pack)
else:
packed_ptrs, packed_indices = packed
packed_groups = group_concurrent_contiguous(
packed_indices, plan.draft_dst_token_indices
)
layers_params.extend(
(src, dst, width, packed_groups)
for src, dst, width in zip(packed_ptrs, draft_dst_ptrs, copy_widths)
)
def process_sliced_draft(params) -> int:
@@ -1284,7 +1316,11 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
executor.submit(process_sliced_draft, [params])
for params in sliced_draft_params
)
return self._await_transfer_futures(futures)
try:
return self._await_transfer_futures(futures)
finally:
if pack_buffer is not None:
concurrent.futures.wait(futures)
transfer_blocks = []
for layer_params in layers_params:
@@ -2492,7 +2528,9 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
decode_kv_args.dst_dcp_size,
)
)
self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size)
self._init_dcp_pack_buffers_once(
decode_kv_args.dst_dcp_size, include_draft=True
)
self.decode_kv_args_table[mooncake_session_id] = decode_kv_args
with self.session_lock:
if mooncake_session_id in self.failed_sessions:
@@ -1,8 +1,13 @@
import concurrent.futures
import unittest
from types import SimpleNamespace
import numpy as np
import torch
from sglang.kernels.ops.kvcache.pd_dcp_gather import copy_mla_rows_into_pack
from sglang.srt.disaggregation.common.staging_buffer import StagingBuffer
from sglang.srt.disaggregation.mooncake.conn import MooncakeKVManager
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -36,6 +41,135 @@ class TestPdDcpGather(CustomTestCase):
torch.testing.assert_close(packed0, kv0[row_indices], rtol=0, atol=0)
torch.testing.assert_close(packed1, kv1[row_indices], rtol=0, atol=0)
def test_packed_tp2_pp2_to_dcp4_preserves_kv(self):
"""Packing must preserve both target rows and draft head shards across PP stages."""
for custom_pool in (False, True):
for capacity in (256 * (2 * 64 + 2 * 512), 256 * 2 * 64 // 4):
for rank in range(4):
with self.subTest(
custom_pool=custom_pool, capacity=capacity, rank=rank
):
self._check_packed_transfer(rank, custom_pool, capacity)
def _check_packed_transfer(self, rank, custom_pool, capacity):
page, tokens, chunk = 64, 521, 256
src_pages = np.array([7, 1, 9, 3, 4, 11, 2, 5, 8], dtype=np.int32)
dst_pages = np.array([4, 1, 6], dtype=np.int32)
layers, widths = [3, 11, 19, 27, 28, 28], [64] * 4 + [256] * 2
logical = torch.arange(tokens, device="cuda")
src_rows = (
torch.as_tensor(src_pages, device="cuda")[logical // page] * page
+ logical % page
)
values = [
(
(logical[:, None] + 256) * 13
+ torch.arange(width, device="cuda") * 7
+ entry * 31
)
.remainder(251)
.to(torch.uint8)
for entry, width in enumerate([64] * 4 + [1024] * 2)
]
destinations = [
torch.full((2048, w), 165, dtype=torch.uint8, device="cuda") for w in widths
]
expected = [x.clone() for x in destinations]
owned = logical[rank::4]
target_rows = (
torch.as_tensor(dst_pages, device="cuda")[owned // 256] * page
+ owned % 256 // 4
)
draft_rows = (
torch.as_tensor(dst_pages, device="cuda")[logical // 256] * 256
+ logical % 256
)
for entry in range(4):
expected[entry][target_rows] = values[entry][owned]
for entry in (4, 5):
expected[entry][draft_rows] = values[entry][
:, rank * 256 : (rank + 1) * 256
]
pack = StagingBuffer(capacity, "cuda:0", 0)
with concurrent.futures.ThreadPoolExecutor(max_workers=3) as executor:
for stage, entries in enumerate(([0, 1], [2, 3, 4, 5])):
sources = []
for entry in entries:
data = values[entry]
if entry >= 4:
start = (rank // 2) * 512
data = data[:, start : start + 512]
source = torch.full(
(1024, data.shape[1]), 165, dtype=torch.uint8, device="cuda"
)
source[src_rows] = data
sources.append(source)
buffers = sources + destinations + [pack.buffer]
def transfer(session, blocks):
def view(ptr, size):
for tensor in buffers:
offset = ptr - tensor.data_ptr()
if 0 <= offset and offset + size <= tensor.numel():
return tensor.flatten()[offset : offset + size]
raise AssertionError(
f"Transfer outside registered buffers: {ptr}, {size}"
)
for src, dst, size in blocks:
view(dst, size).copy_(view(src, size))
torch.cuda.synchronize()
return 0
manager = SimpleNamespace(
is_mla_backend=False,
kv_args=SimpleNamespace(
page_size=page,
kv_layer_ids=[layers[e] for e in entries],
kv_data_ptrs=[x.data_ptr() for x in sources],
num_draft_entries=2 if stage else 0,
engine_rank=stage * 2 + rank // 2,
),
attn_tp_size=2,
max_transfer_batch_indices=37,
enable_custom_mem_pool=custom_pool,
enable_deferred_decode_kv_release=False,
_transfer_data=transfer,
)
manager._await_transfer_futures = lambda futures: (
MooncakeKVManager._await_transfer_futures(manager, futures)
)
for start in range(0, tokens, chunk):
count = min(chunk, tokens - start)
result = MooncakeKVManager.send_kvcache_dcp(
manager,
"session",
src_pages[start // page : (start + count + page - 1) // page],
[x.data_ptr() for x in destinations],
dst_pages,
dcp_token_item_lens=[x.shape[1] for x in sources],
dst_dcp_size=4,
dst_dcp_rank=rank,
src_page_offset=start // page,
decode_prefix_len=256,
num_kv_tokens=count,
executor=executor,
dst_layer_ids=layers,
pack_buffer=pack,
dst_kv_item_lens=[
page * w * (4 if e >= 4 else 1)
for e, w in enumerate(widths)
],
dst_tp_rank=rank,
dst_attn_tp_size=4,
)
self.assertEqual(result, 0)
for entry in entries:
torch.testing.assert_close(
destinations[entry], expected[entry], rtol=0, atol=0
)
if __name__ == "__main__":
unittest.main()
@@ -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()