[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: