已合并
[fix] Refine shape inference for optional outputs of LI and SFA operators. #4400
zzzyh22创建于 3月9日
[fix] Refine shape inference for optional outputs of LI and SFA operators. #4400
已合并
共 2 个文件变更+19-8
| @@ -23,12 +23,13 @@ namespace op_api { | |||
| 23 | const static int64_t DIM_0 = 0; | 23 | const static int64_t DIM_0 = 0; |
| 24 | const static int64_t DIM_1 = 1; | 24 | const static int64_t DIM_1 = 1; |
| 25 | const static int64_t DIM_2 = 2; | 25 | const static int64_t DIM_2 = 2; |
| 26 | +const static int64_t DIM_3 = 3; | ||
| 26 | using namespace at_npu::native; | 27 | using namespace at_npu::native; |
| 27 | using npu_preparation = at_npu::native::OpPreparation; | 28 | using npu_preparation = at_npu::native::OpPreparation; |
| 28 | 29 | ||
| 29 | std::tuple<at::Tensor, at::Tensor> construct_lightning_indexer_output_tensor(const at::Tensor& query, | 30 | std::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 tensor | 76 | // 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 str | 81 | // convert str |
| @@ -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)); |