From 91f386a5b2b37ba11363da8dcbbafeef4877d3c9 Mon Sep 17 00:00:00 2001 From: YAMY <74099316+YAMY1234@users.noreply.github.com> Date: Sat, 25 Jul 2026 13:34:38 -0700 Subject: [PATCH] fix(disagg): support pipeline-parallel hybrid-linear transfer (#32270) --- python/sglang/srt/disaggregation/base/conn.py | 2 + python/sglang/srt/disaggregation/decode.py | 6 + .../srt/disaggregation/mooncake/conn.py | 122 ++++++++++++++--- python/sglang/srt/disaggregation/nixl/conn.py | 128 ++++++++++++++---- python/sglang/srt/disaggregation/prefill.py | 6 + python/sglang/srt/disaggregation/utils.py | 54 ++++++++ python/sglang/srt/mem_cache/memory_pool.py | 21 +++ .../srt/mem_cache/unified_memory_pool.py | 2 + .../srt/model_executor/pool_configurator.py | 17 ++- python/sglang/srt/models/kimi_linear.py | 28 ++-- .../test_disaggregation_kimi_linear.py | 15 +- .../disaggregation/test_nixl_backend_basic.py | 7 +- 12 files changed, 351 insertions(+), 57 deletions(-) diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index 7e68e37e1..2d028a629 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -40,6 +40,7 @@ class KVArgs: kv_data_ptrs: List[int] kv_data_lens: List[int] kv_item_lens: List[int] + kv_layer_ids: List[int] aux_data_ptrs: List[int] aux_data_lens: List[int] aux_item_lens: List[int] @@ -47,6 +48,7 @@ class KVArgs: state_data_ptrs: List[List[int]] state_data_lens: List[List[int]] state_item_lens: List[List[int]] + state_layer_ids: List[List[int]] # Per-tensor TP slice dim, used when prefill/decode attn_tp_size differ. state_dim_per_tensor: List[List[int]] # Number of rows before the slice axis in each per-slot state tensor. diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 676165ba3..00cf04097 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -451,6 +451,12 @@ 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.kv_layer_ids = ( + self.token_to_kv_pool.get_kv_layer_ids() + if self.draft_token_to_kv_pool is None + and hasattr(self.token_to_kv_pool, "get_kv_layer_ids") + else [] + ) if self.transfer_backend == TransferBackend.NIXL: kv_args.kv_data_mem_kinds = kv_data_mem_kinds kv_args.page_size = self.token_to_kv_pool.page_size diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index 182844ce5..a1ccc4bbf 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -43,6 +43,7 @@ from sglang.srt.disaggregation.mooncake.utils import ( ) from sglang.srt.disaggregation.utils import ( DisaggregationMode, + build_transfer_entry_pairs, compute_mamba_state_slice_byte_blocks, ) from sglang.srt.distributed.parallel_state import get_mooncake_transfer_engine @@ -128,6 +129,8 @@ class KVArgsRegisterInfo: # for mamba state different tp slice transfer dst_state_item_lens: List[List[int]] dst_state_dim_per_tensor: List[List[int]] + dst_kv_layer_ids: List[int] + dst_state_layer_ids: List[List[int]] # Note: always put the staging field at the final (since the staging field is optional and contains multiple inputs) staging: Optional[StagingRegisterInfo] = None @@ -150,8 +153,18 @@ class KVArgsRegisterInfo: dst_state_dim_per_tensor=( unpack_int_lists(msg[11], "I") if len(msg) > 11 else [] ), + dst_kv_layer_ids=( + list(struct.unpack(f"{len(msg[12]) // 4}I", msg[12])) + if len(msg) > 12 and msg[12] != b"" + else [] + ), + dst_state_layer_ids=( + unpack_int_lists(msg[13], "I") + if len(msg) > 13 and msg[13] != b"" + else [] + ), # Note: always put the staging field at the final - staging=StagingRegisterInfo.from_zmq_fields(msg, 12), + staging=StagingRegisterInfo.from_zmq_fields(msg, 14), ) @@ -594,6 +607,8 @@ class MooncakeKVManager(CommonKVManager): executor: concurrent.futures.ThreadPoolExecutor, state_type: Optional[StateType] = None, force_flat: bool = False, + src_layer_ids: Optional[List[int]] = None, + dst_layer_ids: Optional[List[int]] = None, ) -> int: """ Generic KV cache transfer supporting both MHA and MLA architectures. @@ -612,17 +627,33 @@ class MooncakeKVManager(CommonKVManager): # Decode pp size should be equal to prefill pp size or 1 if self.is_mla_backend or self.is_hybrid_mla_backend or force_flat: - src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = ( - self.get_mla_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs, state_type) - ) - layers_params = [ - ( - src_kv_ptrs[layer_id], - dst_kv_ptrs[layer_id], - item_lens[layer_id], + # Layer IDs map PP-local buffers to global decode entries. + # Registrations without them retain the existing PP mapping. + if src_layer_ids or dst_layer_ids: + pairs = build_transfer_entry_pairs( + src_layer_ids, + dst_layer_ids, + len(src_data_ptrs), + len(dst_data_ptrs), + allow_positional_fallback=self.pp_size == 1, ) - for layer_id in range(layers_current_pp_stage) - ] + layers_params = [ + (src_data_ptrs[i], dst_data_ptrs[j], item_lens[i]) for i, j in pairs + ] + else: + src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = ( + self.get_mla_kv_ptrs_with_pp( + src_data_ptrs, dst_data_ptrs, state_type + ) + ) + layers_params = [ + ( + src_kv_ptrs[layer_id], + dst_kv_ptrs[layer_id], + item_lens[layer_id], + ) + for layer_id in range(layers_current_pp_stage) + ] else: src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = ( self.get_mha_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs) @@ -704,6 +735,7 @@ class MooncakeKVManager(CommonKVManager): dst_kv_ptrs: list[int], dst_kv_indices: npt.NDArray[np.int32], executor: concurrent.futures.ThreadPoolExecutor, + dst_layer_ids: Optional[List[int]] = None, ): return self._send_kvcache_generic( mooncake_session_id=mooncake_session_id, @@ -713,6 +745,8 @@ class MooncakeKVManager(CommonKVManager): prefill_data_indices=prefill_kv_indices, dst_data_indices=dst_kv_indices, executor=executor, + src_layer_ids=self.kv_args.kv_layer_ids, + dst_layer_ids=dst_layer_ids, ) def send_kvcache_slice( @@ -993,6 +1027,10 @@ class MooncakeKVManager(CommonKVManager): src_slice_outer_counts = ( src_slice_outer_counts[i] if i < len(src_slice_outer_counts) else [] ) + src_state_layer_ids = self.kv_args.state_layer_ids + src_state_layer_ids = ( + src_state_layer_ids[i] if i < len(src_state_layer_ids) else [] + ) if target_rank_registration_info is not None: dst_data_ptrs = ( target_rank_registration_info.dst_state_data_ptrs[i] @@ -1009,8 +1047,14 @@ class MooncakeKVManager(CommonKVManager): if i < len(target_rank_registration_info.dst_state_dim_per_tensor) else [] ) + dst_state_layer_ids = ( + target_rank_registration_info.dst_state_layer_ids[i] + if i < len(target_rank_registration_info.dst_state_layer_ids) + else [] + ) else: dst_data_ptrs, dst_item_lens, dst_dim_per_tensor = [], [], [] + dst_state_layer_ids = [] dst_indices = ( req.dst_state_indices[i] if i < len(req.dst_state_indices) else [] ) @@ -1036,6 +1080,8 @@ class MooncakeKVManager(CommonKVManager): target_rank_registration_info.dst_attn_tp_size, src_conv_shard_groups, src_slice_outer_counts, + src_state_layer_ids, + dst_state_layer_ids, ) or rc ) @@ -1048,6 +1094,8 @@ class MooncakeKVManager(CommonKVManager): src_item_lens, dst_data_ptrs, dst_indices, + src_state_layer_ids, + dst_state_layer_ids, ) or rc ) @@ -1149,11 +1197,21 @@ class MooncakeKVManager(CommonKVManager): src_state_item_lens: list[int], dst_state_data_ptrs: list[int], dst_mamba_index: list, + src_layer_ids: Optional[List[int]] = None, + dst_layer_ids: Optional[List[int]] = None, ): assert len(prefill_mamba_index) == 1, "Mamba should have single state index" transfer_blocks = [] - for i, dst_state_ptr in enumerate(dst_state_data_ptrs): + pairs = build_transfer_entry_pairs( + src_layer_ids or [], + dst_layer_ids or [], + len(src_state_data_ptrs), + len(dst_state_data_ptrs), + allow_positional_fallback=self.pp_size == 1, + ) + for i, j in pairs: + dst_state_ptr = dst_state_data_ptrs[j] length = src_state_item_lens[i] src_addr = src_state_data_ptrs[i] + length * int(prefill_mamba_index[0]) dst_addr = dst_state_ptr + length * int(dst_mamba_index[0]) @@ -1176,6 +1234,8 @@ class MooncakeKVManager(CommonKVManager): dst_attn_tp_size: int, src_state_conv_shard_groups: list = None, src_state_slice_outer_counts: list[int] = None, + src_layer_ids: Optional[List[int]] = None, + dst_layer_ids: Optional[List[int]] = None, ): """Transfer Mamba states with TP slice support. @@ -1205,17 +1265,27 @@ class MooncakeKVManager(CommonKVManager): src_state_item_lens, dst_state_data_ptrs, dst_mamba_index, + src_layer_ids, + dst_layer_ids, ) local_tp_rank_in_group = self.kv_args.engine_rank % self.attn_tp_size dst_tp_rank_in_group = dst_tp_rank % dst_attn_tp_size transfer_blocks = [] - for i, dst_state_ptr in enumerate(dst_state_data_ptrs): + pairs = build_transfer_entry_pairs( + src_layer_ids or [], + dst_layer_ids or [], + len(src_state_data_ptrs), + len(dst_state_data_ptrs), + allow_positional_fallback=self.pp_size == 1, + ) + for i, j in pairs: + dst_state_ptr = dst_state_data_ptrs[j] src_item_len = src_state_item_lens[i] - dst_item_len = dst_state_item_lens[i] + dst_item_len = dst_state_item_lens[j] src_dim = src_state_dim_per_tensor[i] - dst_dim = dst_state_dim_per_tensor[i] + dst_dim = dst_state_dim_per_tensor[j] conv_shard_groups = ( src_state_conv_shard_groups[i] @@ -1370,7 +1440,11 @@ class MooncakeKVManager(CommonKVManager): skip_kv, skip_state = self._get_dsa_cache_transfer_skip_flags( target_rank_registration_info ) - if len(kv_chunk.prefill_kv_indices) == 0 or skip_kv: + if ( + len(kv_chunk.prefill_kv_indices) == 0 + or not self.kv_args.kv_data_ptrs + or skip_kv + ): ret = 0 elif ( self.is_mla_backend @@ -1384,6 +1458,7 @@ class MooncakeKVManager(CommonKVManager): target_rank_registration_info.dst_kv_ptrs, chunked_dst_kv_indice, executor, + target_rank_registration_info.dst_kv_layer_ids, ) elif ( self.enable_staging @@ -1919,9 +1994,20 @@ class MooncakeKVReceiver(CommonKVReceiver): packed_state_dim_per_tensor = pack_int_lists( getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I" ) + packed_state_layer_ids = pack_int_lists( + self.kv_mgr.kv_args.state_layer_ids, "I" + ) + packed_kv_layer_ids = b"".join( + struct.pack("I", layer_id) + for layer_id in self.kv_mgr.kv_args.kv_layer_ids + ) # Note(shangming): No need to add pp rank here since decode pp size should be equal to prefill pp size or 1 tp_rank = self.kv_mgr.kv_args.engine_rank - kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0] + kv_item_len = ( + self.kv_mgr.kv_args.kv_item_lens[0] + if self.kv_mgr.kv_args.kv_item_lens + else 0 + ) 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") @@ -1953,6 +2039,8 @@ class MooncakeKVReceiver(CommonKVReceiver): dst_kv_item_len, packed_state_item_lens, packed_state_dim_per_tensor, + packed_kv_layer_ids, + packed_state_layer_ids, packed_staging_base_ptr, staging_total_size_str, ] diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 0af144c4f..0e1548fe9 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -38,6 +38,7 @@ from sglang.srt.disaggregation.common.utils import ( ) from sglang.srt.disaggregation.utils import ( DisaggregationMode, + build_transfer_entry_pairs, compute_mamba_state_slice_byte_blocks, ) from sglang.srt.environ import envs @@ -209,9 +210,11 @@ class KVArgsRegisterInfo: decode_tp_rank: int dst_kv_item_len: int dst_kv_item_lens: list[int] + dst_kv_layer_ids: list[int] = dataclasses.field(default_factory=list) 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) + dst_state_layer_ids: List[List[int]] = dataclasses.field(default_factory=list) dst_homogeneous_mem_kind: Optional[str] = None kv_xfer_segments: Optional[List[_KVXferPreparedSegment]] = None # Keep last: optional, parsed from a variable-length tail of the ZMQ @@ -249,6 +252,14 @@ class KVArgsRegisterInfo: dst_num_slots = ( int(msg[16].decode("ascii")) if len(msg) > 16 and msg[16] != b"" else None ) + dst_state_layer_ids = ( + unpack_int_lists(msg[19], "I") if len(msg) > 19 and len(msg[19]) > 0 else [] + ) + dst_kv_layer_ids = ( + list(struct.unpack(f"{len(msg[20]) // 4}I", msg[20])) + if len(msg) > 20 and msg[20] != b"" + else [] + ) return cls( room=str(msg[0].decode("ascii")), @@ -265,9 +276,11 @@ class KVArgsRegisterInfo: decode_tp_rank=int(msg[10].decode("ascii")), dst_kv_item_len=dst_kv_item_len, dst_kv_item_lens=dst_kv_item_lens, + dst_kv_layer_ids=dst_kv_layer_ids, dst_num_slots=dst_num_slots, dst_state_item_lens=dst_state_item_lens, dst_state_dim_per_tensor=dst_state_dim_per_tensor, + dst_state_layer_ids=dst_state_layer_ids, staging=StagingRegisterInfo.from_zmq_fields(msg, 14), ) @@ -360,6 +373,9 @@ class NixlKVManager(CommonKVManager): is_mla_backend: Optional[bool] = False, ): super().__init__(args, disaggregation_mode, server_args, is_mla_backend) + self.transfer_source_rank = ( + self.kv_args.pp_rank * self.server_args.tp_size + self.kv_args.engine_rank + ) self.kv_args.kv_data_mem_kinds = _normalize_kv_mem_kinds( getattr(self.kv_args, "kv_data_mem_kinds", None), len(self.kv_args.kv_data_ptrs), @@ -367,6 +383,7 @@ class NixlKVManager(CommonKVManager): self.src_mem_kind = ( _homogeneous_kv_mem_kind(self.kv_args.kv_data_mem_kinds, "source") if disaggregation_mode == DisaggregationMode.PREFILL + and self.kv_args.kv_data_mem_kinds else None ) try: @@ -432,9 +449,10 @@ class NixlKVManager(CommonKVManager): self._num_slots_src: int = 0 if self.disaggregation_mode == DisaggregationMode.PREFILL: - self._num_slots_src = ( - self.kv_args.kv_data_lens[0] // self.kv_args.kv_item_lens[0] - ) + if self.kv_args.kv_item_lens: + self._num_slots_src = ( + self.kv_args.kv_data_lens[0] // self.kv_args.kv_item_lens[0] + ) transfer_queue_size = envs.SGLANG_DISAGGREGATION_QUEUE_SIZE.get() self.transfer_queues: List[FastQueue] = [ FastQueue() for _ in range(transfer_queue_size) @@ -919,9 +937,6 @@ class NixlKVManager(CommonKVManager): peer_info.kv_xfer_segments = prepared_segments def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo): - assert self.src_mem_kind is not None - src_mem_kind = self.src_mem_kind - # If prefill does not run speculative decoding (the usual case), # decode with speculative decoding will have more kv items. # Prefill having more kv items is impossible. @@ -932,6 +947,10 @@ class NixlKVManager(CommonKVManager): "NIXL PD transfer: decode registered fewer KV regions " f"({n_dst}) than prefill ({n_src}); unexpected geometry" ) + if n_src == 0: + return + assert self.src_mem_kind is not None + src_mem_kind = self.src_mem_kind decode_only_spec_dec = n_dst > n_src if ( self.is_mla_backend @@ -979,8 +998,16 @@ class NixlKVManager(CommonKVManager): else self._num_slots_src ) - dst_kv_ptrs = peer_info.dst_kv_ptrs[:n_src] - dst_kv_item_lens = peer_info.dst_kv_item_lens[:n_src] + pairs = build_transfer_entry_pairs( + self.kv_args.kv_layer_ids, + peer_info.dst_kv_layer_ids, + n_src, + n_dst, + allow_positional_fallback=self.pp_size == 1, + ) + dst_indices = [j for _, j in pairs] + dst_kv_ptrs = [peer_info.dst_kv_ptrs[j] for j in dst_indices] + dst_kv_item_lens = [peer_info.dst_kv_item_lens[j] for j in dst_indices] dst_kv_data_lens = [ item_len * dst_num_slots for item_len in dst_kv_item_lens ] @@ -1058,7 +1085,10 @@ class NixlKVManager(CommonKVManager): # Skip KV RDMA transfer when there are no pages to send # (e.g., decode-side radix cache matched the entire prefix). # Aux data is still sent below when is_last_chunk=True. - if len(kv_chunk.prefill_kv_indices) > 0: + if ( + 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 @@ -1077,7 +1107,7 @@ class NixlKVManager(CommonKVManager): notif = ( f"{req.room}_kv_{kv_chunk.chunk_id}" - f"_{int(kv_chunk.is_last_chunk)}_{self.kv_args.engine_rank}" + f"_{int(kv_chunk.is_last_chunk)}_{self.transfer_source_rank}" ) # Decide which kv send path to use: @@ -1163,11 +1193,12 @@ class NixlKVManager(CommonKVManager): dst_info.dst_state_data_ptrs, req.dst_state_indices, dst_info.gpu_id, - f"{req.room}_state_{self.kv_args.engine_rank}", + f"{req.room}_state_{self.transfer_source_rank}", decode_tp_size, decode_tp_rank=dst_info.decode_tp_rank, dst_state_item_lens=dst_info.dst_state_item_lens, dst_state_dim_per_tensor=dst_info.dst_state_dim_per_tensor, + dst_state_layer_ids=dst_info.dst_state_layer_ids, ) handles.extend( h for h in state_xfer_handles if h is not None @@ -1175,12 +1206,13 @@ class NixlKVManager(CommonKVManager): if kv_chunk.prefill_aux_index is None: raise RuntimeError("Missing aux index for last chunk") - # When no KV pages were sent (decode-side cache hit), - # encode pp_rank in aux notif so receiver can mark - # expected_kvs_per_pp[pp_rank] = 0. - if len(kv_chunk.prefill_kv_indices) == 0: + # A no-KV notification still identifies its PP source. + if ( + len(kv_chunk.prefill_kv_indices) == 0 + or not self.kv_args.kv_data_ptrs + ): aux_notif = ( - f"{req.room}_aux_nokv_{self.kv_args.engine_rank}" + f"{req.room}_aux_nokv_{self.transfer_source_rank}" ) else: aux_notif = f"{req.room}_aux" @@ -1267,8 +1299,6 @@ class NixlKVManager(CommonKVManager): f"NIXL memory registration failed for {mem_kind} kv tensors" ) self.kv_descs.append(kv_descs) - if not self.kv_descs: - raise Exception("NIXL memory registration failed for kv tensors") aux_addrs = [] for aux_data_ptr, aux_data_len in zip( self.kv_args.aux_data_ptrs, self.kv_args.aux_data_lens @@ -1766,7 +1796,7 @@ class NixlKVManager(CommonKVManager): notif_tag = ( f"{req.room}_stg_{kv_chunk.chunk_id}_{int(kv_chunk.is_last_chunk)}" - f"_{self.kv_args.engine_rank}_{chunk_idx}" + f"_{self.transfer_source_rank}_{chunk_idx}" f"_{page_start}_{num_pages}_{req.agent_name}" ) handle = self.send_kvcache_staged( @@ -1842,6 +1872,8 @@ class NixlKVManager(CommonKVManager): dst_state_indices: List[int], dst_gpu_id: int, notif: str, + src_layer_ids: list[int] = None, + dst_layer_ids: list[int] = None, ): """Transfer Mamba states via RDMA.""" assert len(prefill_state_indices) == 1, "Mamba should have single state index" @@ -1852,7 +1884,15 @@ class NixlKVManager(CommonKVManager): src_addrs = [] dst_addrs = [] - for i, dst_state_ptr in enumerate(dst_state_data_ptrs): + pairs = build_transfer_entry_pairs( + src_layer_ids or [], + dst_layer_ids or [], + len(src_state_data_ptrs), + len(dst_state_data_ptrs), + allow_positional_fallback=self.pp_size == 1, + ) + for i, j in pairs: + dst_state_ptr = dst_state_data_ptrs[j] length = src_state_item_lens[i] if length == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0: continue @@ -1895,6 +1935,8 @@ class NixlKVManager(CommonKVManager): decode_tp_rank: int, src_state_conv_shard_groups: list = None, src_state_slice_outer_counts: list[int] = None, + src_layer_ids: list[int] = None, + dst_layer_ids: list[int] = None, ): """Transfer Mamba states with TP slice support via RDMA. @@ -1922,6 +1964,8 @@ class NixlKVManager(CommonKVManager): dst_state_indices, dst_gpu_id, notif, + src_layer_ids=src_layer_ids, + dst_layer_ids=dst_layer_ids, ) local_tp_rank_in_group = self.kv_args.engine_rank % self.attn_tp_size @@ -1930,13 +1974,21 @@ class NixlKVManager(CommonKVManager): src_addrs = [] dst_addrs = [] - for i, dst_state_ptr in enumerate(dst_state_data_ptrs): + pairs = build_transfer_entry_pairs( + src_layer_ids or [], + dst_layer_ids or [], + len(src_state_data_ptrs), + len(dst_state_data_ptrs), + allow_positional_fallback=self.pp_size == 1, + ) + for i, j in pairs: + dst_state_ptr = dst_state_data_ptrs[j] src_item_len = src_state_item_lens[i] - dst_item_len = dst_state_item_lens[i] + dst_item_len = dst_state_item_lens[j] if src_item_len == 0 or src_state_data_ptrs[i] == 0 or dst_state_ptr == 0: continue src_dim = src_state_dim_per_tensor[i] - dst_dim = dst_state_dim_per_tensor[i] + dst_dim = dst_state_dim_per_tensor[j] conv_shard_groups = ( src_state_conv_shard_groups[i] @@ -2007,6 +2059,7 @@ class NixlKVManager(CommonKVManager): decode_tp_rank: int = 0, dst_state_item_lens: List[List[int]] | None = None, dst_state_dim_per_tensor: List[List[int]] | None = None, + dst_state_layer_ids: List[List[int]] | None = None, ): """Send state per hybrid component, dispatching by state_type[i].""" state_types = getattr(self.kv_args, "state_types", []) or [] @@ -2021,8 +2074,10 @@ class NixlKVManager(CommonKVManager): src_state_slice_outer_counts = ( getattr(self.kv_args, "state_slice_outer_counts", []) or [] ) + src_state_layer_ids = self.kv_args.state_layer_ids dst_state_item_lens = dst_state_item_lens or [] dst_state_dim_per_tensor = dst_state_dim_per_tensor or [] + dst_state_layer_ids = dst_state_layer_ids or [] handles = [] for i, st in enumerate(state_types): @@ -2046,12 +2101,14 @@ class NixlKVManager(CommonKVManager): if i < len(src_state_slice_outer_counts) else [] ) + src_lids = src_state_layer_ids[i] if i < len(src_state_layer_ids) else [] dst_ptrs = dst_state_data_ptrs[i] if i < len(dst_state_data_ptrs) else [] dst_indices = dst_state_indices[i] if i < len(dst_state_indices) else [] dst_lens = dst_state_item_lens[i] if i < len(dst_state_item_lens) else [] dst_dims = ( dst_state_dim_per_tensor[i] if i < len(dst_state_dim_per_tensor) else [] ) + dst_lids = dst_state_layer_ids[i] if i < len(dst_state_layer_ids) else [] comp_notif = f"{notif}_{i}" if st == StateType.MAMBA: @@ -2072,6 +2129,8 @@ class NixlKVManager(CommonKVManager): decode_tp_rank, src_conv, src_outer_counts, + src_layer_ids=src_lids, + dst_layer_ids=dst_lids, ) else: h = self._send_mamba_state( @@ -2083,6 +2142,8 @@ class NixlKVManager(CommonKVManager): dst_indices, dst_gpu_id, comp_notif, + src_layer_ids=src_lids, + dst_layer_ids=dst_lids, ) elif st in ( StateType.SWA, @@ -2709,6 +2770,10 @@ class NixlKVReceiver(CommonKVReceiver): struct.pack("Q", item_len) for item_len in self.kv_mgr.kv_args.kv_item_lens ) + packed_kv_layer_ids = b"".join( + struct.pack("I", layer_id) + for layer_id in self.kv_mgr.kv_args.kv_layer_ids + ) packed_aux_data_ptrs = b"".join( struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.aux_data_ptrs ) @@ -2721,6 +2786,9 @@ class NixlKVReceiver(CommonKVReceiver): packed_state_dim_per_tensor = pack_int_lists( getattr(self.kv_mgr.kv_args, "state_dim_per_tensor", []) or [], "I" ) + packed_state_layer_ids = pack_int_lists( + self.kv_mgr.kv_args.state_layer_ids, "I" + ) # Include staging allocator metadata if available if ( @@ -2733,10 +2801,12 @@ class NixlKVReceiver(CommonKVReceiver): else: packed_staging_base_ptr = b"" staging_total_size_str = b"" - dst_num_slots = ( - self.kv_mgr.kv_args.kv_data_lens[0] - // self.kv_mgr.kv_args.kv_item_lens[0] - ) + if self.kv_mgr.kv_args.kv_item_lens: + dst_kv_item_len = self.kv_mgr.kv_args.kv_item_lens[0] + dst_num_slots = self.kv_mgr.kv_args.kv_data_lens[0] // dst_kv_item_len + else: + dst_kv_item_len = 0 + dst_num_slots = 0 sock, lock = self._connect_to_bootstrap_server(bootstrap_info) try: @@ -2755,7 +2825,7 @@ class NixlKVReceiver(CommonKVReceiver): str(self.kv_mgr.kv_args.gpu_id).encode("ascii"), str(self.kv_mgr.attn_tp_size).encode("ascii"), str(self.kv_mgr.kv_args.engine_rank).encode("ascii"), - str(self.kv_mgr.kv_args.kv_item_lens[0]).encode("ascii"), + str(dst_kv_item_len).encode("ascii"), packed_state_item_lens, packed_state_dim_per_tensor, packed_staging_base_ptr, @@ -2763,6 +2833,8 @@ class NixlKVReceiver(CommonKVReceiver): str(dst_num_slots).encode("ascii"), packed_kv_data_mem_kinds, packed_kv_item_lens, + packed_state_layer_ids, + packed_kv_layer_ids, ] ) except zmq.ZMQError: diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 1a54a2e04..b17879278 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -219,6 +219,12 @@ 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.kv_layer_ids = ( + self.token_to_kv_pool.get_kv_layer_ids() + if self.draft_token_to_kv_pool is None + and hasattr(self.token_to_kv_pool, "get_kv_layer_ids") + else [] + ) if not self.is_mla_backend: kv_args.kv_head_num = self.token_to_kv_pool.head_num kv_args.total_kv_head_num = ( diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index ab50f0a26..2c47a97eb 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -855,6 +855,53 @@ def compute_mamba_state_slice_byte_blocks( return blocks +def build_transfer_entry_pairs( + src_layer_ids: List[int], + dst_layer_ids: List[int], + n_src: int, + n_dst: int, + allow_positional_fallback: bool = False, +) -> List[Tuple[int, int]]: + """Pair prefill-local transfer entries with decode entries by layer id.""" + if n_src == 0: + return [] + if bool(src_layer_ids) != bool(dst_layer_ids): + if not allow_positional_fallback: + raise RuntimeError( + "Layer metadata must be provided by both PD peers or neither" + ) + src_layer_ids = [] + dst_layer_ids = [] + if src_layer_ids: + if len(src_layer_ids) != n_src or len(dst_layer_ids) != n_dst: + raise RuntimeError( + "Layer metadata length must match transfer entries: " + f"src metadata={len(src_layer_ids)} entries={n_src}, " + f"dst metadata={len(dst_layer_ids)} entries={n_dst}" + ) + # Layer ids can repeat across tensor groups (for example K/V or multiple + # state tensors), so pair occurrences in order rather than by plain lookup. + dst_pos = {} + for j, lid in enumerate(dst_layer_ids): + dst_pos.setdefault(lid, deque()).append(j) + pairs = [] + for i, lid in enumerate(src_layer_ids): + if not dst_pos.get(lid): + raise RuntimeError( + f"Decode peer is missing a transfer entry for model layer {lid}" + ) + pairs.append((i, dst_pos[lid].popleft())) + return pairs + if n_dst < n_src or (n_src != n_dst and not allow_positional_fallback): + # Without layer ids a positional pairing would silently transfer the + # wrong layers (e.g. PP prefill peered with a stale decode server). + raise RuntimeError( + "PP-heterogeneous transfer requires layer ids on " + f"both peers; got src={n_src} dst={n_dst} entries" + ) + return [(i, i) for i in range(n_src)] + + def append_state_component( kv_args: KVArgs, state_type: StateType, @@ -864,6 +911,7 @@ def append_state_component( dim_per_tensor: Optional[List[int]] = None, conv_shard_groups: Optional[List[Optional[List[int]]]] = None, slice_outer_counts: Optional[List[int]] = None, + layer_ids: Optional[List[int]] = None, ) -> None: """Append one state component. Caller orders state_types consistently on prefill and decode sides.""" @@ -874,6 +922,7 @@ def append_state_component( kv_args.state_dim_per_tensor.append(dim_per_tensor or []) kv_args.state_conv_shard_groups.append(conv_shard_groups or []) kv_args.state_slice_outer_counts.append(slice_outer_counts or []) + kv_args.state_layer_ids.append(layer_ids or []) def setup_state_kv_args( @@ -903,6 +952,7 @@ def setup_state_kv_args( kv_args.state_item_lens = [] kv_args.state_dim_per_tensor = [] kv_args.state_slice_outer_counts = [] + kv_args.state_layer_ids = [] kv_args.is_hybrid_mla_backend = False kv_args.state_conv_shard_groups = [] @@ -971,6 +1021,9 @@ def setup_state_kv_args( if hasattr(token_to_kv_pool, "get_state_slice_outer_counts") else None ) + # Global layer ids let the sender pair src/dst entries when the + # prefill PP stage registers only its own subset of mamba layers. + layer_ids = token_to_kv_pool.get_state_layer_ids() append_state_component( kv_args, StateType.MAMBA, @@ -980,6 +1033,7 @@ def setup_state_kv_args( dim, conv_shard_groups, slice_outer_counts, + layer_ids, ) elif isinstance(token_to_kv_pool, (DSATokenToKVPool, NPUMLATokenToKVPool)): if draft_token_to_kv_pool is not None and isinstance( diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index d10115143..44e7c5c3b 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -473,6 +473,7 @@ class MambaPool: enable=enable_memory_saver ) num_mamba_layers = len(mamba_layer_ids) + self.mamba_layer_ids = list(mamba_layer_ids) self.size = size self.device = device @@ -874,6 +875,7 @@ class MambaPool: return ( not _is_npu and len(convs) > 0 + and convs[0].shape[0] > 0 and convs[0].is_cuda and all(c.dtype == torch.bfloat16 and c.is_contiguous() for c in convs) ) @@ -1070,6 +1072,16 @@ class MambaPool: dim_per_tensor += [sliceable_dim] * self.num_mamba_layers return dim_per_tensor + def get_state_layer_ids(self): + """Global model-layer id for each RDMA state entry. + + Aligned element-wise with get_contiguous_buf_infos(), which flattens + the state list tensor-major x layer. Lets PD transfer match entries + by layer id when prefill (PP stage) holds a subset of the mamba layers. + """ + state_tensor_count = sum(1 for _ in self._iter_transfer_state_tensors()) + return list(self.mamba_layer_ids) * state_tensor_count + def get_state_slice_outer_counts(self): """Get the number of rows preceding each tensor's TP slice axis.""" outer_counts = [] @@ -3625,6 +3637,11 @@ class HybridLinearKVPool(KVCache): def get_contiguous_buf_infos(self): return self.full_kv_pool.get_contiguous_buf_infos() + def get_kv_layer_ids(self): + """Global layer ids aligned with the full-attention KV buffers.""" + layer_ids = list(self.full_attention_layer_id_mapping) + return layer_ids if self.use_mla else layer_ids * 2 + def get_state_buf_infos(self): mamba_data_ptrs, mamba_data_lens, mamba_item_lens = ( self.mamba_pool.get_contiguous_buf_infos() @@ -3635,6 +3652,10 @@ class HybridLinearKVPool(KVCache): """Get the sliceable dimension size for each mamba state tensor.""" return self.mamba_pool.get_state_dim_per_tensor() + def get_state_layer_ids(self): + """Global layer id per mamba state entry, aligned with get_state_buf_infos().""" + return self.mamba_pool.get_state_layer_ids() + def get_state_slice_outer_counts(self): """Get the row count preceding each mamba state slice axis.""" return self.mamba_pool.get_state_slice_outer_counts() diff --git a/python/sglang/srt/mem_cache/unified_memory_pool.py b/python/sglang/srt/mem_cache/unified_memory_pool.py index 52e879e54..9e7cbf0df 100644 --- a/python/sglang/srt/mem_cache/unified_memory_pool.py +++ b/python/sglang/srt/mem_cache/unified_memory_pool.py @@ -517,6 +517,8 @@ class UnifiedMambaPool(MambaPool): ): spec = unified_buffer.mamba_spec(sub_pool_name) assert spec.layer_num == len(mamba_layer_ids) + # PP disagg state transfer maps entries by global layer id. + self.mamba_layer_ids = list(mamba_layer_ids) conv_views, temporal_view = unified_buffer.mamba_views_for(sub_pool_name) max_slots = unified_buffer.max_slots(sub_pool_name) diff --git a/python/sglang/srt/model_executor/pool_configurator.py b/python/sglang/srt/model_executor/pool_configurator.py index aad0bf964..9a9b153b5 100644 --- a/python/sglang/srt/model_executor/pool_configurator.py +++ b/python/sglang/srt/model_executor/pool_configurator.py @@ -131,6 +131,17 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): num_layers = kvc.layer_info.num_effective_layers self._cell_size = self._compute_cell_size(kvc, num_layers) + has_kv_on_another_pp_stage = ( + self._cell_size == 0 + and mambaish is not None + and bool(mambaish.full_attention_layer_ids) + and kvc.ps.pp_size > 1 + ) + self._zero_kv_max_tokens = ( + torch.iinfo(torch.int64).max + if has_kv_on_another_pp_stage + else kvc.server_args.max_total_tokens or kvc.model_config.context_len + ) # EAGLE/STANDALONE: scale cell_size to account for draft model KV cache. # Assumes draft and target share the same per-layer KV size (head_dim, @@ -287,7 +298,11 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator): def calculate_pool_sizes( self, available_bytes: int, page_size: int ) -> MemoryPoolConfig: - max_total_num_tokens = available_bytes // self._cell_size + max_total_num_tokens = ( + available_bytes // self._cell_size + if self._cell_size + else self._zero_kv_max_tokens + ) max_total_num_tokens = max_total_num_tokens // page_size * page_size return MemoryPoolConfig(max_total_num_tokens=max_total_num_tokens) diff --git a/python/sglang/srt/models/kimi_linear.py b/python/sglang/srt/models/kimi_linear.py index 7ddc4bfc1..0670f61e3 100644 --- a/python/sglang/srt/models/kimi_linear.py +++ b/python/sglang/srt/models/kimi_linear.py @@ -32,7 +32,7 @@ from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.radix_linear_attention import RadixLinearAttention -from sglang.srt.layers.utils import PPMissingLayer +from sglang.srt.layers.utils import PPMissingLayer, get_layer_id from sglang.srt.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, @@ -628,6 +628,8 @@ class KimiLinearForCausalLM(nn.Module): self.model = KimiLinearModel( config, quant_config, prefix=maybe_prefix(prefix, "model") ) + self.start_layer = self.model.start_layer + self.end_layer = self.model.end_layer self.pp_group = get_pp_group() if self.pp_group.is_last_rank: self.lm_head = ParallelLMHead( @@ -664,6 +666,19 @@ class KimiLinearForCausalLM(nn.Module): else: return hidden_states + def _is_non_local_pp_weight(self, name: str) -> bool: + if self.pp_group.world_size == 1: + return False + + layer_id = get_layer_id(name) + if layer_id is not None: + return not (self.model.start_layer <= layer_id < self.model.end_layer) + if name.startswith("model.embed_tokens."): + return not self.pp_group.is_first_rank + if name.startswith(("model.norm.", "lm_head.")): + return not self.pp_group.is_last_rank + return False + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) @@ -703,6 +718,8 @@ class KimiLinearForCausalLM(nn.Module): for args in weights: name, loaded_weight = args[:2] kwargs = args[2] if len(args) > 2 else {} + if self._is_non_local_pp_weight(name): + continue if "rotary_emb.inv_freq" in name: continue @@ -738,8 +755,6 @@ class KimiLinearForCausalLM(nn.Module): # Skip loading extra bias for GPTQ models. if name.endswith(".bias") and name not in params_dict: continue - # if is_pp_missing_parameter(name, self): - # continue param = params_dict[name] weight_loader = param.weight_loader weight_loader(param, loaded_weight, shard_id) @@ -751,8 +766,6 @@ class KimiLinearForCausalLM(nn.Module): if weight_name not in name: continue name = name.replace(weight_name, param_name) - # if is_pp_missing_parameter(name, self): - # continue param = params_dict[name] weight_loader = param.weight_loader weight_loader( @@ -775,9 +788,6 @@ class KimiLinearForCausalLM(nn.Module): name = maybe_remap_kv_scale_name(name, params_dict) if name is None: continue - # if is_pp_missing_parameter(name, self): - # continue - param = params_dict[name] weight_loader = getattr( param, "weight_loader", default_weight_loader @@ -786,6 +796,8 @@ class KimiLinearForCausalLM(nn.Module): loaded_params.add(name) for layer_id in self.config.full_attention_layer_ids: + if not self.model.start_layer <= layer_id < self.model.end_layer: + continue self_attn = self.model.layers[layer_id].self_attn w_kc, w_vc = self_attn.kv_b_proj.weight.unflatten( 0, (-1, self_attn.qk_nope_head_dim + self_attn.v_head_dim) diff --git a/test/registered/disaggregation/test_disaggregation_kimi_linear.py b/test/registered/disaggregation/test_disaggregation_kimi_linear.py index 7bf447b92..49c93e6ec 100644 --- a/test/registered/disaggregation/test_disaggregation_kimi_linear.py +++ b/test/registered/disaggregation/test_disaggregation_kimi_linear.py @@ -14,7 +14,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=240, stage="base-c", runner_config="4-gpu-h100") +register_cuda_ci(est_time=480, stage="base-c", runner_config="4-gpu-h100") KIMI_LINEAR_MODEL = "yujiepan/kimi-linear-tiny-random" SERVER_ENV = {"SGLANG_BATCH_INVARIANT_OPS_ENABLE_MM_DEEPGEMM": "0"} @@ -38,6 +38,7 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase): prefill_tp_size = 2 decode_tp_size = 1 decode_base_gpu_id = 2 + reference_parallel_args = ["--tp-size", "2"] extra_prefill_args = SERVER_ARGS extra_decode_args = SERVER_ARGS extra_prefill_env = SERVER_ENV @@ -72,7 +73,9 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase): self.model, self.lb_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=["--tp-size", "2", "--trust-remote-code"] + SERVER_ARGS, + other_args=self.reference_parallel_args + + ["--trust-remote-code"] + + SERVER_ARGS, env=SERVER_ENV, ) try: @@ -101,5 +104,13 @@ class TestKimiLinearHeterogeneousTPDisaggregation(PDDisaggregationServerBase): assert_process_healthy(self, "decode", self.process_decode, self.decode_url) +class TestKimiLinearPipelineDisaggregation(TestKimiLinearHeterogeneousTPDisaggregation): + prefill_tp_size = 1 + decode_tp_size = 1 + decode_base_gpu_id = 2 + reference_parallel_args = ["--tp-size", "1", "--pp-size", "2"] + extra_prefill_args = SERVER_ARGS + ["--pp-size", "2"] + + 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 47b369414..8f4220c47 100644 --- a/test/registered/unit/disaggregation/test_nixl_backend_basic.py +++ b/test/registered/unit/disaggregation/test_nixl_backend_basic.py @@ -451,8 +451,10 @@ class TestNixlTransferWorker(CustomTestCase): mgr.enable_staging = False mgr._staging_ctx = None mgr.is_mla_backend = False + mgr.is_hybrid_mla_backend = False mgr.attn_tp_size = 1 - mgr.kv_args = SimpleNamespace(engine_rank=0) + mgr.transfer_source_rank = 0 + mgr.kv_args = SimpleNamespace(engine_rank=0, kv_data_ptrs=[0]) mgr.exceptions = {} mgr.failure_lock = threading.Lock() mgr.failure_records = {} @@ -493,6 +495,7 @@ class TestNixlTransferWorker(CustomTestCase): self.assertEqual(mgr.request_status[room], KVPoll.Failed) self.assertNotIn(room, mgr.transfer_infos) self.assertNotIn(room, mgr.req_to_decode_prefix_len) + mgr.send_aux.assert_called_once() def test_given_non_last_chunk_aborts_mid_transfer_when_worker_finishes_then_failed_status_is_preserved( self, @@ -507,6 +510,7 @@ class TestNixlTransferWorker(CustomTestCase): self.assertEqual(mgr.request_status[room], KVPoll.Failed) self.assertIn(room, mgr.transfer_infos) self.assertIn(room, mgr.req_to_decode_prefix_len) + mgr.send_kvcache.assert_called_once() class TestNixlNotifications(CustomTestCase): @@ -699,6 +703,7 @@ class TestNixlStaging(CustomTestCase): mgr.agent = agent or StagingFakeAgent() mgr.attn_tp_size = 2 mgr.is_mla_backend = False + mgr.transfer_source_rank = 1 mgr.kv_args = SimpleNamespace( gpu_id=1, engine_rank=1,