feat: add nsa and swa disagg support with nixl (#18939)
Signed-off-by: Neal Vaidya <nealv@nvidia.com>
This commit is contained in:
@@ -368,51 +368,55 @@ class NixlKVManager(CommonKVManager):
|
|||||||
self.decode_kv_args_table[agent_name] = decode_kv_args
|
self.decode_kv_args_table[agent_name] = decode_kv_args
|
||||||
self.agent.add_remote_agent(decode_kv_args.agent_metadata)
|
self.agent.add_remote_agent(decode_kv_args.agent_metadata)
|
||||||
|
|
||||||
def send_kvcache(
|
def _send_kvcache_generic(
|
||||||
self,
|
self,
|
||||||
peer_name: str,
|
peer_name: str,
|
||||||
prefill_kv_indices: npt.NDArray[np.int32],
|
src_data_ptrs: list[int],
|
||||||
dst_kv_ptrs: list[int],
|
dst_data_ptrs: list[int],
|
||||||
dst_kv_indices: npt.NDArray[np.int32],
|
item_lens: list[int],
|
||||||
|
prefill_data_indices: npt.NDArray[np.int32],
|
||||||
|
dst_data_indices: npt.NDArray[np.int32],
|
||||||
dst_gpu_id: int,
|
dst_gpu_id: int,
|
||||||
notif: str,
|
notif: str,
|
||||||
):
|
):
|
||||||
|
"""Generic KV cache transfer supporting both MHA and MLA architectures.
|
||||||
|
Used by both send_kvcache and maybe_send_extra."""
|
||||||
# group by indices
|
# group by indices
|
||||||
prefill_kv_blocks, dst_kv_blocks = group_concurrent_contiguous(
|
prefill_kv_blocks, dst_kv_blocks = group_concurrent_contiguous(
|
||||||
prefill_kv_indices, dst_kv_indices
|
prefill_data_indices, dst_data_indices
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.debug(f"sending kvcache to {peer_name} with notif {notif}")
|
logger.debug(f"sending kvcache to {peer_name} with notif {notif}")
|
||||||
# Make descs
|
# Make descs
|
||||||
if self.is_mla_backend:
|
if self.is_mla_backend:
|
||||||
src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = (
|
src_kv_ptrs, dst_kv_ptrs, layers_current_pp_stage = (
|
||||||
self.get_mla_kv_ptrs_with_pp(self.kv_args.kv_data_ptrs, dst_kv_ptrs)
|
self.get_mla_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs)
|
||||||
)
|
)
|
||||||
layers_params = [
|
layers_params = [
|
||||||
(
|
(
|
||||||
src_kv_ptrs[layer_id],
|
src_kv_ptrs[layer_id],
|
||||||
dst_kv_ptrs[layer_id],
|
dst_kv_ptrs[layer_id],
|
||||||
self.kv_args.kv_item_lens[layer_id],
|
item_lens[layer_id],
|
||||||
)
|
)
|
||||||
for layer_id in range(layers_current_pp_stage)
|
for layer_id in range(layers_current_pp_stage)
|
||||||
]
|
]
|
||||||
else:
|
else:
|
||||||
src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = (
|
src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = (
|
||||||
self.get_mha_kv_ptrs_with_pp(self.kv_args.kv_data_ptrs, dst_kv_ptrs)
|
self.get_mha_kv_ptrs_with_pp(src_data_ptrs, dst_data_ptrs)
|
||||||
)
|
)
|
||||||
|
|
||||||
layers_params = [
|
layers_params = [
|
||||||
(
|
(
|
||||||
src_k_ptrs[layer_id],
|
src_k_ptrs[layer_id],
|
||||||
dst_k_ptrs[layer_id],
|
dst_k_ptrs[layer_id],
|
||||||
self.kv_args.kv_item_lens[layer_id],
|
item_lens[layer_id],
|
||||||
)
|
)
|
||||||
for layer_id in range(layers_current_pp_stage)
|
for layer_id in range(layers_current_pp_stage)
|
||||||
] + [
|
] + [
|
||||||
(
|
(
|
||||||
src_v_ptrs[layer_id],
|
src_v_ptrs[layer_id],
|
||||||
dst_v_ptrs[layer_id],
|
dst_v_ptrs[layer_id],
|
||||||
self.kv_args.kv_item_lens[layer_id],
|
item_lens[layer_id],
|
||||||
)
|
)
|
||||||
for layer_id in range(layers_current_pp_stage)
|
for layer_id in range(layers_current_pp_stage)
|
||||||
]
|
]
|
||||||
@@ -455,7 +459,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
dst_reqs = make_req_array(dst_addrs, dst_lens, dst_gpu_id)
|
dst_reqs = make_req_array(dst_addrs, dst_lens, dst_gpu_id)
|
||||||
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"len(src_addrs): before group: {len(prefill_kv_indices)}, after group: {len(src_addrs)}"
|
f"len(src_addrs): before group: {len(prefill_data_indices)}, after group: {len(src_addrs)}"
|
||||||
)
|
)
|
||||||
src_descs = self.agent.get_xfer_descs(src_reqs, "VRAM")
|
src_descs = self.agent.get_xfer_descs(src_reqs, "VRAM")
|
||||||
dst_descs = self.agent.get_xfer_descs(dst_reqs, "VRAM")
|
dst_descs = self.agent.get_xfer_descs(dst_reqs, "VRAM")
|
||||||
@@ -474,6 +478,26 @@ class NixlKVManager(CommonKVManager):
|
|||||||
raise Exception("KVSender failed to post transfer")
|
raise Exception("KVSender failed to post transfer")
|
||||||
return xfer_handle
|
return xfer_handle
|
||||||
|
|
||||||
|
def send_kvcache(
|
||||||
|
self,
|
||||||
|
peer_name: str,
|
||||||
|
prefill_kv_indices: npt.NDArray[np.int32],
|
||||||
|
dst_kv_ptrs: list[int],
|
||||||
|
dst_kv_indices: npt.NDArray[np.int32],
|
||||||
|
dst_gpu_id: int,
|
||||||
|
notif: str,
|
||||||
|
):
|
||||||
|
return self._send_kvcache_generic(
|
||||||
|
peer_name=peer_name,
|
||||||
|
src_data_ptrs=self.kv_args.kv_data_ptrs,
|
||||||
|
dst_data_ptrs=dst_kv_ptrs,
|
||||||
|
item_lens=self.kv_args.kv_item_lens,
|
||||||
|
prefill_data_indices=prefill_kv_indices,
|
||||||
|
dst_data_indices=dst_kv_indices,
|
||||||
|
dst_gpu_id=dst_gpu_id,
|
||||||
|
notif=notif,
|
||||||
|
)
|
||||||
|
|
||||||
def send_kvcache_slice(
|
def send_kvcache_slice(
|
||||||
self,
|
self,
|
||||||
peer_name: str,
|
peer_name: str,
|
||||||
@@ -684,6 +708,59 @@ class NixlKVManager(CommonKVManager):
|
|||||||
raise Exception("Failed to post Mamba state transfer")
|
raise Exception("Failed to post Mamba state transfer")
|
||||||
return xfer_handle
|
return xfer_handle
|
||||||
|
|
||||||
|
def maybe_send_extra(
|
||||||
|
self,
|
||||||
|
peer_name: str,
|
||||||
|
prefill_state_indices: List[int],
|
||||||
|
dst_state_data_ptrs: list[int],
|
||||||
|
dst_state_indices: List[int],
|
||||||
|
dst_gpu_id: int,
|
||||||
|
notif: str,
|
||||||
|
decode_tp_size: int,
|
||||||
|
):
|
||||||
|
"""Send state or extra pool data with type-specific handling."""
|
||||||
|
state_type = getattr(self.kv_args, "state_type", "none")
|
||||||
|
|
||||||
|
if state_type == "mamba":
|
||||||
|
if self.attn_tp_size != decode_tp_size:
|
||||||
|
raise RuntimeError(
|
||||||
|
"PD Disaggregation does NOT support PD different TP sizes for hybrid mamba models yet."
|
||||||
|
)
|
||||||
|
return self._send_mamba_state(
|
||||||
|
peer_name,
|
||||||
|
prefill_state_indices,
|
||||||
|
dst_state_data_ptrs,
|
||||||
|
dst_state_indices,
|
||||||
|
dst_gpu_id,
|
||||||
|
notif,
|
||||||
|
)
|
||||||
|
elif state_type in ["swa", "nsa"]:
|
||||||
|
if not self.is_mla_backend and self.attn_tp_size != decode_tp_size:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"PD Disaggregation does NOT support PD different TP sizes for non-MLA {state_type.upper()} hybrid models yet."
|
||||||
|
)
|
||||||
|
if len(prefill_state_indices) != len(dst_state_indices):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"State index length mismatch: prefill={len(prefill_state_indices)}, "
|
||||||
|
f"dst={len(dst_state_indices)}"
|
||||||
|
)
|
||||||
|
return self._send_kvcache_generic(
|
||||||
|
peer_name=peer_name,
|
||||||
|
src_data_ptrs=self.kv_args.state_data_ptrs,
|
||||||
|
dst_data_ptrs=dst_state_data_ptrs,
|
||||||
|
item_lens=self.kv_args.state_item_lens,
|
||||||
|
prefill_data_indices=np.array(prefill_state_indices, dtype=np.int32),
|
||||||
|
dst_data_indices=np.array(dst_state_indices, dtype=np.int32),
|
||||||
|
dst_gpu_id=dst_gpu_id,
|
||||||
|
notif=notif,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
if state_type != "none":
|
||||||
|
raise RuntimeError(
|
||||||
|
f"PD Disaggregation via NIXL does NOT support {state_type} hybrid models yet."
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
def add_transfer_request(
|
def add_transfer_request(
|
||||||
self,
|
self,
|
||||||
bootstrap_room: int,
|
bootstrap_room: int,
|
||||||
@@ -742,26 +819,17 @@ class NixlKVManager(CommonKVManager):
|
|||||||
# Only the last chunk we need to send the aux data.
|
# Only the last chunk we need to send the aux data.
|
||||||
if is_last:
|
if is_last:
|
||||||
if state_indices is not None:
|
if state_indices is not None:
|
||||||
state_type = getattr(self.kv_args, "state_type", "none")
|
dst_info = self.decode_kv_args_table[req.agent_name]
|
||||||
if (
|
state_xfer_handle = self.maybe_send_extra(
|
||||||
self.attn_tp_size
|
req.agent_name,
|
||||||
!= self.decode_kv_args_table[req.agent_name].decode_tp_size
|
state_indices,
|
||||||
):
|
dst_info.dst_state_data_ptrs,
|
||||||
raise RuntimeError(
|
req.dst_state_indices,
|
||||||
"PD Disaggregation does NOT support PD different TP sizes for hybrid mamba models yet."
|
dst_info.gpu_id,
|
||||||
)
|
f"{req.room}_state_{self.kv_args.pp_rank}",
|
||||||
|
decode_tp_size,
|
||||||
if state_type == "mamba":
|
)
|
||||||
state_xfer_handle = self._send_mamba_state(
|
if state_xfer_handle is not None:
|
||||||
req.agent_name,
|
|
||||||
state_indices,
|
|
||||||
self.decode_kv_args_table[
|
|
||||||
req.agent_name
|
|
||||||
].dst_state_data_ptrs,
|
|
||||||
req.dst_state_indices,
|
|
||||||
self.decode_kv_args_table[req.agent_name].gpu_id,
|
|
||||||
f"{req.room}_state_{self.kv_args.pp_rank}",
|
|
||||||
)
|
|
||||||
handles.append(state_xfer_handle)
|
handles.append(state_xfer_handle)
|
||||||
|
|
||||||
assert aux_index is not None
|
assert aux_index is not None
|
||||||
|
|||||||
Reference in New Issue
Block a user