diff --git a/python/sglang/srt/disaggregation/nixl/conn.py b/python/sglang/srt/disaggregation/nixl/conn.py index 30965357e..5375b7096 100644 --- a/python/sglang/srt/disaggregation/nixl/conn.py +++ b/python/sglang/srt/disaggregation/nixl/conn.py @@ -907,6 +907,19 @@ class NixlKVManager(CommonKVManager): 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. + n_src = len(self.kv_args.kv_item_lens) + n_dst = len(peer_info.dst_kv_item_lens) + if n_dst < n_src: + raise ValueError( + "NIXL PD transfer: decode registered fewer KV regions " + f"({n_dst}) than prefill ({n_src}); unexpected geometry" + ) + decode_only_spec_dec = n_dst > n_src + if self.is_mla_backend or peer_info.decode_tp_size == self.attn_tp_size: dst_mem_kind = None try: @@ -914,6 +927,11 @@ class NixlKVManager(CommonKVManager): peer_info.dst_kv_mem_kinds, "destination" ) except NotImplementedError: + if decode_only_spec_dec: + raise NotImplementedError( + "NIXL PD transfer does not support HiSparse combined with " + "decode-only speculative decoding." + ) mem_segments = _kv_xfer_mem_segments( self.kv_args.kv_data_mem_kinds, peer_info.dst_kv_mem_kinds ) @@ -922,6 +940,12 @@ class NixlKVManager(CommonKVManager): self._init_mixed_equal_tp_prep_handles(peer_info, mem_segments) return + if decode_only_spec_dec and dst_mem_kind != "VRAM": + raise NotImplementedError( + "NIXL PD transfer does not support HiSparse combined with " + "decode-only speculative decoding." + ) + peer_info.dst_homogeneous_mem_kind = dst_mem_kind # Build the shared src dlist on the first equal-TP/MLA peer; later # peers reuse it. Skipped entirely on heterogeneous-TP-only setups. @@ -937,13 +961,15 @@ class NixlKVManager(CommonKVManager): if peer_info.dst_num_slots is not None else self._num_slots_src ) - dst_kv_item_lens = peer_info.dst_kv_item_lens + + dst_kv_ptrs = peer_info.dst_kv_ptrs[:n_src] + dst_kv_item_lens = peer_info.dst_kv_item_lens[:n_src] dst_kv_data_lens = [ item_len * dst_num_slots for item_len in dst_kv_item_lens ] self._init_equal_tp_prep_handle( peer_info.agent_name, - peer_info.dst_kv_ptrs, + dst_kv_ptrs, peer_info.gpu_id, num_slots=peer_info.dst_num_slots, mem_kind=dst_mem_kind,