[NPU] support dsv32 radixcache on ascend (#17964)

This commit is contained in:
khalilzhk
2026-02-03 03:34:12 +08:00
committed by GitHub
parent 677f3c49da
commit b0a6d5244c
4 changed files with 16 additions and 2 deletions

View File

@@ -628,7 +628,7 @@ class AscendAttnBackend(AttentionBackend):
if self.forward_metadata.actual_seq_lengths_q is not None:
actual_seq_qlen = self.forward_metadata.actual_seq_lengths_q
else:
actual_seq_qlen = torch.cumsum(forward_batch.seq_lens, dim=0)
actual_seq_qlen = torch.cumsum(forward_batch.extend_seq_lens, dim=0)
else:
if self.forward_metadata.actual_seq_lengths_q is None:
if (

View File

@@ -221,6 +221,7 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
dtype=self.store_dtype,
device=self.device,
)
self.index_k_buffer = None
if self.index_head_dim is not None:
self.index_k_buffer = torch.zeros(
(

View File

@@ -1244,7 +1244,7 @@ class Indexer(MultiPlatformOp):
)
else:
actual_seq_lengths_kv = forward_batch.seq_lens
actual_seq_lengths_q = forward_batch.seq_lens.cumsum(dim=0)
actual_seq_lengths_q = forward_batch.extend_seq_lens.cumsum(dim=0)
else:
if forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q is None:
if (

View File

@@ -769,6 +769,15 @@ class MLATokenToKVPoolHost(HostKVCache):
pin_memory=self.pin_memory,
allocator=self.allocator,
)
self.index_k_buffer = None
if self.device_pool.index_head_dim is not None:
self.index_k_buffer = alloc_func(
(*base_dims, self.device_pool.index_head_dim),
dtype=self.dtype,
device=self.device,
pin_memory=self.pin_memory,
allocator=self.allocator,
)
# Return k_buffer to preserve original kv_buffer and data_refs init logic,
# though Ascend doesn't use these parameters.
return self.k_buffer
@@ -844,6 +853,8 @@ class MLATokenToKVPoolHost(HostKVCache):
host_k=self.k_buffer,
device_v=device_pool.v_buffer,
host_v=self.v_buffer,
device_index_k=device_pool.index_k_buffer,
host_index_k=self.index_k_buffer,
page_size=self.page_size,
direction=TransferDirection.H2D,
)
@@ -905,6 +916,8 @@ class MLATokenToKVPoolHost(HostKVCache):
host_k=self.k_buffer,
device_v=device_pool.v_buffer,
host_v=self.v_buffer,
device_index_k=device_pool.index_k_buffer,
host_index_k=self.index_k_buffer,
page_size=self.page_size,
direction=TransferDirection.D2H,
)