[NPU]adapt multibatch fia ops (#20177)

This commit is contained in:
McZyWu
2026-05-11 09:44:14 +08:00
committed by GitHub
parent 407665a7d4
commit 4435a23a51
@@ -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]
q_len_offset = 0 ]
for q_len in forward_batch.extend_seq_lens_cpu: attn_output, _ = torch_npu.npu_fused_infer_attention_score(
attn_output[q_len_offset : q_len_offset + q_len] = ( query=q,
torch.ops.npu.npu_fused_infer_attention_score( key=k_cache.view(
q[None, q_len_offset : q_len_offset + q_len], -1, self.page_size, layer.tp_k_head_num * layer.qk_head_dim
k[None, q_len_offset : q_len_offset + q_len], ),
v[None, q_len_offset : q_len_offset + q_len], value=v_cache.view(
num_heads=layer.tp_q_head_num, -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_key_value_heads=layer.tp_k_head_num,
input_layout="BSND", # todo, TND not supports q_heads!=k_heads num_heads=layer.tp_q_head_num,
atten_mask=self.fia_mask.unsqueeze(0),
sparse_mode=3 if q_len != 1 else 0,
scale=layer.scaling, scale=layer.scaling,
next_tokens=0, sparse_mode=3,
)[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 (