From 4435a23a513143f04310282d8a4bf4cb4903b891 Mon Sep 17 00:00:00 2001 From: McZyWu Date: Mon, 11 May 2026 09:44:14 +0800 Subject: [PATCH] [NPU]adapt multibatch fia ops (#20177) --- .../npu/attention/ascend_backend.py | 58 ++++++++++++------- 1 file changed, 36 insertions(+), 22 deletions(-) diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index 7ed67370c..074b8b20c 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -1089,34 +1089,48 @@ class AscendAttnBackend(AttentionBackend): return attn_output 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) - attn_output = torch.empty( - (q.size(0), layer.tp_q_head_num, layer.v_head_dim), - device=q.device, - dtype=q.dtype, + num_token_padding = q.shape[0] + if num_token_padding > forward_batch.num_token_non_padded_cpu: + q, k, v = [ + 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( -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: causal = True if (