[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
@@ -95,6 +95,7 @@ class KVArgs:
hidden_kv_layers: int
# Only used of npu, for decode total kv layers
draft_kv_layers: int
num_draft_entries: int = 0
class KVPoll:
@@ -340,17 +340,34 @@ class CommonKVManager(BaseKVManager):
f"Unsupported PD DCP topology: {self.dcp_size} -> {dst_dcp_size}"
)
def prepare_dcp_token_item_lens(self, dst_page_item_lens: List[int]) -> List[int]:
def prepare_dcp_token_item_lens(
self, dst_page_item_lens: List[Optional[int]], dst_dcp_size: int
) -> List[int]:
page_size = self.kv_args.page_size
num_draft = self.kv_args.num_draft_entries
num_entries = len(self.kv_args.kv_item_lens)
if len(dst_page_item_lens) != num_entries:
raise RuntimeError(
"PD DCP requires the decode to register one KV entry per "
f"prefill entry: src={num_entries} (draft={num_draft}), "
f"dst={len(dst_page_item_lens)}"
)
src_token_lens = [
item_len // page_size for item_len in self.kv_args.kv_item_lens
]
dst_token_lens = [item_len // page_size for item_len in dst_page_item_lens]
if src_token_lens != dst_token_lens:
raise RuntimeError(
"PD DCP source/destination KV geometry differs: "
f"src={src_token_lens}, dst={dst_token_lens}"
for i, dst_item_len in enumerate(dst_page_item_lens):
if dst_item_len is None:
continue
dst_page_scale = page_size * (
dst_dcp_size if i >= num_entries - num_draft else 1
)
if dst_item_len // dst_page_scale != src_token_lens[i]:
raise RuntimeError(
"PD DCP source/destination KV geometry differs at entry "
f"{i}: src token bytes={src_token_lens[i]}, "
f"dst token bytes={dst_item_len // dst_page_scale} "
f"(dst item_len={dst_item_len}, page scale={dst_page_scale})"
)
return src_token_lens
def _register_staging_memory(self, ptr: int, size: int) -> None:
@@ -96,11 +96,14 @@ def init_dcp_pack_buffers(
max_tokens = max_prefill_buffer_tokens()
if max_tokens <= 0:
max_tokens = get_schedule().max_prefill_tokens
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]
# 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.
size_bytes = dcp_pack_buffer_bytes(
kv_args.kv_item_lens, kv_args.page_size, max_tokens, dcp_size
kv_item_lens, kv_args.page_size, max_tokens, dcp_size
)
gpu_id = kv_args.gpu_id
device = f"cuda:{gpu_id}"
@@ -137,8 +137,16 @@ def group_concurrent_contiguous(
@dataclasses.dataclass(frozen=True)
class DCPTokenTransferPlan:
src_token_indices: npt.NDArray[np.int64]
dst_token_indices: npt.NDArray[np.int64]
target_src_token_indices: npt.NDArray[np.int64]
target_dst_token_indices: npt.NDArray[np.int64]
draft_src_token_indices: npt.NDArray[np.int64]
draft_dst_token_indices: npt.NDArray[np.int64]
def empty(self) -> bool:
return (
self.target_src_token_indices.size == 0
and self.draft_src_token_indices.size == 0
)
def build_dcp_token_transfer_plan(
@@ -152,52 +160,38 @@ def build_dcp_token_transfer_plan(
decode_prefix_len: int = 0,
num_kv_tokens: Optional[int] = None,
) -> DCPTokenTransferPlan:
src_pages = np.asarray(src_page_indices, dtype=np.int64)
dst_pages = np.asarray(dst_page_indices, dtype=np.int64)
virtual_page_size = physical_page_size * dcp_size
if decode_prefix_len % virtual_page_size != 0:
raise ValueError(
"PD DCP transfer requires decode_prefix_len to align to the virtual "
f"DCP page size ({virtual_page_size}), got {decode_prefix_len}"
)
src_pages = np.asarray(src_page_indices, dtype=np.int64)
dst_pages = np.asarray(dst_page_indices, dtype=np.int64)
source_capacity = src_pages.size * physical_page_size
if num_kv_tokens is None:
num_kv_tokens = source_capacity
if not 0 <= num_kv_tokens <= source_capacity:
raise ValueError(
"num_kv_tokens must fit in the provided source pages, "
f"got tokens={num_kv_tokens}, capacity={source_capacity}"
)
if src_pages.size == 0:
num_kv_tokens = src_pages.size * physical_page_size
if num_kv_tokens == 0:
empty = np.empty((0,), dtype=np.int64)
return DCPTokenTransferPlan(empty, empty.copy())
return DCPTokenTransferPlan(empty, empty.copy(), empty.copy(), empty.copy())
chunk_start = decode_prefix_len + src_page_offset * physical_page_size
first_owned_offset = (dcp_rank - chunk_start) % dcp_size
owned_offsets = np.arange(
first_owned_offset, num_kv_tokens, dcp_size, dtype=np.int64
)
src_token_indices = (
src_pages[owned_offsets // physical_page_size] * physical_page_size
+ owned_offsets % physical_page_size
)
relative_positions = src_page_offset * physical_page_size + owned_offsets
dst_local_offsets = relative_positions // dcp_size
dst_page_ordinals = dst_local_offsets // physical_page_size
if dst_page_ordinals.size and (
dst_pages.size == 0 or int(dst_page_ordinals.max()) >= dst_pages.size
):
required_pages = int(dst_page_ordinals.max()) + 1
raise ValueError(
"Insufficient destination DCP pages: "
f"required={required_pages}, provided={dst_pages.size}, "
f"src_page_offset={src_page_offset}, dcp_rank={dcp_rank}"
def rows(offsets, dst_page_size, dst_local):
return (
src_pages[offsets // physical_page_size] * physical_page_size
+ offsets % physical_page_size,
dst_pages[dst_local // dst_page_size] * dst_page_size
+ dst_local % dst_page_size,
)
dst_token_indices = (
dst_pages[dst_page_ordinals] * physical_page_size
+ dst_local_offsets % physical_page_size
draft_offsets = np.arange(num_kv_tokens, dtype=np.int64)
draft_local = src_page_offset * physical_page_size + draft_offsets
chunk_start = decode_prefix_len + src_page_offset * physical_page_size
target_offsets = np.arange(
(dcp_rank - chunk_start) % dcp_size,
num_kv_tokens,
dcp_size,
dtype=np.int64,
)
return DCPTokenTransferPlan(src_token_indices, dst_token_indices)
target_local = (src_page_offset * physical_page_size + target_offsets) // dcp_size
target_src, target_dst = rows(target_offsets, physical_page_size, target_local)
draft_src, draft_dst = rows(draft_offsets, virtual_page_size, draft_local)
return DCPTokenTransferPlan(target_src, target_dst, draft_src, draft_dst)
@@ -567,6 +567,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_lens
kv_args.num_draft_entries = num_draft_entries
kv_args.kv_layer_ids = build_kv_layer_ids(
token_to_kv_pool=self.token_to_kv_pool,
draft_token_to_kv_pool=self.draft_token_to_kv_pool,
@@ -982,18 +982,6 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
if num_kv_tokens is None:
raise ValueError("PD DCP transfer requires num_kv_tokens")
physical_page_size = self.kv_args.page_size
plan = build_dcp_token_transfer_plan(
prefill_kv_indices,
dst_kv_indices,
physical_page_size=physical_page_size,
dcp_size=dst_dcp_size,
dcp_rank=dst_dcp_rank,
src_page_offset=src_page_offset,
decode_prefix_len=decode_prefix_len,
num_kv_tokens=num_kv_tokens,
)
if plan.src_token_indices.size == 0:
return 0
src_layer_ids = self.kv_args.kv_layer_ids
if src_layer_ids or dst_layer_ids:
@@ -1010,38 +998,70 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
self.kv_args.kv_data_ptrs,
dst_kv_ptrs,
)
src_token_indices = plan.src_token_indices
dst_token_indices = plan.dst_token_indices
if pack_buffer is not None:
num_draft = self.kv_args.num_draft_entries
num_target = len(src_kv_ptrs) - num_draft
plan = build_dcp_token_transfer_plan(
prefill_kv_indices,
dst_kv_indices,
physical_page_size=physical_page_size,
dcp_size=dst_dcp_size,
dcp_rank=dst_dcp_rank,
src_page_offset=src_page_offset,
decode_prefix_len=decode_prefix_len,
num_kv_tokens=num_kv_tokens,
)
if plan.empty():
return 0
target_src_kv_ptrs = src_kv_ptrs[:num_target]
src_token_indices = plan.target_src_token_indices
if pack_buffer is not None and src_token_indices.size:
from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
packed = try_pack_dcp_src(
pack_buffer=pack_buffer,
kv_data_ptrs=src_kv_ptrs,
kv_data_ptrs=target_src_kv_ptrs,
src_token_indices=src_token_indices,
token_item_lens=dcp_token_item_lens[: len(src_kv_ptrs)],
token_item_lens=dcp_token_item_lens[:num_target],
)
if packed is not None:
src_kv_ptrs, src_token_indices = packed
target_src_kv_ptrs, src_token_indices = packed
layers_current_pp_stage = len(src_kv_ptrs)
src_groups, dst_groups = group_concurrent_contiguous(
src_token_indices,
dst_token_indices,
)
layers_params = [
(
src_kv_ptrs[layer_id],
dst_kv_ptrs[layer_id],
dcp_token_item_lens[layer_id],
layers_params = []
if src_token_indices.size:
target_groups = group_concurrent_contiguous(
src_token_indices,
plan.target_dst_token_indices,
)
for layer_id in range(layers_current_pp_stage)
]
layers_params += [
(
target_src_kv_ptrs[entry],
dst_kv_ptrs[entry],
dcp_token_item_lens[entry],
target_groups,
)
for entry in range(num_target)
]
if num_draft > 0 and plan.draft_src_token_indices.size:
draft_groups = group_concurrent_contiguous(
plan.draft_src_token_indices,
plan.draft_dst_token_indices,
)
layers_params += [
(
src_kv_ptrs[num_target + entry],
dst_kv_ptrs[num_target + entry],
dcp_token_item_lens[num_target + entry],
draft_groups,
)
for entry in range(num_draft)
]
def set_transfer_blocks(
src_ptr: int, dst_ptr: int, token_item_len: int
src_ptr: int, dst_ptr: int, token_item_len: int, groups
) -> List[Tuple[int, int, int]]:
src_groups, dst_groups = groups
return [
(
src_ptr + int(src_group[0]) * token_item_len,
@@ -1051,24 +1071,24 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
for src_group, dst_group in zip(src_groups, dst_groups)
]
def process_layer(src_ptr: int, dst_ptr: int, token_item_len: int) -> int:
def process_layer(
src_ptr: int, dst_ptr: int, token_item_len: int, groups
) -> int:
return self._transfer_data(
mooncake_session_id,
set_transfer_blocks(src_ptr, dst_ptr, token_item_len),
set_transfer_blocks(src_ptr, dst_ptr, token_item_len, groups),
)
if self.enable_custom_mem_pool:
futures = [
executor.submit(process_layer, src_ptr, dst_ptr, token_item_len)
for src_ptr, dst_ptr, token_item_len in layers_params
executor.submit(process_layer, *layer_params)
for layer_params in layers_params
]
return self._await_transfer_futures(futures)
transfer_blocks = []
for src_ptr, dst_ptr, token_item_len in layers_params:
transfer_blocks.extend(
set_transfer_blocks(src_ptr, dst_ptr, token_item_len)
)
for layer_params in layers_params:
transfer_blocks.extend(set_transfer_blocks(*layer_params))
return self._transfer_data(mooncake_session_id, transfer_blocks)
def send_kvcache_slice(
@@ -2174,10 +2194,15 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
decode_kv_args.dst_dcp_rank,
)
if decode_kv_args.requires_dcp_relayout:
num_entries = len(self.kv_args.kv_item_lens)
num_draft = self.kv_args.num_draft_entries
dst_item_lens: List[Optional[int]] = [
decode_kv_args.dst_kv_item_len
] * (num_entries - num_draft) + [None] * num_draft
decode_kv_args.dcp_token_item_lens = (
self.prepare_dcp_token_item_lens(
[decode_kv_args.dst_kv_item_len]
* len(self.kv_args.kv_item_lens)
dst_item_lens,
decode_kv_args.dst_dcp_size,
)
)
self._init_dcp_pack_buffers_once(decode_kv_args.dst_dcp_size)
+74 -32
View File
@@ -1053,7 +1053,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
)
peer_info.dst_homogeneous_mem_kind = dst_mem_kind
peer_info.dcp_token_item_lens = self.prepare_dcp_token_item_lens(
dst_kv_item_lens
dst_kv_item_lens,
peer_info.dst_dcp_size,
)
return
@@ -1315,15 +1316,17 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
packed_src = self._pack_dcp_rank_once(
pack_buffer,
dst_info,
plan.src_token_indices,
plan.target_src_token_indices,
packed_source_by_dcp_rank,
)
kv_xfer_handle = self.send_kvcache_dcp(
req.agent_name,
dst_info,
plan,
notif,
packed_src,
handles.extend(
self.send_kvcache_dcp(
req.agent_name,
dst_info,
plan,
notif,
packed_src,
)
)
elif (
self.is_mla_backend
@@ -1772,12 +1775,13 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
token_item_lens = dst_info.dcp_token_item_lens
assert token_item_lens is not None
num_target = len(self.kv_args.kv_data_ptrs) - self.kv_args.num_draft_entries
rank_stride = pack_buffer.get_size() // dst_info.dst_dcp_size
packed_source_by_dcp_rank[rank] = try_pack_dcp_src(
pack_buffer=pack_buffer,
kv_data_ptrs=self.kv_args.kv_data_ptrs,
kv_data_ptrs=self.kv_args.kv_data_ptrs[:num_target],
src_token_indices=src_token_indices,
token_item_lens=token_item_lens[: len(self.kv_args.kv_data_ptrs)],
token_item_lens=token_item_lens[:num_target],
pack_offset_bytes=rank * rank_stride,
)
return packed_source_by_dcp_rank[rank]
@@ -1794,35 +1798,73 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
raise RuntimeError("Missing NIXL source KV memory kind")
if dst_info.dst_homogeneous_mem_kind is None:
raise RuntimeError("Missing NIXL destination KV memory kind")
if plan.src_token_indices.size == 0:
self.agent.send_notif(peer_name, notif.encode("ascii"))
return None
token_item_lens = dst_info.dcp_token_item_lens
assert token_item_lens is not None
num_draft = self.kv_args.num_draft_entries
num_target = len(self.kv_args.kv_data_ptrs) - num_draft
dst_kv_ptrs = [
dst_info.dst_kv_ptrs[dst_idx] for dst_idx in dst_info.dcp_dst_region_indices
]
src_kv_ptrs = self.kv_args.kv_data_ptrs
src_token_indices = plan.src_token_indices
if packed_src is not None:
src_kv_ptrs, src_token_indices = packed_src
token_item_lens = token_item_lens[: len(src_kv_ptrs)]
return self._send_kvcache_generic(
peer_name=peer_name,
src_data_ptrs=src_kv_ptrs,
dst_data_ptrs=dst_kv_ptrs,
item_lens=token_item_lens,
prefill_data_indices=src_token_indices,
dst_data_indices=plan.dst_token_indices,
dst_gpu_id=dst_info.gpu_id,
notif=notif,
src_mem_kind=self.src_mem_kind,
dst_mem_kind=dst_info.dst_homogeneous_mem_kind,
force_flat=True,
bypass_prepped=True,
)
parts = []
if plan.target_src_token_indices.size:
src_kv_ptrs = self.kv_args.kv_data_ptrs[:num_target]
src_token_indices = plan.target_src_token_indices
if packed_src is not None:
src_kv_ptrs, src_token_indices = packed_src
parts.append(
(
src_kv_ptrs,
dst_kv_ptrs[:num_target],
token_item_lens[:num_target],
src_token_indices,
plan.target_dst_token_indices,
)
)
if num_draft > 0 and plan.draft_src_token_indices.size:
parts.append(
(
self.kv_args.kv_data_ptrs[num_target:],
dst_kv_ptrs[num_target:],
token_item_lens[num_target:],
plan.draft_src_token_indices,
plan.draft_dst_token_indices,
)
)
if not parts:
self.agent.send_notif(peer_name, notif.encode("ascii"))
return []
handles = []
for part_idx, (
src_ptrs,
part_dst_ptrs,
part_item_lens,
src_indices,
dst_indices,
) in enumerate(parts):
part_notif = (
notif if len(parts) == 1 else f"{notif}_part_{part_idx}_{len(parts)}"
)
handles.append(
self._send_kvcache_generic(
peer_name=peer_name,
src_data_ptrs=src_ptrs,
dst_data_ptrs=part_dst_ptrs,
item_lens=part_item_lens,
prefill_data_indices=src_indices,
dst_data_indices=dst_indices,
dst_gpu_id=dst_info.gpu_id,
notif=part_notif,
src_mem_kind=self.src_mem_kind,
dst_mem_kind=dst_info.dst_homogeneous_mem_kind,
force_flat=True,
bypass_prepped=True,
)
)
return handles
def send_kvcache_mixed(
self,
@@ -267,6 +267,7 @@ class PrefillBootstrapQueue:
kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_lens
kv_args.num_draft_entries = num_draft_entries
kv_args.kv_layer_ids = build_kv_layer_ids(
token_to_kv_pool=self.token_to_kv_pool,
draft_token_to_kv_pool=draft_kv_pool,