From 2573190b9377debd135417f25b0c4185b0a3c48a Mon Sep 17 00:00:00 2001 From: Khoa Pham Date: Fri, 31 Jul 2026 02:14:09 -0700 Subject: [PATCH] feat: support Kimi Linear PD disaggregation with DCP (#32837) Co-authored-by: Yangmin Li --- .../srt/arg_groups/pd_disaggregation_hook.py | 25 + python/sglang/srt/disaggregation/base/conn.py | 1 + .../sglang/srt/disaggregation/common/conn.py | 62 ++- .../sglang/srt/disaggregation/common/utils.py | 69 +++ python/sglang/srt/disaggregation/fake/conn.py | 1 + .../srt/disaggregation/mooncake/conn.py | 195 ++++++- python/sglang/srt/disaggregation/mori/conn.py | 1 + python/sglang/srt/disaggregation/nixl/conn.py | 182 ++++++- python/sglang/srt/disaggregation/prefill.py | 15 +- python/sglang/srt/disaggregation/utils.py | 25 + .../test_kimi_linear_pd_dcp4.py | 488 ++++++++++++++++++ .../disaggregation/test_nixl_backend_basic.py | 17 + .../unit/server_args/test_server_args.py | 51 ++ 13 files changed, 1084 insertions(+), 48 deletions(-) create mode 100644 test/registered/disaggregation/test_kimi_linear_pd_dcp4.py diff --git a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py index c64d0086a..f4ea1c4f8 100644 --- a/python/sglang/srt/arg_groups/pd_disaggregation_hook.py +++ b/python/sglang/srt/arg_groups/pd_disaggregation_hook.py @@ -26,6 +26,31 @@ def handle_pd_disaggregation(server_args: ServerArgs) -> None: "with MC_FORCE_TCP=1 (TCP transport, no RDMA)" ) + if server_args.disaggregation_mode == "prefill" and server_args.dcp_size > 1: + logger.warning( + "DCP on a PD prefill server is supported when prefill and decode " + "use the same DCP layout, but it usually adds communication " + "overhead without improving prefill performance." + ) + + if server_args.disaggregation_mode == "decode" and server_args.dcp_size > 1: + if server_args.disaggregation_transfer_backend not in ("mooncake", "nixl"): + raise ValueError( + "PD decode DCP requires --disaggregation-transfer-backend " + "mooncake or nixl, got " + f"{server_args.disaggregation_transfer_backend!r}." + ) + if server_args.disaggregation_decode_enable_radix_cache: + raise ValueError( + "PD decode DCP currently requires chunk cache; " + "--disaggregation-decode-enable-radix-cache is not supported." + ) + if server_args.enable_hierarchical_cache: + raise ValueError( + "PD decode DCP currently requires chunk cache; " + "--enable-hierarchical-cache is not supported." + ) + if server_args.disaggregation_mode == "decode": if server_args.disaggregation_decode_enable_radix_cache: if server_args.enable_hisparse: diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 2d028a629..9ce7b1c7e 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -135,6 +135,7 @@ class BaseKVSender(ABC): self, kv_indices: npt.NDArray[np.int32], state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, ): """ Send the kv cache at the given kv indices and the extra cache/state at the given indices to the decoder server. diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index a7bfcd438..f5d6982a4 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -163,10 +163,13 @@ class CommonKVManager(BaseKVManager): self.bootstrap_host = server_args.host self.bootstrap_port = server_args.disaggregation_bootstrap_port self.dist_init_addr = server_args.dist_init_addr - self.attn_tp_size = get_parallel().attn_tp_size - self.attn_tp_rank = get_parallel().attn_tp_rank - self.attn_cp_size = get_parallel().attn_cp_size - self.attn_cp_rank = get_parallel().attn_cp_rank + parallel = get_parallel() + self.attn_tp_size = parallel.attn_tp_size + self.attn_tp_rank = parallel.attn_tp_rank + self.attn_cp_size = parallel.attn_cp_size + self.attn_cp_rank = parallel.attn_cp_rank + self.dcp_size = server_args.dcp_size + self.dcp_rank = parallel.dcp_rank if self.dcp_size > 1 else 0 self.attn_dp_size = get_attention_dp_size() self.attn_dp_rank = get_attention_dp_rank() self.system_dp_size = ( @@ -268,6 +271,39 @@ class CommonKVManager(BaseKVManager): f"Unsupported DisaggregationMode: {self.disaggregation_mode}" ) + def requires_dcp_relayout(self, dst_dcp_size: int, dst_dcp_rank: int) -> bool: + if self.dcp_size == dst_dcp_size: + if self.dcp_rank != dst_dcp_rank: + raise RuntimeError( + "PD peers must connect matching DCP ranks, got " + f"prefill={self.dcp_rank}, decode={dst_dcp_rank}" + ) + return False + + if ( + self.dcp_size == 1 + and dst_dcp_size > 1 + and (self.is_mla_backend or self.is_hybrid_mla_backend) + ): + return True + + raise RuntimeError( + 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]: + page_size = self.kv_args.page_size + 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}" + ) + return src_token_lens + def check_status(self, bootstrap_room: int) -> KVPoll: return self.request_status[bootstrap_room] @@ -469,7 +505,11 @@ class CommonKVManager(BaseKVManager): info: PrefillServerInfo = None try: - url = f"http://{bootstrap_addr}/route?prefill_dp_rank={-1}&prefill_cp_rank={-1}&target_tp_rank={-1}&target_pp_rank={-1}" + url = ( + f"http://{bootstrap_addr}/route?" + f"prefill_dp_rank={-1}&prefill_cp_rank={-1}&" + f"target_tp_rank={-1}&target_pp_rank={-1}" + ) response = requests.get(url, timeout=5) if response.status_code == 200: data = response.json() @@ -501,6 +541,17 @@ class CommonKVManager(BaseKVManager): f"Both servers must use the same --kv-cache-dtype value." ) + if self.dcp_size > 1: + if not (self.is_mla_backend or self.is_hybrid_mla_backend): + raise RuntimeError( + "PD decode DCP requires an MLA or hybrid-MLA KV pool." + ) + if info.attn_cp_size != 1: + raise RuntimeError( + "PD decode DCP currently requires prefill attention CP=1, " + f"got {info.attn_cp_size}." + ) + self._resolve_rank_mapping(info) self.prefill_info_table[bootstrap_addr] = info logger.debug(f"Prefill parallel info for [{bootstrap_addr}]: {info}") @@ -1137,6 +1188,7 @@ class CommonKVSender(BaseKVSender): self, kv_indices: npt.NDArray[np.int32], state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, ): pass diff --git a/python/sglang/srt/disaggregation/common/utils.py b/python/sglang/srt/disaggregation/common/utils.py index e4571f2d9..650f226e0 100644 --- a/python/sglang/srt/disaggregation/common/utils.py +++ b/python/sglang/srt/disaggregation/common/utils.py @@ -25,6 +25,7 @@ class TransferKVChunk: prefill_aux_index: Optional[int] state_indices: Optional[List] chunk_id: Optional[int] = None + num_kv_tokens: Optional[int] = None trace_ctx: Union[TraceReqContext, TraceNullContext] = dataclasses.field( default_factory=TraceNullContext ) @@ -127,3 +128,71 @@ def group_concurrent_contiguous( dst_groups = [g.tolist() for g in dst_groups] return src_groups, dst_groups + + +@dataclasses.dataclass(frozen=True) +class DCPTokenTransferPlan: + src_token_indices: npt.NDArray[np.int64] + dst_token_indices: npt.NDArray[np.int64] + + +def build_dcp_token_transfer_plan( + src_page_indices: npt.NDArray[np.int32], + dst_page_indices: npt.NDArray[np.int32], + *, + physical_page_size: int, + dcp_size: int, + dcp_rank: int, + src_page_offset: int = 0, + decode_prefix_len: int = 0, + num_kv_tokens: Optional[int] = None, +) -> DCPTokenTransferPlan: + 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: + empty = np.empty((0,), dtype=np.int64) + return DCPTokenTransferPlan(empty, 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}" + ) + + 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) diff --git a/python/sglang/srt/disaggregation/fake/conn.py b/python/sglang/srt/disaggregation/fake/conn.py index b74238359..15a76c039 100644 --- a/python/sglang/srt/disaggregation/fake/conn.py +++ b/python/sglang/srt/disaggregation/fake/conn.py @@ -78,6 +78,7 @@ class FakeKVSender(BaseKVSender): self, kv_indices: npt.NDArray[np.int32], state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, ): self.has_sent = True logger.debug( diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index d4489fd03..be0d3f121 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -34,6 +34,7 @@ from sglang.srt.disaggregation.common.utils import ( AuxDataCodec, FastQueue, TransferKVChunk, + build_dcp_token_transfer_plan, group_concurrent_contiguous, pack_int_lists, unpack_int_lists, @@ -45,6 +46,7 @@ from sglang.srt.disaggregation.utils import ( DisaggregationMode, build_transfer_entry_pairs, compute_mamba_state_slice_byte_blocks, + resolve_dcp_dst_entry_indices, ) from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine from sglang.srt.environ import envs @@ -131,6 +133,10 @@ class KVArgsRegisterInfo: dst_state_dim_per_tensor: List[List[int]] dst_kv_layer_ids: List[int] dst_state_layer_ids: List[List[int]] + dst_dcp_size: int = 1 + dst_dcp_rank: int = 0 + requires_dcp_relayout: bool = False + dcp_token_item_lens: Optional[List[int]] = None # Note: always put the staging field at the final (since the staging field is optional and contains multiple inputs) staging: Optional[StagingRegisterInfo] = None @@ -163,6 +169,13 @@ class KVArgsRegisterInfo: if len(msg) > 13 and msg[13] != b"" else [] ), + # msg[14:16] belong to the staging field below; DCP trails it. + dst_dcp_size=( + int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else 1 + ), + dst_dcp_rank=( + int(msg[17].decode("ascii")) if len(msg) > 17 and msg[17] != b"" else 0 + ), # Note: always put the staging field at the final staging=StagingRegisterInfo.from_zmq_fields(msg, 14), ) @@ -749,6 +762,106 @@ class MooncakeKVManager(CommonKVManager): dst_layer_ids=dst_layer_ids, ) + def send_kvcache_dcp( + self, + mooncake_session_id: str, + prefill_kv_indices: npt.NDArray[np.int32], + dst_kv_ptrs: list[int], + dst_kv_indices: npt.NDArray[np.int32], + *, + dcp_token_item_lens: List[int], + dst_dcp_size: int, + dst_dcp_rank: int, + src_page_offset: int, + decode_prefix_len: int, + num_kv_tokens: int, + executor: concurrent.futures.ThreadPoolExecutor, + dst_layer_ids: List[int], + ) -> int: + 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: + dst_indices = resolve_dcp_dst_entry_indices( + src_layer_ids, + dst_layer_ids, + len(self.kv_args.kv_data_ptrs), + len(dst_kv_ptrs), + ) + src_kv_ptrs = self.kv_args.kv_data_ptrs + dst_kv_ptrs = [dst_kv_ptrs[j] for j in dst_indices] + else: + src_kv_ptrs, dst_kv_ptrs, _ = self.get_mla_kv_ptrs_with_pp( + self.kv_args.kv_data_ptrs, + dst_kv_ptrs, + ) + layers_current_pp_stage = len(src_kv_ptrs) + src_groups, dst_groups = group_concurrent_contiguous( + plan.src_token_indices, + plan.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) + ] + + def set_transfer_blocks( + src_ptr: int, dst_ptr: int, token_item_len: int + ) -> List[Tuple[int, int, int]]: + return [ + ( + src_ptr + int(src_group[0]) * token_item_len, + dst_ptr + int(dst_group[0]) * token_item_len, + len(src_group) * token_item_len, + ) + 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: + return self._transfer_data( + mooncake_session_id, + set_transfer_blocks(src_ptr, dst_ptr, token_item_len), + ) + + 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 + ] + for future in concurrent.futures.as_completed(futures): + status = future.result() + if status != 0: + for pending in futures: + pending.cancel() + return status + return 0 + + 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) + ) + return self._transfer_data(mooncake_session_id, transfer_blocks) + def send_kvcache_slice( self, mooncake_session_id: str, @@ -1437,23 +1550,33 @@ class MooncakeKVManager(CommonKVManager): ) break - chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice] - - # NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices - # is mismatched with the dst_kv_indices when page size > 1, this should never happen. - if len(chunked_dst_kv_indice) < len( - kv_chunk.prefill_kv_indices - ): - logger.warning( - f"len(chunked_dst_kv_indice) = {len(chunked_dst_kv_indice)}, len(kv_chunk.prefill_kv_indices) = {len(kv_chunk.prefill_kv_indices)}" - ) - kv_chunk.prefill_kv_indices = kv_chunk.prefill_kv_indices[ - : len(chunked_dst_kv_indice) - ] - target_rank_registration_info: KVArgsRegisterInfo = ( self.decode_kv_args_table[req.mooncake_session_id] ) + is_dcp_transfer = ( + target_rank_registration_info.requires_dcp_relayout + ) + if is_dcp_transfer: + chunked_dst_kv_indice = req.dst_kv_indices + else: + chunked_dst_kv_indice = req.dst_kv_indices[ + kv_chunk.index_slice + ] + + # NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices + # is mismatched with the dst_kv_indices when page size > 1, this should never happen. + if len(chunked_dst_kv_indice) < len( + kv_chunk.prefill_kv_indices + ): + logger.warning( + f"len(chunked_dst_kv_indice) = {len(chunked_dst_kv_indice)}, len(kv_chunk.prefill_kv_indices) = {len(kv_chunk.prefill_kv_indices)}" + ) + kv_chunk.prefill_kv_indices = ( + kv_chunk.prefill_kv_indices[ + : len(chunked_dst_kv_indice) + ] + ) + skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags( target_rank_registration_info ) @@ -1463,6 +1586,27 @@ class MooncakeKVManager(CommonKVManager): or skip_kv ): ret = 0 + elif is_dcp_transfer: + dcp_token_item_lens = ( + target_rank_registration_info.dcp_token_item_lens + ) + assert dcp_token_item_lens is not None + ret = self.send_kvcache_dcp( + req.mooncake_session_id, + kv_chunk.prefill_kv_indices, + target_rank_registration_info.dst_kv_ptrs, + chunked_dst_kv_indice, + dcp_token_item_lens=dcp_token_item_lens, + dst_dcp_size=target_rank_registration_info.dst_dcp_size, + dst_dcp_rank=target_rank_registration_info.dst_dcp_rank, + src_page_offset=kv_chunk.index_slice.start or 0, + decode_prefix_len=req.decode_prefix_len or 0, + num_kv_tokens=kv_chunk.num_kv_tokens, + executor=executor, + dst_layer_ids=( + target_rank_registration_info.dst_kv_layer_ids + ), + ) elif ( self.is_mla_backend or self.is_hybrid_mla_backend @@ -1670,9 +1814,19 @@ class MooncakeKVManager(CommonKVManager): continue mooncake_session_id = waiting_req_bytes[3].decode("ascii") if room == "None": - self.decode_kv_args_table[mooncake_session_id] = ( - KVArgsRegisterInfo.from_zmq(waiting_req_bytes) + decode_kv_args = KVArgsRegisterInfo.from_zmq(waiting_req_bytes) + decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout( + decode_kv_args.dst_dcp_size, + decode_kv_args.dst_dcp_rank, ) + if decode_kv_args.requires_dcp_relayout: + 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) + ) + ) + self.decode_kv_args_table[mooncake_session_id] = decode_kv_args with self.session_lock: if mooncake_session_id in self.failed_sessions: self.failed_sessions.remove(mooncake_session_id) @@ -1786,6 +1940,7 @@ class MooncakeKVManager(CommonKVManager): is_last_chunk: bool, aux_index: Optional[int] = None, state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, trace_ctx: Optional[Union[TraceReqContext, TraceNullContext]] = None, ): assert self.disaggregation_mode == DisaggregationMode.PREFILL @@ -1824,6 +1979,7 @@ class MooncakeKVManager(CommonKVManager): is_last_chunk=is_last_chunk, prefill_aux_index=aux_index, state_indices=state_indices, + num_kv_tokens=num_kv_tokens, trace_ctx=trace_ctx, ) ) @@ -1904,6 +2060,7 @@ class MooncakeKVSender(CommonKVSender): self, kv_indices: npt.NDArray[np.int32], state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, ): kv_indices, index_slice, is_last_chunk, should_skip = ( self._prepare_send_indices(kv_indices, state_indices) @@ -1917,6 +2074,7 @@ class MooncakeKVSender(CommonKVSender): kv_indices, index_slice, False, + num_kv_tokens=num_kv_tokens, trace_ctx=self.trace_ctx.copy_for_thread(), ) else: @@ -1927,6 +2085,7 @@ class MooncakeKVSender(CommonKVSender): True, aux_index=self.aux_index, state_indices=state_indices, + num_kv_tokens=num_kv_tokens, trace_ctx=self.trace_ctx.copy_for_thread(), ) self._record_transfer_indices(kv_indices, state_indices) @@ -2030,6 +2189,8 @@ class MooncakeKVReceiver(CommonKVReceiver): dst_tp_rank = str(tp_rank).encode("ascii") dst_attn_tp_size = str(self.kv_mgr.attn_tp_size).encode("ascii") dst_kv_item_len = str(kv_item_len).encode("ascii") + dst_dcp_size = str(self.kv_mgr.dcp_size).encode("ascii") + dst_dcp_rank = str(self.kv_mgr.dcp_rank).encode("ascii") if ( self.kv_mgr.enable_staging and self.kv_mgr._staging_ctx.allocator is not None @@ -2062,6 +2223,8 @@ class MooncakeKVReceiver(CommonKVReceiver): packed_state_layer_ids, packed_staging_base_ptr, staging_total_size_str, + dst_dcp_size, + dst_dcp_rank, ] ) except zmq.ZMQError: diff --git a/python/sglang/srt/disaggregation/mori/conn.py b/python/sglang/srt/disaggregation/mori/conn.py index ac591179f..7954877f7 100644 --- a/python/sglang/srt/disaggregation/mori/conn.py +++ b/python/sglang/srt/disaggregation/mori/conn.py @@ -1417,6 +1417,7 @@ class MoriKVSender(CommonKVSender): self, kv_indices: npt.NDArray[np.int32], state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, ): kv_indices, index_slice, is_last_chunk, should_skip = ( self._prepare_send_indices(kv_indices, state_indices) diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 831d2f267..ff22ecfde 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -32,6 +32,7 @@ from sglang.srt.disaggregation.common.staging_handler import ( from sglang.srt.disaggregation.common.utils import ( FastQueue, TransferKVChunk, + build_dcp_token_transfer_plan, group_concurrent_contiguous, pack_int_lists, unpack_int_lists, @@ -40,6 +41,7 @@ from sglang.srt.disaggregation.utils import ( DisaggregationMode, build_transfer_entry_pairs, compute_mamba_state_slice_byte_blocks, + resolve_dcp_dst_entry_indices, ) from sglang.srt.environ import envs from sglang.srt.server_args import ServerArgs @@ -158,6 +160,7 @@ class TransferInfo: required_dst_info_num: int dst_state_indices: List[List[int]] decode_prefix_len: Optional[int] = None # for decode radix cache + is_dummy_rank: Optional[bool] = None # NOTE: optional staging field; populated via STAGING_RSP. Keep at the # end so positional construction in from_zmq() continues to work. staging: Optional[StagingTransferInfo] = None @@ -167,6 +170,8 @@ class TransferInfo: # When dst_kv_indices is empty due to a decode-side radix cache # full hit (decode_prefix_len > 0), the transfer is NOT dummy -- # aux/state data still needs to be sent. + if self.is_dummy_rank is not None: + return self.is_dummy_rank if self.dst_kv_indices.size == 0 and self.decode_prefix_len: return False return self.dst_kv_indices.size == 0 @@ -189,6 +194,11 @@ class TransferInfo: decode_prefix_len=( int(msg[8].decode("ascii")) if len(msg) > 8 and msg[8] != b"" else None ), # hacky just add it into the message that will be sent + is_dummy_rank=( + bool(int(msg[9].decode("ascii"))) + if len(msg) > 9 and msg[9] != b"" + else None + ), ) @@ -211,6 +221,11 @@ class KVArgsRegisterInfo: dst_kv_item_len: int dst_kv_item_lens: list[int] dst_kv_layer_ids: list[int] = dataclasses.field(default_factory=list) + dst_dcp_size: int = 1 + dst_dcp_rank: int = 0 + requires_dcp_relayout: bool = False + dcp_token_item_lens: Optional[List[int]] = None + dcp_dst_region_indices: Optional[List[int]] = None dst_num_slots: Optional[int] = None dst_state_item_lens: List[List[int]] = dataclasses.field(default_factory=list) dst_state_dim_per_tensor: List[List[int]] = dataclasses.field(default_factory=list) @@ -277,6 +292,13 @@ class KVArgsRegisterInfo: dst_kv_item_len=dst_kv_item_len, dst_kv_item_lens=dst_kv_item_lens, dst_kv_layer_ids=dst_kv_layer_ids, + # msg[19:21] are the layer-id frames above; DCP trails them. + dst_dcp_size=( + int(msg[21].decode("ascii")) if len(msg) > 21 and msg[21] != b"" else 1 + ), + dst_dcp_rank=( + int(msg[22].decode("ascii")) if len(msg) > 22 and msg[22] != b"" else 0 + ), dst_num_slots=dst_num_slots, dst_state_item_lens=dst_state_item_lens, dst_state_dim_per_tensor=dst_state_dim_per_tensor, @@ -973,6 +995,27 @@ class NixlKVManager(CommonKVManager): assert self.src_mem_kind is not None src_mem_kind = self.src_mem_kind decode_only_spec_dec = n_dst > n_src + + if peer_info.requires_dcp_relayout: + dst_indices = resolve_dcp_dst_entry_indices( + self.kv_args.kv_layer_ids, + peer_info.dst_kv_layer_ids, + n_src, + n_dst, + ) + peer_info.dcp_dst_region_indices = dst_indices + dst_kv_mem_kinds = [peer_info.dst_kv_mem_kinds[j] for j in dst_indices] + dst_kv_item_lens = [peer_info.dst_kv_item_lens[j] for j in dst_indices] + dst_mem_kind = _homogeneous_kv_mem_kind( + dst_kv_mem_kinds, + "PD DCP destination", + ) + 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 + ) + return + if ( self.is_mla_backend or self.is_hybrid_mla_backend @@ -1110,20 +1153,28 @@ class NixlKVManager(CommonKVManager): len(kv_chunk.prefill_kv_indices) > 0 and self.kv_args.kv_data_ptrs ): - chunked_dst_kv_indice = req.dst_kv_indices[kv_chunk.index_slice] - - # NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices - # is mismatched with the dst_kv_indices when page size > 1, this should never happen. - if len(chunked_dst_kv_indice) < len( - kv_chunk.prefill_kv_indices - ): - logger.warning( - f"len(chunked_dst_kv_indice) = {len(chunked_dst_kv_indice)}, len(kv_chunk.prefill_kv_indices) = {len(kv_chunk.prefill_kv_indices)}" - ) - kv_chunk.prefill_kv_indices = kv_chunk.prefill_kv_indices[ - : len(chunked_dst_kv_indice) + is_dcp_transfer = dst_info.requires_dcp_relayout + if is_dcp_transfer: + chunked_dst_kv_indice = req.dst_kv_indices + else: + chunked_dst_kv_indice = req.dst_kv_indices[ + kv_chunk.index_slice ] + # NOTE: This is temporarily a workaround to deal with the case where the prefill_kv_indices + # is mismatched with the dst_kv_indices when page size > 1, this should never happen. + if len(chunked_dst_kv_indice) < len( + kv_chunk.prefill_kv_indices + ): + logger.warning( + f"len(chunked_dst_kv_indice) = {len(chunked_dst_kv_indice)}, len(kv_chunk.prefill_kv_indices) = {len(kv_chunk.prefill_kv_indices)}" + ) + kv_chunk.prefill_kv_indices = ( + kv_chunk.prefill_kv_indices[ + : len(chunked_dst_kv_indice) + ] + ) + src_prefill_kv_indices = kv_chunk.prefill_kv_indices notif = ( @@ -1164,7 +1215,18 @@ class NixlKVManager(CommonKVManager): break if kv_xfer_handle is None: - if ( + if is_dcp_transfer: + kv_xfer_handle = self.send_kvcache_dcp( + req.agent_name, + src_prefill_kv_indices, + dst_info, + chunked_dst_kv_indice, + src_page_offset=kv_chunk.index_slice.start or 0, + decode_prefix_len=req.decode_prefix_len or 0, + num_kv_tokens=kv_chunk.num_kv_tokens, + notif=notif, + ) + elif ( self.is_mla_backend or self.is_hybrid_mla_backend or decode_tp_size == self.attn_tp_size @@ -1228,15 +1290,17 @@ class NixlKVManager(CommonKVManager): if kv_chunk.prefill_aux_index is None: raise RuntimeError("Missing aux index for last chunk") # A no-KV notification still identifies its PP source. + # Empty non-final chunks do not consume chunk IDs, so a + # final no-KV chunk_id equals the prior KV chunk count. + aux_notif = f"{req.room}_aux" if ( len(kv_chunk.prefill_kv_indices) == 0 or not self.kv_args.kv_data_ptrs ): - aux_notif = ( - f"{req.room}_aux_nokv_{self.transfer_source_rank}" + aux_notif += ( + f"_nokv_{self.transfer_source_rank}" + f"_{kv_chunk.chunk_id}" ) - else: - aux_notif = f"{req.room}_aux" aux_xfer_handle = self.send_aux( req.agent_name, kv_chunk.prefill_aux_index, @@ -1354,6 +1418,9 @@ class NixlKVManager(CommonKVManager): if agent_name in self.decode_kv_args_table: logger.info(f"Peer {agent_name} was already registered, ignoring.") return + decode_kv_args.requires_dcp_relayout = self.requires_dcp_relayout( + decode_kv_args.dst_dcp_size, decode_kv_args.dst_dcp_rank + ) self.decode_kv_args_table[agent_name] = decode_kv_args self.agent.add_remote_agent(decode_kv_args.agent_metadata) if self.disaggregation_mode == DisaggregationMode.PREFILL: @@ -1373,6 +1440,7 @@ class NixlKVManager(CommonKVManager): src_mem_kind: str = "VRAM", dst_mem_kind: str = "VRAM", force_flat: bool = False, + bypass_prepped: bool = False, ): """Generic KV cache transfer supporting both MHA and MLA architectures. Used by both send_kvcache and maybe_send_extra. @@ -1382,7 +1450,8 @@ class NixlKVManager(CommonKVManager): index) whose per-layer list must not be half-split into K/V.""" # Prepped path (KV only; state transfers use the non-prepped path below). if ( - src_data_ptrs is self.kv_args.kv_data_ptrs + not bypass_prepped + and src_data_ptrs is self.kv_args.kv_data_ptrs and "" in self.prep_handles and peer_name in self.prep_handles ): @@ -1550,6 +1619,63 @@ class NixlKVManager(CommonKVManager): dst_mem_kind=dst_mem_kind, ) + def send_kvcache_dcp( + self, + peer_name: str, + prefill_kv_indices: npt.NDArray[np.int32], + dst_info: KVArgsRegisterInfo, + dst_kv_indices: npt.NDArray[np.int32], + *, + src_page_offset: int, + decode_prefix_len: int, + num_kv_tokens: int, + notif: str, + ): + if self.src_mem_kind is None: + 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 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_info.dst_dcp_size, + dcp_rank=dst_info.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: + 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 + dst_kv_ptrs = [ + dst_info.dst_kv_ptrs[dst_idx] for dst_idx in dst_info.dcp_dst_region_indices + ] + + # Prepared handles encode page-level offsets, while DCP relayout needs + # flat descriptors for the selected token rows. + return self._send_kvcache_generic( + peer_name=peer_name, + src_data_ptrs=self.kv_args.kv_data_ptrs, + dst_data_ptrs=dst_kv_ptrs, + item_lens=token_item_lens, + prefill_data_indices=plan.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, + ) + def send_kvcache_mixed( self, peer_name: str, @@ -2243,6 +2369,7 @@ class NixlKVManager(CommonKVManager): chunk_id: int, aux_index: Optional[int] = None, state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, ): assert self.disaggregation_mode == DisaggregationMode.PREFILL assert not is_last_chunk or (is_last_chunk and aux_index is not None) @@ -2273,6 +2400,7 @@ class NixlKVManager(CommonKVManager): chunk_id=chunk_id, prefill_aux_index=aux_index, state_indices=state_indices, + num_kv_tokens=num_kv_tokens, ) ) return None @@ -2313,8 +2441,8 @@ class NixlKVManager(CommonKVManager): elif tag == "stg": self._handle_stg_notification(components, room) elif tag == "aux": - # main's "nokv" marker (decode-side radix cache hit): - # mark expected_kvs_per_pp[pp_rank] = 0 for this rank. + # Main's "nokv" marker carries the number of earlier KV + # chunks expected from this PP rank. self._handle_aux_notification(room, components) elif tag == "state": pp_rank = int(components[2]) if len(components) > 2 else 0 @@ -2342,15 +2470,16 @@ class NixlKVManager(CommonKVManager): Notification tag layouts: aux: {room}_aux -> 2 fields - aux (nokv): {room}_aux_nokv_{pp_rank} -> 4 fields - (decode-side radix cache hit; this pp_rank sent - no KV pages, so expected_kvs_per_pp[pp_rank] = 0) + aux (nokv): {room}_aux_nokv_{pp_rank}_{expected} -> 5 fields + (the last chunk had no KV pages for this rank; + `expected` is the number of prior KV chunks) """ self.transfer_statuses[room].received_aux = True # main's "nokv" marker (decode-side radix cache hit, see #19746). if len(components) > 3 and components[2] == "nokv": pp_rank = int(components[3]) - self.transfer_statuses[room].expected_kvs_per_pp[pp_rank] = 0 + expected = int(components[4]) if len(components) > 4 else 0 + self.transfer_statuses[room].expected_kvs_per_pp[pp_rank] = expected if self.transfer_statuses[room].num_pp_ranks_expected is None: self.transfer_statuses[room].num_pp_ranks_expected = ( self.required_prefill_response_num_table.get(room, 1) @@ -2597,6 +2726,7 @@ class NixlKVSender(CommonKVSender): self, kv_indices: npt.NDArray[np.int32], state_indices: Optional[List] = None, + num_kv_tokens: Optional[int] = None, ): if self._send_failed: return @@ -2620,6 +2750,7 @@ class NixlKVSender(CommonKVSender): self.chunk_id, self.aux_index, state_indices, + num_kv_tokens, ) self._record_transfer_indices(kv_indices, state_indices) self.chunk_id += 1 @@ -2737,6 +2868,7 @@ class NixlKVReceiver(CommonKVReceiver): str(self.required_dst_info_num).encode("ascii"), packed_state_indices, str(decode_prefix_len or 0).encode("ascii"), + str(int(is_dummy)).encode("ascii"), ] ) except zmq.ZMQError: @@ -2865,6 +2997,8 @@ class NixlKVReceiver(CommonKVReceiver): packed_kv_item_lens, packed_state_layer_ids, packed_kv_layer_ids, + str(self.kv_mgr.dcp_size).encode("ascii"), + str(self.kv_mgr.dcp_rank).encode("ascii"), ] ) except zmq.ZMQError: diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index b53f6e9b5..87acd602c 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -332,7 +332,8 @@ class PrefillBootstrapQueue: req.start_send_idx = decode_prefix_len num_kv_indices_to_send = num_kv_indices - decode_prefix_len num_pages = kv_to_page_num( - num_kv_indices_to_send, self.token_to_kv_pool.page_size + num_kv_indices_to_send, + self.scheduler.token_to_kv_pool_allocator.page_size, ) req.disagg_kv_sender.init(num_pages, req.metadata_buffer_index) req.pending_bootstrap = False @@ -1082,7 +1083,11 @@ class SchedulerDisaggregationPrefillMixin: cached_end = len(req.prefix_indices) - req.host_hit_length if cached_end <= req.start_send_idx: return - assert cached_end % self.token_to_kv_pool_allocator.page_size == 0 + if cached_end % self.token_to_kv_pool_allocator.page_size != 0: + # DCP radix hits can end on a logical cache-page boundary that is + # not a complete physical DCP page. The regular final send covers + # the full range; only skip this optional early-send optimization. + return # Early-send issues the KV read before this step's forward is enqueued, # but under overlap scheduling the PRIOR step's prefill forward may still # be writing these prefix pages on forward_stream. Record a completion @@ -1235,7 +1240,11 @@ class SchedulerDisaggregationPrefillMixin: page_indices = kv_to_page_indices(kv_indices, page_size) if not req.disagg_kv_sender.should_send_kv_chunk(len(page_indices), last_chunk): return - req.disagg_kv_sender.send(page_indices, state_indices) + req.disagg_kv_sender.send( + page_indices, + state_indices, + num_kv_tokens=end_idx - start_idx, + ) req.start_send_idx = end_idx def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None: diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index d2a5fd1f6..237f293ec 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -934,6 +934,31 @@ def build_transfer_entry_pairs( return [(i, i) for i in range(n_src)] +def resolve_dcp_dst_entry_indices( + src_layer_ids: List[int], + dst_layer_ids: List[int], + n_src: int, + n_dst: int, +) -> List[int]: + """Destination entry index for each local KV entry, for a DCP relayout. + + DCP re-splits the KV by context while PP re-splits it by layer, so the two + index spaces only line up when neither peer is pipelined. Both backends + need the same resolution, hence the shared helper. + """ + if not src_layer_ids and not dst_layer_ids: + # Legacy/non-PP layout. n_dst may exceed n_src when the decode side + # runs speculative decoding and the prefill side does not. + return list(range(n_src)) + # A one-sided mapping is rejected by build_transfer_entry_pairs itself. + return [ + j + for _, j in build_transfer_entry_pairs( + src_layer_ids, dst_layer_ids, n_src, n_dst + ) + ] + + def append_state_component( kv_args: KVArgs, state_type: StateType, diff --git a/test/registered/disaggregation/test_kimi_linear_pd_dcp4.py b/test/registered/disaggregation/test_kimi_linear_pd_dcp4.py new file mode 100644 index 000000000..0d55179ab --- /dev/null +++ b/test/registered/disaggregation/test_kimi_linear_pd_dcp4.py @@ -0,0 +1,488 @@ +import math +import os +import time +import unittest + +import requests +import torch +from transformers import AutoTokenizer + +from sglang.srt.utils import kill_process_tree +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, +) +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + popen_launch_pd_server, + popen_launch_server, +) + +register_cuda_ci(est_time=1200, suite="nightly-8-gpu-b200", nightly=True) + +KIMI_LINEAR_MODEL = "moonshotai/Kimi-Linear-48B-A3B-Instruct" +PHYSICAL_PAGE_SIZE = 64 +CHUNKED_PREFILL_SIZE = 8192 +PARITY_PROMPT_LENGTHS = (63, 64, 65, 255, 256, 257, 8191, 8192, 8193) +LONG_CONTEXT_TOKENS = int(os.environ.get("SGLANG_TEST_PD_LONG_CONTEXT_TOKENS", "32768")) +LONG_CONTEXT_DEPTHS = tuple( + float(value) + for value in os.environ.get("SGLANG_TEST_PD_NIAH_DEPTHS", "0.1,0.5,0.9").split(",") +) +LOGPROB_ATOL = 0.20 +NIAH_KEY = "739391" +KIMI_LINEAR_MAX_CONTEXT = 1_048_576 + + +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) + ) + + +@unittest.skipUnless( + _has_eight_blackwell_gpus(), + "Kimi-Linear PD+DCP acceptance requires eight Blackwell GPUs", +) +class TestKimiLinearPDDCP4(GSM8KMixin, PDDisaggregationServerBase): + model = KIMI_LINEAR_MODEL + gsm8k_score_threshold = 0.88 + gsm8k_num_examples = 200 + gsm8k_num_threads = 4 + gsm8k_num_shots = 5 + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls._collect_monolithic_references() + cls.launch_all() + + @classmethod + def _monolithic_reference_args(cls): + return [ + "--tp-size", + "4", + "--ep-size", + "4", + "--attention-backend", + "tokenspeed_mla", + "--kv-cache-dtype", + "fp8_e4m3", + "--trust-remote-code", + "--random-seed", + "0", + "--dtype", + "bfloat16", + "--page-size", + str(PHYSICAL_PAGE_SIZE), + "--chunked-prefill-size", + str(CHUNKED_PREFILL_SIZE), + "--cuda-graph-max-bs-decode", + "64", + "--cuda-graph-backend-prefill", + "disabled", + "--mem-fraction-static", + "0.80", + ] + + @classmethod + def _tokenize(cls, text: str, *, add_special_tokens: bool): + return cls.tokenizer.encode(text, add_special_tokens=add_special_tokens) + + @classmethod + def _repeat_to_length(cls, token_ids, length: int): + if length == 0: + return [] + assert token_ids, "filler must tokenize to at least one token" + return (token_ids * math.ceil(length / len(token_ids)))[:length] + + @classmethod + def _chat_template_parts(cls): + marker = "SGLANGPDCONTENTSENTINEL739391" + marker_ids = cls._tokenize(marker, add_special_tokens=False) + templated = cls.tokenizer.apply_chat_template( + [{"role": "user", "content": marker}], + tokenize=True, + add_generation_prompt=True, + ) + full_ids = templated if isinstance(templated, list) else templated["input_ids"] + marker_starts = [ + index + for index in range(len(full_ids) - len(marker_ids) + 1) + if full_ids[index : index + len(marker_ids)] == marker_ids + ] + assert len(marker_starts) == 1, ( + "Kimi chat template must contain the user-content marker exactly once, " + f"found offsets={marker_starts}" + ) + marker_start = marker_starts[0] + return ( + full_ids[:marker_start], + full_ids[marker_start + len(marker_ids) :], + ) + + @classmethod + def _build_boundary_prompt(cls, target_length: int): + prefix = cls._tokenize( + "Read the repeated facts carefully.\n", + add_special_tokens=True, + ) + filler = cls._tokenize( + "The sky is blue and the grass is green. ", + add_special_tokens=False, + ) + suffix = cls._tokenize( + "\nAnswer with one word. The capital of France is", + add_special_tokens=False, + ) + middle_length = target_length - len(prefix) - len(suffix) + assert middle_length >= 0, ( + f"target prompt length {target_length} is shorter than fixed prompt " + f"content ({len(prefix) + len(suffix)})" + ) + prompt_ids = prefix + cls._repeat_to_length(filler, middle_length) + suffix + assert len(prompt_ids) == target_length + assert target_length < KIMI_LINEAR_MAX_CONTEXT + return prompt_ids + + @classmethod + def _build_niah_prompt(cls, target_length: int, depth: float): + assert 0 <= depth <= 1 + chat_prefix, chat_suffix = cls._chat_template_parts() + prefix = cls._tokenize( + ( + "You will read a long collection of mundane records. One record " + "contains an access code. Remember that code exactly.\n" + ), + add_special_tokens=False, + ) + filler = cls._tokenize( + ( + "Archive note: the weather was mild, the office lights were on, " + "and no unusual event was reported.\n" + ), + add_special_tokens=False, + ) + needle = cls._tokenize( + f"IMPORTANT RECORD: The access code is {NIAH_KEY}.\n", + add_special_tokens=False, + ) + query = cls._tokenize( + "\nWhat is the access code? Reply with the digits only.", + add_special_tokens=False, + ) + filler_length = ( + target_length + - len(chat_prefix) + - len(prefix) + - len(needle) + - len(query) + - len(chat_suffix) + ) + assert filler_length >= 0 + before_length = int(filler_length * depth) + after_length = filler_length - before_length + prompt_ids = ( + chat_prefix + + prefix + + cls._repeat_to_length(filler, before_length) + + needle + + cls._repeat_to_length(filler, after_length) + + query + + chat_suffix + ) + assert len(prompt_ids) == target_length + assert target_length < KIMI_LINEAR_MAX_CONTEXT + return prompt_ids + + @classmethod + def _generate( + cls, + base_url: str, + input_ids, + *, + max_new_tokens: int, + ignore_eos: bool = True, + ): + response = requests.post( + base_url + "/generate", + json={ + "input_ids": input_ids, + "sampling_params": { + "temperature": 0, + "max_new_tokens": max_new_tokens, + "ignore_eos": ignore_eos, + }, + "return_logprob": True, + }, + timeout=900, + ) + response.raise_for_status() + return response.json() + + @staticmethod + def _flush_cache(base_url: str): + response = requests.post( + base_url + "/flush_cache", + params={"timeout": 30}, + timeout=120, + ) + response.raise_for_status() + + @classmethod + def _flush_pd_caches(cls): + cls._flush_cache(cls.prefill_url) + cls._flush_cache(cls.decode_url) + + @classmethod + def _collect_monolithic_references(cls): + reference_process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5, + other_args=cls._monolithic_reference_args(), + ) + try: + cls.tokenizer = AutoTokenizer.from_pretrained( + cls.model, trust_remote_code=True + ) + cls.parity_prompts = [ + cls._build_boundary_prompt(target_length) + for target_length in PARITY_PROMPT_LENGTHS + ] + cls._flush_cache(cls.base_url) + cls.parity_references = [ + cls._generate(cls.base_url, prompt, max_new_tokens=1) + for prompt in cls.parity_prompts + ] + cls._flush_cache(cls.base_url) + cls.parity_batch_references = cls._generate( + cls.base_url, cls.parity_prompts, max_new_tokens=1 + ) + + cls.niah_prompts = [ + cls._build_niah_prompt(LONG_CONTEXT_TOKENS, needle_depth) + for needle_depth in LONG_CONTEXT_DEPTHS + ] + cls._flush_cache(cls.base_url) + cls.niah_references = [ + cls._generate( + cls.base_url, + prompt, + max_new_tokens=16, + ignore_eos=False, + ) + for prompt in cls.niah_prompts + ] + finally: + kill_process_tree(reference_process.pid, wait_timeout=60) + time.sleep(5) + + @classmethod + def start_prefill(cls): + args = [ + "--trust-remote-code", + "--disaggregation-mode", + "prefill", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp-size", + "4", + "--ep-size", + "4", + "--attention-backend", + "tokenspeed_mla", + "--kv-cache-dtype", + "fp8_e4m3", + "--dtype", + "bfloat16", + "--random-seed", + "0", + "--page-size", + str(PHYSICAL_PAGE_SIZE), + "--chunked-prefill-size", + str(CHUNKED_PREFILL_SIZE), + "--mem-fraction-static", + "0.80", + ] + args += cls.transfer_backend + cls.rdma_devices + cls.process_prefill = popen_launch_pd_server( + cls.model, + cls.prefill_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5, + other_args=args, + ) + + @classmethod + def start_decode(cls): + args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--disaggregation-bootstrap-port", + cls.bootstrap_port, + "--tp-size", + "4", + "--dcp-size", + "4", + "--base-gpu-id", + "4", + "--attention-backend", + "tokenspeed_mla", + "--kv-cache-dtype", + "fp8_e4m3", + "--dcp-comm-backend", + "a2a", + "--dcp-replicate-q-proj", + "--dtype", + "bfloat16", + "--random-seed", + "0", + "--page-size", + str(PHYSICAL_PAGE_SIZE), + "--cuda-graph-max-bs-decode", + "64", + "--cuda-graph-backend-prefill", + "disabled", + "--mem-fraction-static", + "0.80", + ] + args += cls.transfer_backend + cls.rdma_devices + cls.process_decode = popen_launch_pd_server( + cls.model, + cls.decode_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 5, + other_args=args, + ) + + def _assert_output_parity( + self, reference, actual, *, label: str, check_logprobs: bool = True + ): + self.assertEqual( + actual["output_ids"], + reference["output_ids"], + f"{label}: generated token IDs diverged", + ) + if not check_logprobs: + return + + reference_logprobs = reference["meta_info"]["output_token_logprobs"] + actual_logprobs = actual["meta_info"]["output_token_logprobs"] + self.assertEqual( + len(actual_logprobs), + len(reference_logprobs), + f"{label}: output logprob lengths diverged", + ) + + max_delta = 0.0 + for token_index, (expected, observed) in enumerate( + zip(reference_logprobs, actual_logprobs) + ): + self.assertEqual( + int(observed[1]), + int(expected[1]), + f"{label}: logprob token ID diverged at output {token_index}", + ) + max_delta = max(max_delta, abs(float(observed[0]) - float(expected[0]))) + self.assertLessEqual( + max_delta, + LOGPROB_ATOL, + f"{label}: max output logprob delta {max_delta:.6f}", + ) + + def test_monolithic_pd_token_and_logprob_parity(self): + self._flush_pd_caches() + for target_length, prompt, reference in zip( + PARITY_PROMPT_LENGTHS, self.parity_prompts, self.parity_references + ): + with self.subTest(prompt_tokens=target_length): + actual = self._generate(self.base_url, prompt, max_new_tokens=1) + self.assertEqual(actual["meta_info"]["prompt_tokens"], target_length) + self._assert_output_parity( + reference, + actual, + label=f"sequential prompt_tokens={target_length}", + ) + + def test_chunk_and_virtual_page_boundary_batch_parity(self): + self._flush_pd_caches() + actual_outputs = self._generate( + self.base_url, self.parity_prompts, max_new_tokens=1 + ) + self.assertEqual(len(actual_outputs), len(self.parity_batch_references)) + for target_length, reference, actual in zip( + PARITY_PROMPT_LENGTHS, self.parity_batch_references, actual_outputs + ): + with self.subTest(prompt_tokens=target_length): + self.assertEqual(actual["meta_info"]["prompt_tokens"], target_length) + self._assert_output_parity( + reference, + actual, + label=f"batched prompt_tokens={target_length}", + check_logprobs=False, + ) + + def test_long_context_needle_parity(self): + self._flush_pd_caches() + for needle_depth, prompt, reference in zip( + LONG_CONTEXT_DEPTHS, self.niah_prompts, self.niah_references + ): + with self.subTest( + prompt_tokens=LONG_CONTEXT_TOKENS, needle_depth=needle_depth + ): + self.assertIn(NIAH_KEY, reference["text"]) + actual = self._generate( + self.base_url, + prompt, + max_new_tokens=16, + ignore_eos=False, + ) + self.assertIn(NIAH_KEY, actual["text"]) + self._assert_output_parity( + reference, + actual, + label=( + f"niah prompt_tokens={LONG_CONTEXT_TOKENS} " + f"depth={needle_depth}" + ), + ) + + def _assert_batch_completes(self, batch_size: int): + response = requests.post( + self.base_url + "/generate", + json={ + "text": [ + f"Reply with one short word for request {index}: the sky is" + for index in range(batch_size) + ], + "sampling_params": { + "temperature": 0, + "max_new_tokens": 8, + "ignore_eos": True, + }, + }, + timeout=300, + ) + response.raise_for_status() + outputs = response.json() + self.assertIsInstance(outputs, list) + self.assertEqual(len(outputs), batch_size) + self.assertTrue(all(output["text"].strip() for output in outputs)) + + def test_decode_cuda_graph_and_eager_batch(self): + self._assert_batch_completes(2) + self._assert_batch_completes(2) + self._assert_batch_completes(65) + + def test_decode_physical_capacity_sanity(self): + response = requests.get(self.decode_url + "/server_info", timeout=30) + response.raise_for_status() + self.assertGreater(response.json()["max_total_num_tokens"], 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/disaggregation/test_nixl_backend_basic.py b/test/registered/unit/disaggregation/test_nixl_backend_basic.py index 67e727f41..61c8d2f5c 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -215,6 +215,10 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase): b"64", b"DRAM,DRAM", b"".join(struct.pack("Q", item_len) for item_len in [1024, 2048]), + pack_int_lists([[4], [4, 5]], "I"), + b"".join(struct.pack("I", layer_id) for layer_id in [2, 7]), + b"4", + b"3", ] info = KVArgsRegisterInfo.from_zmq(msg) @@ -236,6 +240,10 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase): self.assertEqual(info.dst_kv_mem_kinds, ["DRAM", "DRAM"]) self.assertEqual(info.dst_state_item_lens, state_item_lens) self.assertEqual(info.dst_state_dim_per_tensor, state_dims) + self.assertEqual(info.dst_dcp_size, 4) + self.assertEqual(info.dst_dcp_rank, 3) + self.assertEqual(info.dst_state_layer_ids, [[4], [4, 5]]) + self.assertEqual(info.dst_kv_layer_ids, [2, 7]) self.assertIsNotNone(info.staging) self.assertEqual(info.staging.base_ptr, staging_ptr) self.assertEqual(info.staging.total_size, 1048576) @@ -262,6 +270,8 @@ class TestNixlKVArgsRegisterInfo(CustomTestCase): self.assertEqual(info.dst_state_item_lens, []) self.assertEqual(info.dst_state_dim_per_tensor, []) self.assertEqual(info.dst_kv_item_lens, [256]) + self.assertEqual(info.dst_dcp_size, 1) + self.assertEqual(info.dst_dcp_rank, 0) self.assertIsNone(info.staging) @@ -445,6 +455,12 @@ class TestNixlTransferWorker(CustomTestCase): staging=None, kv_xfer_segments=None, dst_homogeneous_mem_kind="VRAM", + # Non-DCP peer. Without this the worker raises AttributeError + # and lands in the same Failed status the assertions expect, + # so the transfer path would go unexercised. + requires_dcp_relayout=False, + dcp_dst_region_indices=None, + dcp_token_item_lens=None, ) } mgr.req_to_decode_prefix_len = {room: 4} @@ -496,6 +512,7 @@ class TestNixlTransferWorker(CustomTestCase): self.assertNotIn(room, mgr.transfer_infos) self.assertNotIn(room, mgr.req_to_decode_prefix_len) mgr.send_aux.assert_called_once() + self.assertEqual(mgr.send_aux.call_args.args[-1], "21_aux_nokv_0_0") def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved( self, diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index e6bb7b405..648ffe52c 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -8,6 +8,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch import sglang.srt.server_args as server_args_module +from sglang.srt.arg_groups import pd_disaggregation_hook from sglang.srt.arg_groups.speculative_hook import handle_speculative_decoding from sglang.srt.entrypoints.sidecar import ( SGLANG_GRPC_ENDPOINT_ENV, @@ -219,6 +220,56 @@ class TestLoadBalanceMethod(unittest.TestCase): server_args = self._load_balance_args(disaggregation_mode="decode") self.assertEqual(server_args.load_balance_method, "round_robin") + def test_pd_prefill_dcp_warns_about_performance(self): + server_args = ServerArgs( + model_path="dummy", + disaggregation_mode="prefill", + dcp_size=4, + ) + with self.assertLogs(pd_disaggregation_hook.logger, level="WARNING") as logs: + server_args._handle_pd_disaggregation() + self.assertIn("without improving prefill performance", "\n".join(logs.output)) + + def test_pd_decode_dcp_forces_chunk_cache(self): + server_args = self._load_balance_args( + disaggregation_mode="decode", + disaggregation_transfer_backend="mooncake", + dcp_size=4, + ) + self.assertTrue(server_args.disable_radix_cache) + + def test_pd_decode_dcp_rejects_unsupported_transfer_backend(self): + server_args = ServerArgs( + model_path="dummy", + disaggregation_mode="decode", + disaggregation_transfer_backend="fake", + dcp_size=4, + ) + with self.assertRaisesRegex(ValueError, "mooncake or nixl"): + server_args._handle_pd_disaggregation() + + def test_pd_decode_dcp_rejects_radix_cache(self): + server_args = ServerArgs( + model_path="dummy", + disaggregation_mode="decode", + disaggregation_transfer_backend="nixl", + disaggregation_decode_enable_radix_cache=True, + dcp_size=4, + ) + with self.assertRaisesRegex(ValueError, "currently requires chunk cache"): + server_args._handle_pd_disaggregation() + + def test_pd_decode_dcp_rejects_hierarchical_cache(self): + server_args = ServerArgs( + model_path="dummy", + disaggregation_mode="decode", + disaggregation_transfer_backend="nixl", + enable_hierarchical_cache=True, + dcp_size=4, + ) + with self.assertRaisesRegex(ValueError, "--enable-hierarchical-cache"): + server_args._handle_pd_disaggregation() + def test_pd_decode_radix_cache_rejects_hisparse(self): server_args = ServerArgs( model_path="dummy",