[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:
co-authored by
Claude Fable 5
Cursor
parent
7d9c57da6e
commit
a207786205
@@ -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,16 +340,33 @@ 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:
|
||||||
|
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(
|
raise RuntimeError(
|
||||||
"PD DCP source/destination KV geometry differs: "
|
"PD DCP source/destination KV geometry differs at entry "
|
||||||
f"src={src_token_lens}, dst={dst_token_lens}"
|
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
|
||||||
|
|
||||||
|
|||||||
@@ -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())
|
||||||
|
|
||||||
|
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,
|
||||||
|
)
|
||||||
|
|
||||||
|
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
|
chunk_start = decode_prefix_len + src_page_offset * physical_page_size
|
||||||
first_owned_offset = (dcp_rank - chunk_start) % dcp_size
|
target_offsets = np.arange(
|
||||||
owned_offsets = np.arange(
|
(dcp_rank - chunk_start) % dcp_size,
|
||||||
first_owned_offset, num_kv_tokens, dcp_size, dtype=np.int64
|
num_kv_tokens,
|
||||||
|
dcp_size,
|
||||||
|
dtype=np.int64,
|
||||||
)
|
)
|
||||||
src_token_indices = (
|
target_local = (src_page_offset * physical_page_size + target_offsets) // dcp_size
|
||||||
src_pages[owned_offsets // physical_page_size] * physical_page_size
|
target_src, target_dst = rows(target_offsets, physical_page_size, target_local)
|
||||||
+ owned_offsets % physical_page_size
|
draft_src, draft_dst = rows(draft_offsets, virtual_page_size, draft_local)
|
||||||
)
|
return DCPTokenTransferPlan(target_src, target_dst, draft_src, draft_dst)
|
||||||
|
|
||||||
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 = (
|
|
||||||
dst_pages[dst_page_ordinals] * physical_page_size
|
|
||||||
+ dst_local_offsets % physical_page_size
|
|
||||||
)
|
|
||||||
return DCPTokenTransferPlan(src_token_indices, dst_token_indices)
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
target_groups = group_concurrent_contiguous(
|
||||||
src_token_indices,
|
src_token_indices,
|
||||||
dst_token_indices,
|
plan.target_dst_token_indices,
|
||||||
)
|
)
|
||||||
|
layers_params += [
|
||||||
layers_params = [
|
|
||||||
(
|
(
|
||||||
src_kv_ptrs[layer_id],
|
target_src_kv_ptrs[entry],
|
||||||
dst_kv_ptrs[layer_id],
|
dst_kv_ptrs[entry],
|
||||||
dcp_token_item_lens[layer_id],
|
dcp_token_item_lens[entry],
|
||||||
|
target_groups,
|
||||||
)
|
)
|
||||||
for layer_id in range(layers_current_pp_stage)
|
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)
|
||||||
|
|||||||
@@ -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,16 +1316,18 @@ 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(
|
||||||
|
self.send_kvcache_dcp(
|
||||||
req.agent_name,
|
req.agent_name,
|
||||||
dst_info,
|
dst_info,
|
||||||
plan,
|
plan,
|
||||||
notif,
|
notif,
|
||||||
packed_src,
|
packed_src,
|
||||||
)
|
)
|
||||||
|
)
|
||||||
elif (
|
elif (
|
||||||
self.is_mla_backend
|
self.is_mla_backend
|
||||||
or self.is_hybrid_mla_backend
|
or self.is_hybrid_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
|
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:
|
if packed_src is not None:
|
||||||
src_kv_ptrs, src_token_indices = packed_src
|
src_kv_ptrs, src_token_indices = packed_src
|
||||||
token_item_lens = token_item_lens[: len(src_kv_ptrs)]
|
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,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
return self._send_kvcache_generic(
|
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,
|
peer_name=peer_name,
|
||||||
src_data_ptrs=src_kv_ptrs,
|
src_data_ptrs=src_ptrs,
|
||||||
dst_data_ptrs=dst_kv_ptrs,
|
dst_data_ptrs=part_dst_ptrs,
|
||||||
item_lens=token_item_lens,
|
item_lens=part_item_lens,
|
||||||
prefill_data_indices=src_token_indices,
|
prefill_data_indices=src_indices,
|
||||||
dst_data_indices=plan.dst_token_indices,
|
dst_data_indices=dst_indices,
|
||||||
dst_gpu_id=dst_info.gpu_id,
|
dst_gpu_id=dst_info.gpu_id,
|
||||||
notif=notif,
|
notif=part_notif,
|
||||||
src_mem_kind=self.src_mem_kind,
|
src_mem_kind=self.src_mem_kind,
|
||||||
dst_mem_kind=dst_info.dst_homogeneous_mem_kind,
|
dst_mem_kind=dst_info.dst_homogeneous_mem_kind,
|
||||||
force_flat=True,
|
force_flat=True,
|
||||||
bypass_prepped=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)
|
|
||||||
dst_pages = np.array([7], dtype=np.int32)
|
|
||||||
plan = build_dcp_token_transfer_plan(
|
|
||||||
src_pages,
|
|
||||||
dst_pages,
|
|
||||||
physical_page_size=page_size,
|
physical_page_size=page_size,
|
||||||
dcp_size=dcp_size,
|
dcp_size=dcp_size,
|
||||||
dcp_rank=0,
|
dcp_rank=dcp_rank,
|
||||||
num_kv_tokens=256,
|
**kwargs,
|
||||||
)
|
)
|
||||||
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,
|
||||||
|
)
|
||||||
|
np.testing.assert_array_equal(plan.draft_src_token_indices, [18, 19, 6, 7])
|
||||||
|
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
|
||||||
)
|
)
|
||||||
self.assertEqual(len(packed_groups), 1)
|
|
||||||
self.assertEqual(len(packed_groups[0]), 64)
|
|
||||||
|
|
||||||
|
|
||||||
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 = []
|
||||||
|
|||||||
Reference in New Issue
Block a user