[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
@@ -628,7 +628,7 @@ class AscendAttnBackend(AttentionBackend):
if self.forward_metadata.actual_seq_lengths_q is not None: if self.forward_metadata.actual_seq_lengths_q is not None:
actual_seq_qlen = self.forward_metadata.actual_seq_lengths_q actual_seq_qlen = self.forward_metadata.actual_seq_lengths_q
else: 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: else:
if self.forward_metadata.actual_seq_lengths_q is None: if self.forward_metadata.actual_seq_lengths_q is None:
if ( if (
@@ -221,6 +221,7 @@ class NPUMLATokenToKVPool(MLATokenToKVPool):
dtype=self.store_dtype, dtype=self.store_dtype,
device=self.device, device=self.device,
) )
self.index_k_buffer = None
if self.index_head_dim is not None: if self.index_head_dim is not None:
self.index_k_buffer = torch.zeros( self.index_k_buffer = torch.zeros(
( (
@@ -1244,7 +1244,7 @@ class Indexer(MultiPlatformOp):
) )
else: else:
actual_seq_lengths_kv = forward_batch.seq_lens 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: else:
if forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q is None: if forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q is None:
if ( if (
@@ -769,6 +769,15 @@ class MLATokenToKVPoolHost(HostKVCache):
pin_memory=self.pin_memory, pin_memory=self.pin_memory,
allocator=self.allocator, 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, # Return k_buffer to preserve original kv_buffer and data_refs init logic,
# though Ascend doesn't use these parameters. # though Ascend doesn't use these parameters.
return self.k_buffer return self.k_buffer
@@ -844,6 +853,8 @@ class MLATokenToKVPoolHost(HostKVCache):
host_k=self.k_buffer, host_k=self.k_buffer,
device_v=device_pool.v_buffer, device_v=device_pool.v_buffer,
host_v=self.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, page_size=self.page_size,
direction=TransferDirection.H2D, direction=TransferDirection.H2D,
) )
@@ -905,6 +916,8 @@ class MLATokenToKVPoolHost(HostKVCache):
host_k=self.k_buffer, host_k=self.k_buffer,
device_v=device_pool.v_buffer, device_v=device_pool.v_buffer,
host_v=self.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, page_size=self.page_size,
direction=TransferDirection.D2H, direction=TransferDirection.D2H,
) )