已合并
[fix] Refine shape inference for optional outputs of LI and SFA operators. #4400
[fix] Refine shape inference for optional outputs of LI and SFA operators. #4400
已合并
zzzyh22创建于 3月9日
2 个文件变更+19-8
Mop_plugin/ops/opapi/LightningIndexerKernelNpuOpApi.cpp+9-3
@@ -23,12 +23,13 @@ namespace op_api {
23const static int64_t DIM_0 = 0;23const static int64_t DIM_0 = 0;
24const static int64_t DIM_1 = 1;24const static int64_t DIM_1 = 1;
25const static int64_t DIM_2 = 2;25const static int64_t DIM_2 = 2;
26+const static int64_t DIM_3 = 3;
26using namespace at_npu::native;27using namespace at_npu::native;
27using npu_preparation = at_npu::native::OpPreparation;28using npu_preparation = at_npu::native::OpPreparation;
28 29 
29std::tuple<at::Tensor, at::Tensor> construct_lightning_indexer_output_tensor(const at::Tensor& query,30std::tuple<at::Tensor, at::Tensor> construct_lightning_indexer_output_tensor(const at::Tensor& query,
30 const at::Tensor& key, const c10::optional<at::Tensor> &actual_seq_lengths_query, int64_t sparse_count,31 const at::Tensor& key, const c10::optional<at::Tensor> &actual_seq_lengths_query, int64_t sparse_count,
31- std::string query_layout_str, std::string key_layout_str)32+ std::string query_layout_str, std::string key_layout_str, bool return_value)
32{33{
33 at::SmallVector<int64_t, SIZE> output_size;34 at::SmallVector<int64_t, SIZE> output_size;
34 for (size_t i = 0; i < query.sizes().size(); i++) {35 for (size_t i = 0; i < query.sizes().size(); i++) {
@@ -48,7 +49,12 @@ std::tuple<at::Tensor, at::Tensor> construct_lightning_indexer_output_tensor(con
48 output_size = {query.size(DIM_0), key.size(n_dim_index), sparse_count};49 output_size = {query.size(DIM_0), key.size(n_dim_index), sparse_count};
49 }50 }
50 at::Tensor sparse_indices_out = npu_preparation::apply_tensor_without_format(output_size, at::kInt);51 at::Tensor sparse_indices_out = npu_preparation::apply_tensor_without_format(output_size, at::kInt);
51- at::Tensor sparse_values_out = npu_preparation::apply_tensor_without_format(output_size, query.dtype());52+ at::Tensor sparse_values_out;
53+ if (return_value) {
54+ sparse_values_out = npu_preparation::apply_tensor_without_format(output_size, query.dtype());
55+ } else {
56+ sparse_values_out = npu_preparation::apply_tensor_without_format({0}, query.dtype());
57+ }
52 58 
53 return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);59 return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out);
54}60}
@@ -69,7 +75,7 @@ std::tuple<at::Tensor, at::Tensor> npu_lightning_indexer(
69 75 
70 // construct the output tensor76 // construct the output tensor
71 std::tuple<at::Tensor, at::Tensor> lightning_indexer_output = op_api::construct_lightning_indexer_output_tensor(77 std::tuple<at::Tensor, at::Tensor> lightning_indexer_output = op_api::construct_lightning_indexer_output_tensor(
72- query, key, actual_seq_lengths_query, sparse_count, query_layout_str, key_layout_str);78+ query, key, actual_seq_lengths_query, sparse_count, query_layout_str, key_layout_str, return_value);
73 at::Tensor sparse_indices_out = std::get<0>(lightning_indexer_output);79 at::Tensor sparse_indices_out = std::get<0>(lightning_indexer_output);
74 at::Tensor sparse_values_out = std::get<1>(lightning_indexer_output);80 at::Tensor sparse_values_out = std::get<1>(lightning_indexer_output);
75 // convert str81 // convert str
Mop_plugin/ops/opapi/SparseFlashAttentionKernelNpuOpApi.cpp+10-5
@@ -83,12 +83,17 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_sparse_flash_attention(
83 at::Tensor softmax_sum;83 at::Tensor softmax_sum;
84 at::SmallVector<int64_t, SIZE> softmax_max_size;84 at::SmallVector<int64_t, SIZE> softmax_max_size;
85 at::SmallVector<int64_t, SIZE> softmax_sum_size;85 at::SmallVector<int64_t, SIZE> softmax_sum_size;
86- if (query.dim() == DIM_3) {86+ if (return_softmax_lse) {
87- softmax_max_size = {key.size(1), query.size(0), query.size(1) / key.size(1)};87+ if (query.dim() == DIM_3) {
88- softmax_sum_size = {key.size(1), query.size(0), query.size(1) / key.size(1)};88+ softmax_max_size = {key.size(1), query.size(0), query.size(1) / key.size(1)};
89+ softmax_sum_size = {key.size(1), query.size(0), query.size(1) / key.size(1)};
90+ } else {
91+ softmax_max_size = {query.size(0), key.size(2), query.size(1), query.size(2) / key.size(2)};
92+ softmax_sum_size = {query.size(0), key.size(2), query.size(1), query.size(2) / key.size(2)};
93+ }
89 } else {94 } else {
90- softmax_max_size = {query.size(0), key.size(2), query.size(1), query.size(2) / key.size(2)};95+ softmax_max_size = {0};
91- softmax_sum_size = {query.size(0), key.size(2), query.size(1), query.size(2) / key.size(2)};96+ softmax_sum_size = {0};
92 }97 }
93 softmax_max = at::empty(softmax_max_size, query.options().dtype(at::kFloat));98 softmax_max = at::empty(softmax_max_size, query.options().dtype(at::kFloat));
94 softmax_sum = at::empty(softmax_sum_size, query.options().dtype(at::kFloat));99 softmax_sum = at::empty(softmax_sum_size, query.options().dtype(at::kFloat));