[NPU] support dsv32 radixcache on ascend (#17964)
This commit is contained in:
@@ -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,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user