已合并
[bugfix][master]同步BSA算子适配aclnn参数校验 #4795
[bugfix][master]同步BSA算子适配aclnn参数校验 #4795
已合并
Sunshine_Youngster创建于 4月22日
3 个文件变更+3-8
@@ -69,11 +69,9 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_block_sparse_attention_backwa
69 const at::IntArrayRef block_shape_value = (block_shape.has_value() && block_shape->size() >= 2)69 const at::IntArrayRef block_shape_value = (block_shape.has_value() && block_shape->size() >= 2)
70 ? *block_shape70 ? *block_shape
71 : at::IntArrayRef(kDefaultBlockShape, 2);71 : at::IntArrayRef(kDefaultBlockShape, 2);
72- const at::IntArrayRef actual_seq_lengths_value = actual_seq_lengths.value_or(at::IntArrayRef{});
73- const at::IntArrayRef actual_seq_lengths_kv_value = actual_seq_lengths_kv.value_or(at::IntArrayRef{});
74 72 
75 // 初始化 aclnn 中的暂不支持参数73 // 初始化 aclnn 中的暂不支持参数
76- const at::Tensor atten_mask = at::Tensor();74+ const at::Tensor atten_mask{nullptr};
77 const int64_t mask_type = 0;75 const int64_t mask_type = 0;
78 const int64_t pre_tokens = 2147483647;76 const int64_t pre_tokens = 2147483647;
79 const int64_t next_tokens = 2147483647;77 const int64_t next_tokens = 2147483647;
@@ -88,7 +86,7 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_block_sparse_attention_backwa
88 d_out, query, key, value,86 d_out, query, key, value,
89 attention_out, softmax_lse,87 attention_out, softmax_lse,
90 block_sparse_mask, atten_mask, block_shape_value,88 block_sparse_mask, atten_mask, block_shape_value,
91- actual_seq_lengths_value, actual_seq_lengths_kv_value,89+ actual_seq_lengths, actual_seq_lengths_kv,
92 q_input_layout_ptr, kv_input_layout_ptr,90 q_input_layout_ptr, kv_input_layout_ptr,
93 num_key_value_heads, mask_type, scale_value,91 num_key_value_heads, mask_type, scale_value,
94 pre_tokens, next_tokens,92 pre_tokens, next_tokens,
@@ -67,8 +67,6 @@ std::tuple<at::Tensor, at::Tensor> npu_block_sparse_attention(
67 67 
68 // 获取参数68 // 获取参数
69 const at::IntArrayRef block_shape_value = block_shape;69 const at::IntArrayRef block_shape_value = block_shape;
70- const at::IntArrayRef actual_seq_lengths_q_value = actual_seq_lengths.value_or(at::IntArrayRef{});
71- const at::IntArrayRef actual_seq_lengths_kv_value = actual_seq_lengths_kv.value_or(at::IntArrayRef{});
72 const int64_t softmax_lse_flag_value = softmax_lse_flag.value_or(0);70 const int64_t softmax_lse_flag_value = softmax_lse_flag.value_or(0);
73 71 
74 // 初始化aclnn中的暂不支持参数72 // 初始化aclnn中的暂不支持参数
@@ -86,7 +84,7 @@ std::tuple<at::Tensor, at::Tensor> npu_block_sparse_attention(
86 // 调用aclnn接口84 // 调用aclnn接口
87 EXEC_NPU_NO_FORMAT_CHECK_CMD(85 EXEC_NPU_NO_FORMAT_CHECK_CMD(
88 aclnnBlockSparseAttention, query, key, value, block_sparse_mask, atten_mask,86 aclnnBlockSparseAttention, query, key, value, block_sparse_mask, atten_mask,
89- block_shape_value, actual_seq_lengths_q_value, actual_seq_lengths_kv_value, block_table,87+ block_shape_value, actual_seq_lengths, actual_seq_lengths_kv, block_table,
90 q_input_layout_ptr, kv_input_layout_ptr, num_key_value_heads, mask_type, scale_value,88 q_input_layout_ptr, kv_input_layout_ptr, num_key_value_heads, mask_type, scale_value,
91 inner_precise, block_size, pre_tokens, next_tokens, softmax_lse_flag_value,89 inner_precise, block_size, pre_tokens, next_tokens, softmax_lse_flag_value,
92 attention_out, softmax_lse_out);90 attention_out, softmax_lse_out);
@@ -189,7 +189,6 @@ class TestNPUBlockSparseAttention(TestCase):
189 query, key, value, block_sparse_mask, block_shape,189 query, key, value, block_sparse_mask, block_shape,
190 q_input_layout="BNSD", kv_input_layout="BNSD",190 q_input_layout="BNSD", kv_input_layout="BNSD",
191 num_key_value_heads=num_kv_heads, scale_value=scale_value, inner_precise=1,191 num_key_value_heads=num_kv_heads, scale_value=scale_value, inner_precise=1,
192- actual_seq_lengths=[S] * B, actual_seq_lengths_kv=[S] * B,
193 softmax_lse_flag=1,192 softmax_lse_flag=1,
194 )193 )
195 194