@@ -53,37 +53,28 @@ def vision_fusion_attention(
5353 Attention output with same shape as ``q``.
5454 """
5555 import torch_npu
56- from torchair .ops import npu_fused_infer_attention_score
57-
58- if torch .compiler .is_compiling ():
59- actual_seq_lengths = torch .tensor (
60- actual_seq_qlen , dtype = torch .int64 , device = q .device
61- )
62- actual_seq_lengths_kv = torch .tensor (
63- actual_seq_kvlen , dtype = torch .int64 , device = q .device
64- )
65- return npu_fused_infer_attention_score (
66- q ,
67- k ,
68- v ,
69- actual_seq_lengths = actual_seq_lengths ,
70- actual_seq_lengths_kv = actual_seq_lengths_kv ,
71- num_heads = num_heads ,
72- scale = scale ,
73- input_layout = input_layout ,
74- )[0 ]
75- else :
76- return torch_npu .npu_fusion_attention (
77- q ,
78- k ,
79- v ,
80- actual_seq_qlen = actual_seq_qlen ,
81- actual_seq_kvlen = actual_seq_kvlen ,
82- head_num = num_heads ,
83- scale = scale ,
84- input_layout = input_layout ,
85- )[0 ]
86-
56+ return torch_npu .npu_fused_infer_attention_score (
57+ q ,
58+ k ,
59+ v ,
60+ actual_seq_lengths = actual_seq_qlen ,
61+ actual_seq_lengths_kv = actual_seq_kvlen ,
62+ num_heads = num_heads ,
63+ scale = scale ,
64+ input_layout = input_layout ,
65+ )[0 ]
66+ '''
67+ return torch_npu.npu_fusion_attention(
68+ q,
69+ k,
70+ v,
71+ actual_seq_qlen=actual_seq_qlen,
72+ actual_seq_kvlen=actual_seq_kvlen,
73+ head_num=num_heads,
74+ scale=scale,
75+ input_layout=input_layout,
76+ )[0]
77+ '''
8778
8879__all__ = [
8980 "reshape_paged_cache" ,
0 commit comments