From 9bd92ba0f6e8040c2f9fb211a50c5094d4cb5d53 Mon Sep 17 00:00:00 2001 From: xjx471258437 <151611486+xjx471258437@users.noreply.github.com> Date: Sun, 25 Jan 2026 15:34:02 +0800 Subject: [PATCH] Support PD disaggregation with different TP/DP size for Qwen3-Next (#16056) Co-authored-by: xjx392321 --- python/sglang/srt/disaggregation/base/conn.py | 2 + python/sglang/srt/disaggregation/decode.py | 5 + .../srt/disaggregation/mooncake/conn.py | 154 ++++++++++++++++-- python/sglang/srt/disaggregation/prefill.py | 5 + python/sglang/srt/mem_cache/memory_pool.py | 31 ++++ .../test_disaggregation_hybrid_attention.py | 74 +++++++++ 6 files changed, 258 insertions(+), 13 deletions(-) diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index f454dc645..da4629e52 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -24,6 +24,8 @@ class KVArgs: state_data_lens: List[int] state_item_lens: List[int] state_type: str # "none", "mamba", "swa" + # for mamba state different tp slice transfer + state_dim_per_tensor: List[int] # dimension to slice for each state tensor ib_device: str ib_traffic_class: str gpu_id: int diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 4738b032f..9e9633d4d 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -292,6 +292,11 @@ class DecodePreallocQueue: kv_args.state_type = "swa" elif isinstance(self.token_to_kv_pool, HybridLinearKVPool): kv_args.state_type = "mamba" + # Get state dimension info for cross-TP slice transfer + if hasattr(self.token_to_kv_pool, "get_state_dim_per_tensor"): + kv_args.state_dim_per_tensor = ( + self.token_to_kv_pool.get_state_dim_per_tensor() + ) elif isinstance(self.token_to_kv_pool, NSATokenToKVPool): kv_args.state_type = "nsa" else: diff --git a/python/sglang/srt/disaggregation/mooncake/conn.py b/python/sglang/srt/disaggregation/mooncake/conn.py index a0c80e0d1..cf5f64e14 100644 --- a/python/sglang/srt/disaggregation/mooncake/conn.py +++ b/python/sglang/srt/disaggregation/mooncake/conn.py @@ -114,6 +114,9 @@ class KVArgsRegisterInfo: dst_tp_rank: int dst_attn_tp_size: int dst_kv_item_len: int + # for mamba state different tp slice transfer + dst_state_item_lens: list[int] + dst_state_dim_per_tensor: list[int] @classmethod def from_zmq(cls, msg: List[bytes]): @@ -128,6 +131,16 @@ class KVArgsRegisterInfo: dst_tp_rank=int(msg[7].decode("ascii")), dst_attn_tp_size=int(msg[8].decode("ascii")), dst_kv_item_len=int(msg[9].decode("ascii")), + dst_state_item_lens=( + list(struct.unpack(f"{len(msg[10])//4}I", msg[10])) + if len(msg) > 10 and len(msg[10]) > 0 + else [] + ), + dst_state_dim_per_tensor=( + list(struct.unpack(f"{len(msg[11])//4}I", msg[11])) + if len(msg) > 11 and len(msg[11]) > 0 + else [] + ), ) @@ -641,17 +654,42 @@ class MooncakeKVManager(CommonKVManager): prefill_state_indices: list[int], dst_state_data_ptrs: list[int], executor: concurrent.futures.ThreadPoolExecutor, + target_rank_registration_info: Optional[KVArgsRegisterInfo] = None, ): """Send state or extra pool data with type-specific handling.""" state_type = getattr(self.kv_args, "state_type", "none") if state_type == "mamba": - return self._send_mamba_state( - req, - prefill_state_indices, - dst_state_data_ptrs, - ) + # Check if we need slice transfer for different TP sizes + if ( + target_rank_registration_info is not None + and self.attn_tp_size != target_rank_registration_info.dst_attn_tp_size + ): + return self._send_mamba_state_slice( + req, + prefill_state_indices, + dst_state_data_ptrs, + target_rank_registration_info.dst_state_item_lens, + target_rank_registration_info.dst_state_dim_per_tensor, + target_rank_registration_info.dst_tp_rank, + target_rank_registration_info.dst_attn_tp_size, + ) + else: + return self._send_mamba_state( + req, + prefill_state_indices, + dst_state_data_ptrs, + ) elif state_type in ["swa", "nsa"]: + # SWA and NSA hybrid models do not support different TP sizes yet + if ( + target_rank_registration_info is not None + and not self.is_mla_backend + and self.attn_tp_size != target_rank_registration_info.dst_attn_tp_size + ): + raise RuntimeError( + f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet." + ) # Reuse _send_kvcache_generic interface to send extra pool data prefill_state_indices = np.array(prefill_state_indices, dtype=np.int32) dst_state_indices = np.array(req.dst_state_indices, dtype=np.int32) @@ -688,6 +726,90 @@ class MooncakeKVManager(CommonKVManager): return self._transfer_data(req.mooncake_session_id, transfer_blocks) + def _send_mamba_state_slice( + self, + req: TransferInfo, + prefill_mamba_index: list[int], + dst_state_data_ptrs: list[int], + dst_state_item_lens: list[int], + dst_state_dim_per_tensor: list[int], + dst_tp_rank: int, + dst_attn_tp_size: int, + ): + """Transfer Mamba states with TP slice support. + + Mamba state layout: + - conv_state: [num_layers, size+1, conv_dim/tp, conv_kernel-1] + - temporal_state: [num_layers, size+1, num_heads/tp, head_dim, state_size] + + The 3rd dimension is sliced by TP. When prefill and decode have different + attn_tp_size, we need to slice the state accordingly. + """ + logger.warning_once( + "Using Mamba state slice transfer for different TP sizes between prefill and decode. " + f"Prefill attn_tp_size={self.attn_tp_size}, Decode attn_tp_size={dst_attn_tp_size}. " + "Performance may be affected." + ) + assert len(prefill_mamba_index) == 1, "Mamba should have single state index" + + transfer_blocks = [] + prefill_state_data_ptrs = self.kv_args.state_data_ptrs + prefill_state_item_lens = self.kv_args.state_item_lens + src_state_dim_per_tensor = getattr(self.kv_args, "state_dim_per_tensor", []) + + # If no dimension info available, fall back to regular transfer + if not src_state_dim_per_tensor or not dst_state_dim_per_tensor: + return self._send_mamba_state(req, prefill_mamba_index, dst_state_data_ptrs) + + 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 + + for i, dst_state_ptr in enumerate(dst_state_data_ptrs): + src_item_len = prefill_state_item_lens[i] + dst_item_len = dst_state_item_lens[i] + src_dim = src_state_dim_per_tensor[i] + dst_dim = dst_state_dim_per_tensor[i] + + # Calculate bytes per dimension slice + # item_len = dim * trailing_dims_size, so trailing_dims_size = item_len / dim + src_bytes_per_dim = src_item_len // src_dim + dst_bytes_per_dim = dst_item_len // dst_dim + + # Determine slicing parameters based on TP configuration + if self.attn_tp_size > dst_attn_tp_size: + # Multiple prefill ranks send to 1 decode rank + # Each prefill sends all its dims to the appropriate offset in decode + src_dim_start = 0 + num_dims_to_send = src_dim + dst_dim_start = local_tp_rank_in_group * src_dim + else: + # 1 prefill rank sends to multiple decode ranks + # Prefill sends a slice of its dims to each decode rank + src_dim_start = (dst_tp_rank_in_group * dst_dim) % src_dim + num_dims_to_send = dst_dim + dst_dim_start = 0 + + # Calculate byte offsets + src_dim_offset = src_dim_start * src_bytes_per_dim + dst_dim_offset = dst_dim_start * dst_bytes_per_dim + bytes_to_send = num_dims_to_send * src_bytes_per_dim + + # Calculate addresses for this state tensor + src_addr = ( + prefill_state_data_ptrs[i] + + src_item_len * int(prefill_mamba_index[0]) + + src_dim_offset + ) + dst_addr = ( + dst_state_ptr + + dst_item_len * int(req.dst_state_indices[0]) + + dst_dim_offset + ) + + transfer_blocks.append((src_addr, dst_addr, bytes_to_send)) + + return self._transfer_data(req.mooncake_session_id, transfer_blocks) + def sync_status_to_decode_endpoint( self, remote: str, dst_port: int, room: int, status: int, prefill_rank: int ): @@ -798,19 +920,12 @@ class MooncakeKVManager(CommonKVManager): if kv_chunk.is_last: if kv_chunk.state_indices is not None: - if not self.is_mla_backend and ( - self.attn_tp_size - != target_rank_registration_info.dst_attn_tp_size - ): - raise RuntimeError( - f"PD Disaggregation does NOT support PD different TP sizes for non-MLA hybrid models yet." - ) - self.maybe_send_extra( req, kv_chunk.state_indices, target_rank_registration_info.dst_state_data_ptrs, executor, + target_rank_registration_info, ) # Only the last chunk we need to send the aux data @@ -1204,6 +1319,17 @@ class MooncakeKVReceiver(CommonKVReceiver): packed_state_data_ptrs = b"".join( struct.pack("Q", ptr) for ptr in self.kv_mgr.kv_args.state_data_ptrs ) + # Pack state_item_lens and state_dim_per_tensor for mamba state slice transfer + packed_state_item_lens = b"".join( + struct.pack("I", item_len) + for item_len in self.kv_mgr.kv_args.state_item_lens + ) + state_dim_per_tensor = getattr( + self.kv_mgr.kv_args, "state_dim_per_tensor", [] + ) + packed_state_dim_per_tensor = b"".join( + struct.pack("I", dim) for dim in state_dim_per_tensor + ) # 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] @@ -1225,6 +1351,8 @@ class MooncakeKVReceiver(CommonKVReceiver): dst_tp_rank, dst_attn_tp_size, dst_kv_item_len, + packed_state_item_lens, + packed_state_dim_per_tensor, ] ) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 39a824c3a..35762a744 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -161,6 +161,11 @@ class PrefillBootstrapQueue: kv_args.state_type = "swa" elif isinstance(self.token_to_kv_pool, HybridLinearKVPool): kv_args.state_type = "mamba" + # Get state dimension info for cross-TP slice transfer + if hasattr(self.token_to_kv_pool, "get_state_dim_per_tensor"): + kv_args.state_dim_per_tensor = ( + self.token_to_kv_pool.get_state_dim_per_tensor() + ) elif isinstance(self.token_to_kv_pool, NSATokenToKVPool): kv_args.state_type = "nsa" else: diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 3734e902a..0afbb15fd 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -374,6 +374,33 @@ class MambaPool: ] return data_ptrs, data_lens, item_lens + def get_state_dim_per_tensor(self): + """Get the sliceable dimension size for each state tensor. + + For mamba state, the layout is: + - conv_state: [num_layers, size+1, conv_dim/tp, conv_kernel-1] + - temporal_state: [num_layers, size+1, num_heads/tp, head_dim, state_size] + + The 3rd dimension (index 2) is the one that gets sliced by TP. + Returns the size of this dimension for each tensor (repeated for each layer). + """ + state_tensors = [] + for field in vars(self.mamba_cache): + value = getattr(self.mamba_cache, field) + if isinstance(value, list): + state_tensors.extend(value) + else: + state_tensors.append(value) + + dim_per_tensor = [] + for state_tensor in state_tensors: + # state_tensor shape: [num_layers, size+1, sliceable_dim, ...] + # The sliceable dimension is at index 2 (after num_layers and size) + sliceable_dim = state_tensor.shape[2] + # Repeat for each layer since we have per-layer data_ptrs + dim_per_tensor += [sliceable_dim] * self.num_mamba_layers + return dim_per_tensor + class HybridReqToTokenPool(ReqToTokenPool): """A memory pool that maps a request to its token locations.""" @@ -1232,6 +1259,10 @@ class HybridLinearKVPool(KVCache): ) return mamba_data_ptrs, mamba_data_lens, mamba_item_lens + def get_state_dim_per_tensor(self): + """Get the sliceable dimension size for each mamba state tensor.""" + return self.mamba_pool.get_state_dim_per_tensor() + def maybe_get_custom_mem_pool(self): return self.full_kv_pool.maybe_get_custom_mem_pool() diff --git a/test/srt/test_disaggregation_hybrid_attention.py b/test/srt/test_disaggregation_hybrid_attention.py index 276325982..c8154e94d 100644 --- a/test/srt/test_disaggregation_hybrid_attention.py +++ b/test/srt/test_disaggregation_hybrid_attention.py @@ -153,5 +153,79 @@ class TestDisaggregationHybridAttentionMambaExtraBuffer(PDDisaggregationServerBa self.assertGreater(metrics["accuracy"], 0.93) +class TestDisaggregationHybridAttentionMambaDPDecode(PDDisaggregationServerBase): + """Test with prefill tp=2 and decode tp=2/dp=2 with dp-attention enabled.""" + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.model = "Qwen/Qwen3-Next-80B-A3B-Instruct" + + # Non blocking start servers + cls.start_prefill() + cls.start_decode() + + # Block until both + cls.wait_server_ready(cls.prefill_url + "/health") + cls.wait_server_ready(cls.decode_url + "/health") + + cls.launch_lb() + + @classmethod + def start_prefill(cls): + prefill_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "prefill", + "--tp", + "2", + ] + prefill_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, + other_args=prefill_args, + ) + + @classmethod + def start_decode(cls): + decode_args = [ + "--trust-remote-code", + "--disaggregation-mode", + "decode", + "--tp", + "2", + "--dp", + "2", + "--enable-dp-attention", + "--enable-dp-lm-head", + "--base-gpu-id", + "2", + ] + decode_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, + other_args=decode_args, + ) + + def test_gsm8k(self): + args = SimpleNamespace( + num_shots=5, + data_path=None, + num_questions=200, + max_new_tokens=512, + parallel=128, + host=f"http://{self.base_host}", + port=int(self.lb_port), + ) + metrics = run_eval_few_shot_gsm8k(args) + print(f"Evaluation metrics: {metrics}") + + self.assertGreater(metrics["accuracy"], 0.93) + + if __name__ == "__main__": unittest.main()