Fix disagg speculative decoding with NIXL connector (#30222)

This commit is contained in:
nvjullin
2026-07-06 15:51:33 -07:00
committed by GitHub
parent 1c23954cb9
commit b41552334d
+28 -2
View File
@@ -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,