[NPU]adapt multibatch fia ops (#20177)
This commit is contained in:
@@ -1089,34 +1089,48 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
return attn_output
|
return attn_output
|
||||||
|
|
||||||
if self.use_fia:
|
if self.use_fia:
|
||||||
"""FIA will support multi-bs in the later version of CANN"""
|
"""FIA supports multi-bs in the current version of CANN"""
|
||||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||||
attn_output = torch.empty(
|
num_token_padding = q.shape[0]
|
||||||
(q.size(0), layer.tp_q_head_num, layer.v_head_dim),
|
if num_token_padding > forward_batch.num_token_non_padded_cpu:
|
||||||
device=q.device,
|
q, k, v = [
|
||||||
dtype=q.dtype,
|
data[: forward_batch.num_token_non_padded_cpu]
|
||||||
|
for data in [q, k, v]
|
||||||
|
]
|
||||||
|
attn_output, _ = torch_npu.npu_fused_infer_attention_score(
|
||||||
|
query=q,
|
||||||
|
key=k_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
|
||||||
|
),
|
||||||
|
value=v_cache.view(
|
||||||
|
-1, self.page_size, layer.tp_v_head_num * layer.v_head_dim
|
||||||
|
),
|
||||||
|
block_table=self.forward_metadata.block_tables,
|
||||||
|
block_size=self.page_size,
|
||||||
|
atten_mask=self.fia_mask,
|
||||||
|
input_layout="TND",
|
||||||
|
actual_seq_lengths=self.forward_metadata.seq_lens_list_cumsum,
|
||||||
|
actual_seq_lengths_kv=self.forward_metadata.seq_lens_cpu_int,
|
||||||
|
num_key_value_heads=layer.tp_k_head_num,
|
||||||
|
num_heads=layer.tp_q_head_num,
|
||||||
|
scale=layer.scaling,
|
||||||
|
sparse_mode=3,
|
||||||
)
|
)
|
||||||
q_len_offset = 0
|
|
||||||
for q_len in forward_batch.extend_seq_lens_cpu:
|
|
||||||
attn_output[q_len_offset : q_len_offset + q_len] = (
|
|
||||||
torch.ops.npu.npu_fused_infer_attention_score(
|
|
||||||
q[None, q_len_offset : q_len_offset + q_len],
|
|
||||||
k[None, q_len_offset : q_len_offset + q_len],
|
|
||||||
v[None, q_len_offset : q_len_offset + q_len],
|
|
||||||
num_heads=layer.tp_q_head_num,
|
|
||||||
num_key_value_heads=layer.tp_k_head_num,
|
|
||||||
input_layout="BSND", # todo, TND not supports q_heads!=k_heads
|
|
||||||
atten_mask=self.fia_mask.unsqueeze(0),
|
|
||||||
sparse_mode=3 if q_len != 1 else 0,
|
|
||||||
scale=layer.scaling,
|
|
||||||
next_tokens=0,
|
|
||||||
)[0]
|
|
||||||
)
|
|
||||||
q_len_offset += q_len
|
|
||||||
attn_output = attn_output.view(
|
attn_output = attn_output.view(
|
||||||
-1, layer.tp_q_head_num * layer.v_head_dim
|
-1, layer.tp_q_head_num * layer.v_head_dim
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if num_token_padding != forward_batch.num_token_non_padded_cpu:
|
||||||
|
attn_output = torch.cat(
|
||||||
|
[
|
||||||
|
attn_output,
|
||||||
|
attn_output.new_zeros(
|
||||||
|
num_token_padding - attn_output.shape[0],
|
||||||
|
*attn_output.shape[1:],
|
||||||
|
),
|
||||||
|
],
|
||||||
|
dim=0,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
causal = True
|
causal = True
|
||||||
if (
|
if (
|
||||||
|
|||||||
Reference in New Issue
Block a user