Fix disagg speculative decoding with NIXL connector (#30222)
This commit is contained in:
@@ -907,6 +907,19 @@ class NixlKVManager(CommonKVManager):
|
|||||||
def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo):
|
def _prepare_payload_xfer(self, peer_info: KVArgsRegisterInfo):
|
||||||
assert self.src_mem_kind is not None
|
assert self.src_mem_kind is not None
|
||||||
src_mem_kind = self.src_mem_kind
|
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:
|
if self.is_mla_backend or peer_info.decode_tp_size == self.attn_tp_size:
|
||||||
dst_mem_kind = None
|
dst_mem_kind = None
|
||||||
try:
|
try:
|
||||||
@@ -914,6 +927,11 @@ class NixlKVManager(CommonKVManager):
|
|||||||
peer_info.dst_kv_mem_kinds, "destination"
|
peer_info.dst_kv_mem_kinds, "destination"
|
||||||
)
|
)
|
||||||
except NotImplementedError:
|
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(
|
mem_segments = _kv_xfer_mem_segments(
|
||||||
self.kv_args.kv_data_mem_kinds, peer_info.dst_kv_mem_kinds
|
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)
|
self._init_mixed_equal_tp_prep_handles(peer_info, mem_segments)
|
||||||
return
|
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
|
peer_info.dst_homogeneous_mem_kind = dst_mem_kind
|
||||||
# Build the shared src dlist on the first equal-TP/MLA peer; later
|
# Build the shared src dlist on the first equal-TP/MLA peer; later
|
||||||
# peers reuse it. Skipped entirely on heterogeneous-TP-only setups.
|
# 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
|
if peer_info.dst_num_slots is not None
|
||||||
else self._num_slots_src
|
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 = [
|
dst_kv_data_lens = [
|
||||||
item_len * dst_num_slots for item_len in dst_kv_item_lens
|
item_len * dst_num_slots for item_len in dst_kv_item_lens
|
||||||
]
|
]
|
||||||
self._init_equal_tp_prep_handle(
|
self._init_equal_tp_prep_handle(
|
||||||
peer_info.agent_name,
|
peer_info.agent_name,
|
||||||
peer_info.dst_kv_ptrs,
|
dst_kv_ptrs,
|
||||||
peer_info.gpu_id,
|
peer_info.gpu_id,
|
||||||
num_slots=peer_info.dst_num_slots,
|
num_slots=peer_info.dst_num_slots,
|
||||||
mem_kind=dst_mem_kind,
|
mem_kind=dst_mem_kind,
|
||||||
|
|||||||
Reference in New Issue
Block a user