[npu] [bugfix] Fix HiCache MHA backup for NPU (#34341)

This commit is contained in:
gjsheu
2026-08-11 21:09:45 +08:00
committed by GitHub
parent 3add7e19ff
commit a50ab9cec7
2 changed files with 17 additions and 8 deletions
@@ -109,13 +109,17 @@ class AscendKVManager(MooncakeKVManager):
executor: concurrent.futures.ThreadPoolExecutor, executor: concurrent.futures.ThreadPoolExecutor,
dst_layer_ids: Optional[List[int]] = None, dst_layer_ids: Optional[List[int]] = None,
dst_device_kv_indices: Optional[npt.NDArray[np.int32]] = None, dst_device_kv_indices: Optional[npt.NDArray[np.int32]] = None,
dst_kv_item_len: Optional[int] = None,
dst_attn_tp_size: Optional[int] = None,
): ):
if dst_device_kv_indices is not None: if dst_device_kv_indices is not None:
raise NotImplementedError( raise NotImplementedError(
"Ascend PD transfer does not support HiSparse " "Ascend PD transfer does not support HiSparse "
"destination device KV indices" "destination device KV indices"
) )
self._validate_envelope_kv_layout(
dst_kv_ptrs, dst_kv_item_len, dst_attn_tp_size
)
# 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_kv_indices, dst_kv_indices
+12 -7
View File
@@ -384,13 +384,18 @@ class MHATokenToKVPoolHost(HostKVCache):
def backup_from_device_all_layer( def backup_from_device_all_layer(
self, device_pool, host_indices, device_indices, io_backend self, device_pool, host_indices, device_indices, io_backend
): ):
( if io_backend == "kernel_ascend":
device_k_data_ptrs, # NPU pools use contiguous multi-layer tensors and intentionally do
device_v_data_ptrs, # not build the CUDA-style k_data_ptrs/v_data_ptrs arrays.
device_k_buffers, device_kv_buffers = None
device_v_buffers, else:
) = self._resolve_device_transfer_buffers(device_pool) (
device_kv_buffers = device_k_buffers + device_v_buffers device_k_data_ptrs,
device_v_data_ptrs,
device_k_buffers,
device_v_buffers,
) = self._resolve_device_transfer_buffers(device_pool)
device_kv_buffers = device_k_buffers + device_v_buffers
if io_backend == "kernel": if io_backend == "kernel":
if self.layout == "layer_first": if self.layout == "layer_first":
if self.can_use_jit: if self.can_use_jit: