[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 hidden_kv_layers: int
# Only used of npu, for decode total kv layers # Only used of npu, for decode total kv layers
draft_kv_layers: int draft_kv_layers: int
num_draft_entries: int = 0
class KVPoll: class KVPoll:
@@ -340,17 +340,34 @@ class CommonKVManager(BaseKVManager):
f"Unsupported PD DCP topology: {self.dcp_size} -> {dst_dcp_size}" 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 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 = [ src_token_lens = [
item_len // page_size for item_len in self.kv_args.kv_item_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] for i, dst_item_len in enumerate(dst_page_item_lens):
if src_token_lens != dst_token_lens: if dst_item_len is None:
raise RuntimeError( continue
"PD DCP source/destination KV geometry differs: " dst_page_scale = page_size * (
f"src={src_token_lens}, dst={dst_token_lens}" 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 return src_token_lens
def _register_staging_memory(self, ptr: int, size: int) -> None: 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() max_tokens = max_prefill_buffer_tokens()
if max_tokens <= 0: if max_tokens <= 0:
max_tokens = get_schedule().max_prefill_tokens 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) # 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 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. # x 576 bf16 dims x 2 B: 2.14 GiB/buffer, 8.58 GiB for 4 queues.
size_bytes = dcp_pack_buffer_bytes( 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 gpu_id = kv_args.gpu_id
device = f"cuda:{gpu_id}" device = f"cuda:{gpu_id}"
@@ -137,8 +137,16 @@ def group_concurrent_contiguous(
@dataclasses.dataclass(frozen=True) @dataclasses.dataclass(frozen=True)
class DCPTokenTransferPlan: class DCPTokenTransferPlan:
src_token_indices: npt.NDArray[np.int64] target_src_token_indices: npt.NDArray[np.int64]
dst_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( def build_dcp_token_transfer_plan(
@@ -152,52 +160,38 @@ def build_dcp_token_transfer_plan(
decode_prefix_len: int = 0, decode_prefix_len: int = 0,
num_kv_tokens: Optional[int] = None, num_kv_tokens: Optional[int] = None,
) -> DCPTokenTransferPlan: ) -> 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 virtual_page_size = physical_page_size * dcp_size
if decode_prefix_len % virtual_page_size != 0: if decode_prefix_len % virtual_page_size != 0:
raise ValueError( raise ValueError(
"PD DCP transfer requires decode_prefix_len to align to the virtual " "PD DCP transfer requires decode_prefix_len to align to the virtual "
f"DCP page size ({virtual_page_size}), got {decode_prefix_len}" 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: if num_kv_tokens is None:
num_kv_tokens = source_capacity num_kv_tokens = src_pages.size * physical_page_size
if not 0 <= num_kv_tokens <= source_capacity: if num_kv_tokens == 0:
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:
empty = np.empty((0,), dtype=np.int64) 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 def rows(offsets, dst_page_size, dst_local):
first_owned_offset = (dcp_rank - chunk_start) % dcp_size return (
owned_offsets = np.arange( src_pages[offsets // physical_page_size] * physical_page_size
first_owned_offset, num_kv_tokens, dcp_size, dtype=np.int64 + offsets % physical_page_size,
) dst_pages[dst_local // dst_page_size] * dst_page_size
src_token_indices = ( + dst_local % dst_page_size,
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}"
) )
dst_token_indices = ( draft_offsets = np.arange(num_kv_tokens, dtype=np.int64)
dst_pages[dst_page_ordinals] * physical_page_size draft_local = src_page_offset * physical_page_size + draft_offsets
+ dst_local_offsets % physical_page_size 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_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_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( kv_args.kv_layer_ids = build_kv_layer_ids(
token_to_kv_pool=self.token_to_kv_pool, token_to_kv_pool=self.token_to_kv_pool,
draft_token_to_kv_pool=self.draft_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: if num_kv_tokens is None:
raise ValueError("PD DCP transfer requires num_kv_tokens") raise ValueError("PD DCP transfer requires num_kv_tokens")
physical_page_size = self.kv_args.page_size 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 src_layer_ids = self.kv_args.kv_layer_ids
if src_layer_ids or dst_layer_ids: if src_layer_ids or dst_layer_ids:
@@ -1010,38 +998,70 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
self.kv_args.kv_data_ptrs, self.kv_args.kv_data_ptrs,
dst_kv_ptrs, dst_kv_ptrs,
) )
src_token_indices = plan.src_token_indices num_draft = self.kv_args.num_draft_entries
dst_token_indices = plan.dst_token_indices num_target = len(src_kv_ptrs) - num_draft
if pack_buffer is not None:
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 from sglang.srt.disaggregation.common.dcp_pack import try_pack_dcp_src
packed = try_pack_dcp_src( packed = try_pack_dcp_src(
pack_buffer=pack_buffer, pack_buffer=pack_buffer,
kv_data_ptrs=src_kv_ptrs, kv_data_ptrs=target_src_kv_ptrs,
src_token_indices=src_token_indices, 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: 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) layers_params = []
src_groups, dst_groups = group_concurrent_contiguous( if src_token_indices.size:
src_token_indices, target_groups = group_concurrent_contiguous(
dst_token_indices, src_token_indices,
) plan.target_dst_token_indices,
layers_params = [
(
src_kv_ptrs[layer_id],
dst_kv_ptrs[layer_id],
dcp_token_item_lens[layer_id],
) )
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( 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]]: ) -> List[Tuple[int, int, int]]:
src_groups, dst_groups = groups
return [ return [
( (
src_ptr + int(src_group[0]) * token_item_len, 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) 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( return self._transfer_data(
mooncake_session_id, 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: if self.enable_custom_mem_pool:
futures = [ futures = [
executor.submit(process_layer, src_ptr, dst_ptr, token_item_len) executor.submit(process_layer, *layer_params)
for src_ptr, dst_ptr, token_item_len in layers_params for layer_params in layers_params
] ]
return self._await_transfer_futures(futures) return self._await_transfer_futures(futures)
transfer_blocks = [] transfer_blocks = []
for src_ptr, dst_ptr, token_item_len in layers_params: for layer_params in layers_params:
transfer_blocks.extend( transfer_blocks.extend(set_transfer_blocks(*layer_params))
set_transfer_blocks(src_ptr, dst_ptr, token_item_len)
)
return self._transfer_data(mooncake_session_id, transfer_blocks) return self._transfer_data(mooncake_session_id, transfer_blocks)
def send_kvcache_slice( def send_kvcache_slice(
@@ -2174,10 +2194,15 @@ class MooncakeKVManager(StagingManagerMixin, CommonKVManager):
decode_kv_args.dst_dcp_rank, decode_kv_args.dst_dcp_rank,
) )
if decode_kv_args.requires_dcp_relayout: 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 = ( decode_kv_args.dcp_token_item_lens = (
self.prepare_dcp_token_item_lens( self.prepare_dcp_token_item_lens(
[decode_kv_args.dst_kv_item_len] dst_item_lens,
* len(self.kv_args.kv_item_lens) 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)
+74 -32
View File
@@ -1053,7 +1053,8 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
) )
peer_info.dst_homogeneous_mem_kind = dst_mem_kind peer_info.dst_homogeneous_mem_kind = dst_mem_kind
peer_info.dcp_token_item_lens = self.prepare_dcp_token_item_lens( 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 return
@@ -1315,15 +1316,17 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
packed_src = self._pack_dcp_rank_once( packed_src = self._pack_dcp_rank_once(
pack_buffer, pack_buffer,
dst_info, dst_info,
plan.src_token_indices, plan.target_src_token_indices,
packed_source_by_dcp_rank, packed_source_by_dcp_rank,
) )
kv_xfer_handle = self.send_kvcache_dcp( handles.extend(
req.agent_name, self.send_kvcache_dcp(
dst_info, req.agent_name,
plan, dst_info,
notif, plan,
packed_src, notif,
packed_src,
)
) )
elif ( elif (
self.is_mla_backend self.is_mla_backend
@@ -1772,12 +1775,13 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
token_item_lens = dst_info.dcp_token_item_lens token_item_lens = dst_info.dcp_token_item_lens
assert token_item_lens is not None 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 rank_stride = pack_buffer.get_size() // dst_info.dst_dcp_size
packed_source_by_dcp_rank[rank] = try_pack_dcp_src( packed_source_by_dcp_rank[rank] = try_pack_dcp_src(
pack_buffer=pack_buffer, 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, 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, pack_offset_bytes=rank * rank_stride,
) )
return packed_source_by_dcp_rank[rank] return packed_source_by_dcp_rank[rank]
@@ -1794,35 +1798,73 @@ class NixlKVManager(StagingManagerMixin, CommonKVManager):
raise RuntimeError("Missing NIXL source KV memory kind") raise RuntimeError("Missing NIXL source KV memory kind")
if dst_info.dst_homogeneous_mem_kind is None: if dst_info.dst_homogeneous_mem_kind is None:
raise RuntimeError("Missing NIXL destination KV memory kind") 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 token_item_lens = dst_info.dcp_token_item_lens
assert token_item_lens is not None 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_kv_ptrs = [
dst_info.dst_kv_ptrs[dst_idx] for dst_idx in dst_info.dcp_dst_region_indices 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( parts = []
peer_name=peer_name, if plan.target_src_token_indices.size:
src_data_ptrs=src_kv_ptrs, src_kv_ptrs = self.kv_args.kv_data_ptrs[:num_target]
dst_data_ptrs=dst_kv_ptrs, src_token_indices = plan.target_src_token_indices
item_lens=token_item_lens, if packed_src is not None:
prefill_data_indices=src_token_indices, src_kv_ptrs, src_token_indices = packed_src
dst_data_indices=plan.dst_token_indices, parts.append(
dst_gpu_id=dst_info.gpu_id, (
notif=notif, src_kv_ptrs,
src_mem_kind=self.src_mem_kind, dst_kv_ptrs[:num_target],
dst_mem_kind=dst_info.dst_homogeneous_mem_kind, token_item_lens[:num_target],
force_flat=True, src_token_indices,
bypass_prepped=True, 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( def send_kvcache_mixed(
self, self,
@@ -267,6 +267,7 @@ class PrefillBootstrapQueue:
kv_args.kv_data_ptrs = kv_data_ptrs kv_args.kv_data_ptrs = kv_data_ptrs
kv_args.kv_data_lens = kv_data_lens kv_args.kv_data_lens = kv_data_lens
kv_args.kv_item_lens = kv_item_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( kv_args.kv_layer_ids = build_kv_layer_ids(
token_to_kv_pool=self.token_to_kv_pool, token_to_kv_pool=self.token_to_kv_pool,
draft_token_to_kv_pool=draft_kv_pool, draft_token_to_kv_pool=draft_kv_pool,
@@ -0,0 +1,175 @@
import json
import os
import shutil
import tempfile
import unittest
from pathlib import Path
import requests
import torch
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.server_fixtures.disaggregation_fixture import (
PDDisaggregationServerBase,
)
register_cuda_ci(est_time=500, stage="nightly", runner_config="8-gpu-b200")
KIMI_LINEAR_MODEL = "moonshotai/Kimi-Linear-48B-A3B-Instruct"
PHYSICAL_PAGE_SIZE = 64
CHUNKED_PREFILL_SIZE = 8192
def _has_eight_blackwell_gpus() -> bool:
if not torch.cuda.is_available() or torch.cuda.device_count() < 8:
return False
return all(
torch.cuda.get_device_capability(device_index) >= (10, 0)
for device_index in range(8)
)
def _write_dummy_qwen3_dspark_draft(root: Path) -> str:
draft_dir = root / "qwen3-dspark-kimi-proxy"
draft_dir.mkdir()
config = {
"architectures": ["Qwen3DSparkModel"],
"model_type": "qwen3",
"dtype": "bfloat16",
"hidden_size": 2304,
"intermediate_size": 9216,
"num_hidden_layers": 5,
"num_attention_heads": 16,
"num_key_value_heads": 4,
"head_dim": 128,
"hidden_act": "silu",
"rms_norm_eps": 1e-5,
"attention_bias": False,
"attention_dropout": 0.0,
"max_position_embeddings": 1048576,
"rope_parameters": {
"rope_theta": 10000.0,
"rope_type": "default",
},
"vocab_size": 163840,
"bos_token_id": 163584,
"eos_token_id": 163586,
"mask_token_id": 163839,
"block_size": 7,
"markov_rank": 256,
"markov_head_type": "vanilla",
"enable_confidence_head": True,
"confidence_head_with_markov": True,
"num_target_layers": 27,
"target_layer_ids": [1, 7, 13, 19, 26],
"layer_types": ["full_attention"] * 5,
"tie_word_embeddings": False,
"use_cache": True,
}
(draft_dir / "config.json").write_text(json.dumps(config), encoding="utf-8")
return str(draft_dir)
@unittest.skipUnless(
_has_eight_blackwell_gpus(),
"Kimi-Linear PD DCP4 + DSPARK requires eight Blackwell GPUs",
)
class TestKimiLinearPDDCP4DSpark(GSM8KMixin, PDDisaggregationServerBase):
model = KIMI_LINEAR_MODEL
gsm8k_score_threshold = 0.88
gsm8k_num_examples = 400
gsm8k_num_threads = 64
gsm8k_num_shots = 5
@classmethod
def setUpClass(cls):
super().setUpClass()
os.environ["MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER"] = "65535"
os.environ["MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER"] = "65535"
cls._draft_root = tempfile.mkdtemp(prefix="dspark_pd_dcp_draft_")
draft_path = _write_dummy_qwen3_dspark_draft(Path(cls._draft_root))
dspark_args = [
"--speculative-algorithm",
"DSPARK",
"--speculative-draft-model-path",
draft_path,
"--speculative-draft-load-format",
"dummy",
"--speculative-attention-mode",
"decode",
"--speculative-draft-attention-backend",
"trtllm_mha",
]
common_args = [
"--attention-backend",
"tokenspeed_mla",
"--kv-cache-dtype",
"fp8_e4m3",
"--dtype",
"bfloat16",
"--random-seed",
"0",
"--page-size",
str(PHYSICAL_PAGE_SIZE),
"--cuda-graph-backend-prefill",
"disabled",
"--mem-fraction-static",
"0.80",
] + dspark_args
cls.prefill_tp_size = 4
cls.decode_tp_size = 4
cls.decode_base_gpu_id = 4
cls.extra_prefill_args = common_args + [
"--ep-size",
"4",
"--chunked-prefill-size",
str(CHUNKED_PREFILL_SIZE),
]
cls.extra_decode_args = common_args + [
"--dcp-size",
"4",
"--dcp-comm-backend",
"a2a",
"--dcp-replicate-q-proj",
"--cuda-graph-max-bs-decode",
"64",
]
cls.extra_prefill_env = {"SGLANG_RAGGED_VERIFY_MODE": "static"}
cls.extra_decode_env = {"SGLANG_RAGGED_VERIFY_MODE": "static"}
cls.launch_all()
@classmethod
def tearDownClass(cls):
os.environ.pop("MC_TCP_MAX_QUEUED_TRANSFERS_PER_PEER", None)
os.environ.pop("MC_TCP_MAX_PENDING_ADMISSIONS_PER_PEER", None)
shutil.rmtree(cls._draft_root, ignore_errors=True)
super().tearDownClass()
def test_spec_verify_runs_on_decode(self):
response = requests.post(
self.base_url + "/generate",
json={
"text": "The capital of France is",
"sampling_params": {
"temperature": 0,
"max_new_tokens": 32,
"ignore_eos": True,
},
},
timeout=300,
)
response.raise_for_status()
meta_info = response.json()["meta_info"]
self.assertGreater(
meta_info.get("spec_verify_ct", 0),
0,
"DSPARK verify did not run on the decode side",
)
self.assertGreater(meta_info["completion_tokens"], 0)
if __name__ == "__main__":
unittest.main()
@@ -1,10 +1,12 @@
import unittest import unittest
from contextlib import nullcontext from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import Mock, patch from unittest.mock import Mock, patch
import numpy as np import numpy as np
import torch import torch
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.common.dcp_pack import ( from sglang.srt.disaggregation.common.dcp_pack import (
dcp_pack_buffer_bytes, dcp_pack_buffer_bytes,
try_pack_dcp_src, try_pack_dcp_src,
@@ -19,32 +21,170 @@ from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=12, suite="base-a-test-cpu") register_cpu_ci(est_time=12, suite="base-a-test-cpu")
class TestPackedDcpGrouping(CustomTestCase): def _plan(*, src, dst, page_size, dcp_size, dcp_rank, **kwargs):
def test_packed_groups_collapse_cyclic_src(self): return build_dcp_token_transfer_plan(
page_size = 64 np.asarray(src, dtype=np.int32),
dcp_size = 4 np.asarray(dst, dtype=np.int32),
src_pages = np.arange(4, dtype=np.int32) physical_page_size=page_size,
dst_pages = np.array([7], dtype=np.int32) dcp_size=dcp_size,
plan = build_dcp_token_transfer_plan( dcp_rank=dcp_rank,
src_pages, **kwargs,
dst_pages, )
physical_page_size=page_size,
dcp_size=dcp_size,
dcp_rank=0,
num_kv_tokens=256,
)
raw_src, _ = group_concurrent_contiguous(
plan.src_token_indices, plan.dst_token_indices
)
self.assertEqual(len(raw_src), 64)
self.assertTrue(all(len(group) == 1 for group in raw_src))
packed_src = np.arange(plan.dst_token_indices.size, dtype=np.int64)
packed_groups, _ = group_concurrent_contiguous( class TestDcpTokenTransferPlan(CustomTestCase):
packed_src, plan.dst_token_indices def test_one_virtual_page_explicit_rows(self):
# P=2, N=4. Prefill pages 5,2,11,4; decode virtual page 7.
# pos 0..7 src rows: 10,11, 4,5, 22,23, 8,9
# draft dest page is P*N=8 → 56..63
# each rank stores local rows 14,15 (page P=2)
expected_draft_src = [10, 11, 4, 5, 22, 23, 8, 9]
expected_draft_dst = list(range(56, 64))
expected_target_src = {
0: [10, 22],
1: [11, 23],
2: [4, 8],
3: [5, 9],
}
seen_src = []
for rank, src in expected_target_src.items():
plan = _plan(
src=[5, 2, 11, 4],
dst=[7],
page_size=2,
dcp_size=4,
dcp_rank=rank,
num_kv_tokens=8,
)
np.testing.assert_array_equal(
plan.draft_src_token_indices, expected_draft_src
)
np.testing.assert_array_equal(
plan.draft_dst_token_indices, expected_draft_dst
)
np.testing.assert_array_equal(plan.target_src_token_indices, src)
np.testing.assert_array_equal(plan.target_dst_token_indices, [14, 15])
seen_src.extend(plan.target_src_token_indices.tolist())
self.assertEqual(sorted(seen_src), sorted(expected_draft_src))
def test_second_chunk_crosses_dest_pages(self):
# P=2, N=2 (virtual page = 4). Decode already holds a 4-token prefix;
# dst=[4, 6] is the full send-range page list. This chunk is the second
# prefill page of the send range (src_page_offset=1), so its 4 tokens
# sit at send-range pos 2..5 (absolute 6..9) and straddle virtual page
# 4 (rows 16..19) and virtual page 6 (rows 24..27).
plan = _plan(
src=[9, 3],
dst=[4, 6],
page_size=2,
dcp_size=2,
dcp_rank=0,
src_page_offset=1,
decode_prefix_len=4,
num_kv_tokens=4,
) )
self.assertEqual(len(packed_groups), 1) np.testing.assert_array_equal(plan.draft_src_token_indices, [18, 19, 6, 7])
self.assertEqual(len(packed_groups[0]), 64) np.testing.assert_array_equal(plan.draft_dst_token_indices, [18, 19, 24, 25])
# rank 0 owns absolute pos 6, 8 -> per-rank slots 1, 2 -> pages 4, 6.
np.testing.assert_array_equal(plan.target_src_token_indices, [18, 6])
np.testing.assert_array_equal(plan.target_dst_token_indices, [9, 12])
plan_r1 = _plan(
src=[9, 3],
dst=[4, 6],
page_size=2,
dcp_size=2,
dcp_rank=1,
src_page_offset=1,
decode_prefix_len=4,
num_kv_tokens=4,
)
np.testing.assert_array_equal(plan_r1.draft_src_token_indices, [18, 19, 6, 7])
np.testing.assert_array_equal(plan_r1.draft_dst_token_indices, [18, 19, 24, 25])
np.testing.assert_array_equal(plan_r1.target_src_token_indices, [19, 7])
np.testing.assert_array_equal(plan_r1.target_dst_token_indices, [9, 12])
def test_rejects_unaligned_prefix(self):
with self.assertRaisesRegex(ValueError, "align"):
_plan(
src=[0],
dst=[0],
page_size=2,
dcp_size=4,
dcp_rank=0,
decode_prefix_len=1,
num_kv_tokens=2,
)
def test_empty_tokens(self):
plan = _plan(
src=[0], dst=[0], page_size=2, dcp_size=4, dcp_rank=0, num_kv_tokens=0
)
self.assertTrue(plan.empty())
class TestPackedDcpGrouping(CustomTestCase):
def test_target_needs_pack_draft_does_not(self):
plan = _plan(
src=[0, 1, 2, 3],
dst=[0],
page_size=2,
dcp_size=4,
dcp_rank=0,
num_kv_tokens=8,
)
np.testing.assert_array_equal(plan.target_src_token_indices, [0, 4])
np.testing.assert_array_equal(plan.target_dst_token_indices, [0, 1])
target_src, _ = group_concurrent_contiguous(
plan.target_src_token_indices, plan.target_dst_token_indices
)
self.assertEqual(target_src, [[0], [4]])
packed_src, packed_dst = group_concurrent_contiguous(
np.arange(2, dtype=np.int64), plan.target_dst_token_indices
)
self.assertEqual(packed_src, [[0, 1]])
self.assertEqual(packed_dst, [[0, 1]])
draft_src, draft_dst = group_concurrent_contiguous(
plan.draft_src_token_indices, plan.draft_dst_token_indices
)
self.assertEqual(draft_src, [[0, 1, 2, 3, 4, 5, 6, 7]])
self.assertEqual(draft_dst, [[0, 1, 2, 3, 4, 5, 6, 7]])
def _dcp_kv_manager_stub(*, page_size, kv_item_lens, num_draft_entries):
return SimpleNamespace(
kv_args=SimpleNamespace(
page_size=page_size,
kv_item_lens=kv_item_lens,
num_draft_entries=num_draft_entries,
)
)
class TestPrepareDcpTokenItemLens(CustomTestCase):
def test_draft_tail_scales_by_dst_dcp_size(self):
mgr = _dcp_kv_manager_stub(
page_size=64,
kv_item_lens=[64 * 32, 64 * 32, 64 * 16],
num_draft_entries=1,
)
token_lens = CommonKVManager.prepare_dcp_token_item_lens(
mgr, [64 * 32, 64 * 32, 4 * 64 * 16], dst_dcp_size=4
)
self.assertEqual(token_lens, [32, 32, 16])
def test_rejects_unscaled_draft_item_len(self):
mgr = _dcp_kv_manager_stub(
page_size=64,
kv_item_lens=[64 * 32, 64 * 16],
num_draft_entries=1,
)
with self.assertRaisesRegex(RuntimeError, "geometry differs at entry 1"):
CommonKVManager.prepare_dcp_token_item_lens(
mgr, [64 * 32, 64 * 16], dst_dcp_size=4
)
class TestDcpPackBufferBytes(CustomTestCase): class TestDcpPackBufferBytes(CustomTestCase):
@@ -575,7 +575,9 @@ class TestNixlTransferWorker(CustomTestCase):
mgr.is_hybrid_mla_backend = False mgr.is_hybrid_mla_backend = False
mgr.attn_tp_size = 1 mgr.attn_tp_size = 1
mgr.transfer_source_rank = 0 mgr.transfer_source_rank = 0
mgr.kv_args = SimpleNamespace(engine_rank=0, kv_data_ptrs=[0]) mgr.kv_args = SimpleNamespace(
engine_rank=0, kv_data_ptrs=[0], num_draft_entries=0
)
mgr.exceptions = {} mgr.exceptions = {}
mgr.failure_lock = threading.Lock() mgr.failure_lock = threading.Lock()
mgr.failure_records = {} mgr.failure_records = {}
@@ -674,6 +676,7 @@ class TestNixlTransferWorker(CustomTestCase):
engine_rank=0, engine_rank=0,
kv_data_ptrs=[0x1000], kv_data_ptrs=[0x1000],
page_size=4, page_size=4,
num_draft_entries=0,
) )
mgr._dcp_pack_buffers = [SimpleNamespace(get_size=lambda: 16)] mgr._dcp_pack_buffers = [SimpleNamespace(get_size=lambda: 16)]
@@ -686,7 +689,8 @@ class TestNixlTransferWorker(CustomTestCase):
def send_kvcache_dcp(*args, **kwargs): def send_kvcache_dcp(*args, **kwargs):
submitted.append((args[0], args[-1])) submitted.append((args[0], args[-1]))
return f"handle-{args[0]}" # One handle per transfer part; the worker extends its handle list.
return [f"handle-{args[0]}"]
mgr.send_kvcache_dcp = MagicMock(side_effect=send_kvcache_dcp) mgr.send_kvcache_dcp = MagicMock(side_effect=send_kvcache_dcp)
submitted_counts_at_poll = [] submitted_counts_at_poll = []