diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index dcaf4c95d..2e3f5dfe2 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -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: diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index b2864ae4b..78fe9d585 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -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: diff --git a/python/sglang/srt/disaggregation/common/dcp_pack.py b/python/sglang/srt/disaggregation/common/dcp_pack.py index db0a6c534..779b2b92d 100644 --- a/python/sglang/srt/disaggregation/common/dcp_pack.py +++ b/python/sglang/srt/disaggregation/common/dcp_pack.py @@ -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}" diff --git a/python/sglang/srt/disaggregation/common/utils.py b/python/sglang/srt/disaggregation/common/utils.py index 1571c3d19..95991d874 100644 --- a/python/sglang/srt/disaggregation/common/utils.py +++ b/python/sglang/srt/disaggregation/common/utils.py @@ -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) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index b16c5cec2..ebb553272 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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, diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 970824f99..98a1c289a 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -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) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 980568e2e..2b88fe520 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -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, diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 1afebb284..05c08ff7d 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -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, diff --git a/test/registered/e2e/disaggregation/test_kimi_linear_pd_dcp4_dspark.py b/test/registered/e2e/disaggregation/test_kimi_linear_pd_dcp4_dspark.py new file mode 100644 index 000000000..c280e5863 --- /dev/null +++ b/test/registered/e2e/disaggregation/test_kimi_linear_pd_dcp4_dspark.py @@ -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() diff --git a/test/registered/unit/disaggregation/test_dcp_pack.py b/test/registered/unit/disaggregation/test_dcp_pack.py index c8990ef39..cef4295a8 100644 --- a/test/registered/unit/disaggregation/test_dcp_pack.py +++ b/test/registered/unit/disaggregation/test_dcp_pack.py @@ -1,10 +1,12 @@ import unittest from contextlib import nullcontext +from types import SimpleNamespace from unittest.mock import Mock, patch import numpy as np import torch +from sglang.srt.disaggregation.common.conn import CommonKVManager from sglang.srt.disaggregation.common.dcp_pack import ( dcp_pack_buffer_bytes, 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") -class TestPackedDcpGrouping(CustomTestCase): - def test_packed_groups_collapse_cyclic_src(self): - page_size = 64 - dcp_size = 4 - 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, - 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)) +def _plan(*, src, dst, page_size, dcp_size, dcp_rank, **kwargs): + return build_dcp_token_transfer_plan( + np.asarray(src, dtype=np.int32), + np.asarray(dst, dtype=np.int32), + physical_page_size=page_size, + dcp_size=dcp_size, + dcp_rank=dcp_rank, + **kwargs, + ) - packed_src = np.arange(plan.dst_token_indices.size, dtype=np.int64) - packed_groups, _ = group_concurrent_contiguous( - packed_src, plan.dst_token_indices + +class TestDcpTokenTransferPlan(CustomTestCase): + 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) - self.assertEqual(len(packed_groups[0]), 64) + 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 + ) class TestDcpPackBufferBytes(CustomTestCase): diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index 58e9fa18f..503a33dc0 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -575,7 +575,9 @@ class TestNixlTransferWorker(CustomTestCase): mgr.is_hybrid_mla_backend = False mgr.attn_tp_size = 1 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.failure_lock = threading.Lock() mgr.failure_records = {} @@ -674,6 +676,7 @@ class TestNixlTransferWorker(CustomTestCase): engine_rank=0, kv_data_ptrs=[0x1000], page_size=4, + num_draft_entries=0, ) mgr._dcp_pack_buffers = [SimpleNamespace(get_size=lambda: 16)] @@ -686,7 +689,8 @@ class TestNixlTransferWorker(CustomTestCase): def send_kvcache_dcp(*args, **kwargs): 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) submitted_counts_at_poll = []