已合并
refactor: Remove SFA/SFAG/SLI #3591
JialiZheng创建于 6月30日
refactor: Remove SFA/SFAG/SLI #3591
已合并
共 10 个文件变更+31-764
| @@ -24,4 +24,34 @@ from .matmul_add_builder import MatmulAddOpBuilder | |||
| 24 | from .groupmatmul_add_builder import GroupMatmulAddOpBuilder | 24 | from .groupmatmul_add_builder import GroupMatmulAddOpBuilder |
| 25 | from .fused_ema_adamw_builder import FusedEmaAdamWOpBuilder | 25 | from .fused_ema_adamw_builder import FusedEmaAdamWOpBuilder |
| 26 | from .smart_swap_builder import SmartSwapBuilder | 26 | from .smart_swap_builder import SmartSwapBuilder |
| 27 | -from .npu_sparse_lightning_indexer_grad_kl_loss_builder import NPUSparseLIGradKlLossOpBuilder | 27 | + |
| 28 | + | ||
| 29 | +__all__ = [ | ||
| 30 | + "FusionAttentionV2OpBuilder", | ||
| 31 | + "AlgorithmOpBuilder", | ||
| 32 | + "AdaptiveRecomputingPluggableAllocatorBuilder", | ||
| 33 | + "NpuDropoutAddLayerNormOpBuilder", | ||
| 34 | + "AtbOpBuilder", | ||
| 35 | + "SwigluOpBuilder", | ||
| 36 | + "LcalOpBuilder", | ||
| 37 | + "RmsNormOpBuilder", | ||
| 38 | + "GroupedMatMulAllReduceOpBuilder", | ||
| 39 | + "GMMOpBuilder", | ||
| 40 | + "GMMV2OpBuilder", | ||
| 41 | + "QuantGMMOpBuilder", | ||
| 42 | + "WeightQuantGMMOpBuilder", | ||
| 43 | + "FFNOpBuilder", | ||
| 44 | + "MatmulAllReduceAddRmsNormOpBuilder", | ||
| 45 | + "InplaceMatmulAllReduceAddRmsNormOpBuilder", | ||
| 46 | + "RotaryPositionEmbeddingOpBuilder", | ||
| 47 | + "MoeTokenPermuteOpBuilder", | ||
| 48 | + "MoeTokenUnpermuteOpBuilder", | ||
| 49 | + "RingAttentionUpdateOpBuilder", | ||
| 50 | + "BatchMatMulReduceScatterAlltoAllOpBuilder", | ||
| 51 | + "AllToAllAllGatherBatchMatMulOpBuilder", | ||
| 52 | + "AdaptiveCpOpBuilder", | ||
| 53 | + "MatmulAddOpBuilder", | ||
| 54 | + "GroupMatmulAddOpBuilder", | ||
| 55 | + "FusedEmaAdamWOpBuilder", | ||
| 56 | + "SmartSwapBuilder", | ||
| 57 | +] | ||
| @@ -1,32 +0,0 @@ | |||
| 1 | -import os | ||
| 2 | -import torch | ||
| 3 | - | ||
| 4 | -from mindspeed.op_builder.builder import MindSpeedOpBuilder | ||
| 5 | - | ||
| 6 | - | ||
| 7 | -class NPULightningIndexerOpBuilder(MindSpeedOpBuilder): | ||
| 8 | - OP_NAME = "npu_lightning_indexer" | ||
| 9 | - _torch_path = None | ||
| 10 | - | ||
| 11 | - def __init__(self): | ||
| 12 | - from sysconfig import get_paths | ||
| 13 | - self._torch_path = os.path.dirname(os.path.abspath(torch.__file__)) | ||
| 14 | - super(NPULightningIndexerOpBuilder, self).__init__(self.OP_NAME) | ||
| 15 | - | ||
| 16 | - def sources(self): | ||
| 17 | - return ['ops/csrc/cann/npu_lightning_indexer.cpp'] | ||
| 18 | - | ||
| 19 | - def include_paths(self): | ||
| 20 | - paths = super().include_paths() | ||
| 21 | - paths += ['ops/csrc/cann/inc', | ||
| 22 | - os.path.join(self._torch_path, 'include'), | ||
| 23 | - os.path.join(self._torch_path, 'include/torch/csrc/api/include'), | ||
| 24 | - os.path.join(self._torch_npu_path, 'include/torch_npu/csrc/framework/utils'), | ||
| 25 | - os.path.join(self._torch_npu_path, 'include/torch_npu/csrc/aten'), | ||
| 26 | - ] | ||
| 27 | - return paths | ||
| 28 | - | ||
| 29 | - def cxx_args(self): | ||
| 30 | - args = super().cxx_args() | ||
| 31 | - args += ['-Wno-narrowing'] | ||
| 32 | - return args | ||
| @@ -1,32 +0,0 @@ | |||
| 1 | -import os | ||
| 2 | -import torch | ||
| 3 | - | ||
| 4 | -from mindspeed.op_builder.builder import MindSpeedOpBuilder | ||
| 5 | - | ||
| 6 | - | ||
| 7 | -class NPUSparseLIGradKlLossOpBuilder(MindSpeedOpBuilder): | ||
| 8 | - OP_NAME = "npu_sparse_lightning_indexer_grad_kl_loss" | ||
| 9 | - _torch_path = None | ||
| 10 | - | ||
| 11 | - def __init__(self): | ||
| 12 | - from sysconfig import get_paths | ||
| 13 | - self._torch_path = os.path.dirname(os.path.abspath(torch.__file__)) | ||
| 14 | - super(NPUSparseLIGradKlLossOpBuilder, self).__init__(self.OP_NAME) | ||
| 15 | - | ||
| 16 | - def sources(self): | ||
| 17 | - return ['ops/csrc/cann/npu_sparse_lightning_indexer_grad_kl_loss.cpp'] | ||
| 18 | - | ||
| 19 | - def include_paths(self): | ||
| 20 | - paths = super().include_paths() | ||
| 21 | - paths += ['ops/csrc/cann/inc', | ||
| 22 | - os.path.join(self._torch_path, 'include'), | ||
| 23 | - os.path.join(self._torch_path, 'include/torch/csrc/api/include'), | ||
| 24 | - os.path.join(self._torch_npu_path, 'include/torch_npu/csrc/framework/utils'), | ||
| 25 | - os.path.join(self._torch_npu_path, 'include/torch_npu/csrc/aten'), | ||
| 26 | - ] | ||
| 27 | - return paths | ||
| 28 | - | ||
| 29 | - def cxx_args(self): | ||
| 30 | - args = super().cxx_args() | ||
| 31 | - args += ['-Wno-narrowing'] | ||
| 32 | - return args | ||
| @@ -1,121 +0,0 @@ | |||
| 1 | - | ||
| 2 | - | ||
| 3 | - | ||
| 4 | - | ||
| 5 | - | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | -const static int64_t DIM_0 = 0; | ||
| 10 | -const static int64_t DIM_1 = 1; | ||
| 11 | -const static int64_t DIM_2 = 2; | ||
| 12 | -const static int64_t SIZE = 4; | ||
| 13 | - | ||
| 14 | -std::tuple<at::Tensor, at::Tensor> construct_lightning_indexer_output_tensor( | ||
| 15 | - const at::Tensor& query, | ||
| 16 | - const at::Tensor& key, | ||
| 17 | - const c10::optional<at::Tensor> &actual_seq_lengths_query, | ||
| 18 | - int64_t sparse_count, | ||
| 19 | - std::string query_layout_str, | ||
| 20 | - std::string key_layout_str) | ||
| 21 | -{ | ||
| 22 | - at::SmallVector<int64_t, SIZE> output_size; | ||
| 23 | - | ||
| 24 | - if (query_layout_str == "BSND") { | ||
| 25 | - output_size = {query.size(DIM_0), query.size(DIM_1), key.size(DIM_2), sparse_count}; | ||
| 26 | - } else { | ||
| 27 | - int n_dim_index = 0; | ||
| 28 | - n_dim_index = (key_layout_str == "TND") ? DIM_1 : DIM_2; | ||
| 29 | - output_size = {query.size(DIM_0), key.size(n_dim_index), sparse_count}; | ||
| 30 | - } | ||
| 31 | - at::Tensor sparse_indices_out = at::empty(output_size, at::kInt); | ||
| 32 | - at::Tensor sparse_values_out = at::empty(output_size, query.dtype()); | ||
| 33 | - | ||
| 34 | - return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out); | ||
| 35 | -} | ||
| 36 | - | ||
| 37 | -std::tuple<at::Tensor, at::Tensor> npu_lightning_indexer( | ||
| 38 | - const at::Tensor &query, | ||
| 39 | - const at::Tensor &key, | ||
| 40 | - const at::Tensor &weights, | ||
| 41 | - const c10::optional<at::Tensor> &actual_seq_lengths_query, | ||
| 42 | - const c10::optional<at::Tensor> &actual_seq_lengths_key, | ||
| 43 | - const c10::optional<at::Tensor> &block_table, | ||
| 44 | - c10::string_view layout_query, | ||
| 45 | - c10::string_view layout_key, | ||
| 46 | - int64_t sparse_count, | ||
| 47 | - int64_t sparse_mode, | ||
| 48 | - int64_t pre_tokens, | ||
| 49 | - int64_t next_tokens, | ||
| 50 | - int64_t cmp_ratio, | ||
| 51 | - bool return_value) | ||
| 52 | -{ | ||
| 53 | - TORCH_CHECK(query.numel() > 0, "Tensor query is empty.") | ||
| 54 | - TORCH_CHECK(key.numel() > 0, "Tensor key is empty.") | ||
| 55 | - | ||
| 56 | - std::string query_layout_str = std::string(layout_query); | ||
| 57 | - std::string key_layout_str = std::string(layout_key); | ||
| 58 | - | ||
| 59 | - at::SmallVector<int64_t, SIZE> output_size; | ||
| 60 | - // convert str | ||
| 61 | - char *query_layout_ptr = const_cast<char *>(query_layout_str.c_str()); | ||
| 62 | - char *key_layout_ptr = const_cast<char *>(key_layout_str.c_str()); | ||
| 63 | - | ||
| 64 | - if (query_layout_str == "BSND") { | ||
| 65 | - output_size = {query.size(DIM_0), query.size(DIM_1), key.size(DIM_2), sparse_count}; | ||
| 66 | - } else { | ||
| 67 | - int n_dim_index = 0; | ||
| 68 | - n_dim_index = (key_layout_str == "TND") ? DIM_1 : DIM_2; | ||
| 69 | - output_size = {query.size(DIM_0), key.size(n_dim_index), sparse_count}; | ||
| 70 | - } | ||
| 71 | - | ||
| 72 | - at::Tensor sparse_values_out = at::empty(output_size, query.options()); | ||
| 73 | - at::Tensor sparse_indices_out = at::empty(output_size, query.options().dtype(at::kInt)); | ||
| 74 | - | ||
| 75 | - ACLNN_CMD(aclnnLightningIndexer, query, key, weights, | ||
| 76 | - actual_seq_lengths_query, actual_seq_lengths_key, block_table, | ||
| 77 | - query_layout_ptr, key_layout_ptr, sparse_count, sparse_mode, pre_tokens, next_tokens, cmp_ratio, | ||
| 78 | - return_value, sparse_indices_out, sparse_values_out); | ||
| 79 | - | ||
| 80 | - return std::tuple<at::Tensor, at::Tensor>(sparse_indices_out, sparse_values_out); | ||
| 81 | -} | ||
| 82 | - | ||
| 83 | -std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_lightning_indexer_grad( | ||
| 84 | - const at::Tensor &query, | ||
| 85 | - const at::Tensor &key, | ||
| 86 | - const at::Tensor &dy, | ||
| 87 | - const at::Tensor &sparse_indices, | ||
| 88 | - const at::Tensor &weights, | ||
| 89 | - const c10::optional<at::Tensor> &actual_seq_lengths_query, | ||
| 90 | - const c10::optional<at::Tensor> &actual_seq_lengths_key, | ||
| 91 | - const c10::optional<std::string> layout, | ||
| 92 | - c10::optional<int64_t> sparse_mode, | ||
| 93 | - c10::optional<int64_t> pre_tokens, | ||
| 94 | - c10::optional<int64_t> next_tokens, | ||
| 95 | - c10::optional<int64_t> cmp_ratio) | ||
| 96 | -{ | ||
| 97 | - at::Tensor d_query = at::zeros(query.sizes(), query.options()); | ||
| 98 | - at::Tensor d_key = at::zeros(key.sizes(), key.options()); | ||
| 99 | - at::Tensor d_weights = at::zeros(weights.sizes(), weights.options()); | ||
| 100 | - | ||
| 101 | - std::string layout_str_view = layout.value_or("BSND"); | ||
| 102 | - char *layout_ptr = const_cast<char *>(layout_str_view.c_str()); | ||
| 103 | - const int64_t sparse_mode_const = sparse_mode.value_or(0); | ||
| 104 | - const int64_t pre_tokens_const = pre_tokens.value_or(9223372036854775807); | ||
| 105 | - const int64_t next_tokens_const = next_tokens.value_or(9223372036854775807); | ||
| 106 | - const int64_t cmp_ratio_const = cmp_ratio.value_or(1); | ||
| 107 | - const int64_t head_num = 64; | ||
| 108 | - const bool deterministic = false; | ||
| 109 | - | ||
| 110 | - ACLNN_CMD(aclnnLightningIndexerGrad, query, key, dy, sparse_indices, weights, | ||
| 111 | - actual_seq_lengths_query, actual_seq_lengths_key, | ||
| 112 | - head_num, layout_ptr, sparse_mode_const, pre_tokens_const, next_tokens_const, deterministic, cmp_ratio_const, | ||
| 113 | - d_query, d_key, d_weights); | ||
| 114 | - return std::make_tuple(d_query, d_key, d_weights); | ||
| 115 | -} | ||
| 116 | - | ||
| 117 | - | ||
| 118 | -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { | ||
| 119 | - m.def("npu_lightning_indexer", &npu_lightning_indexer, "npu_lightning_indexer forward"); | ||
| 120 | - m.def("npu_lightning_indexer_grad", &npu_lightning_indexer_grad, "npu_lightning_indexer_grad backward"); | ||
| 121 | -} | ||
| @@ -1,83 +0,0 @@ | |||
| 1 | - | ||
| 2 | - | ||
| 3 | - | ||
| 4 | - | ||
| 5 | - | ||
| 6 | - | ||
| 7 | - | ||
| 8 | - | ||
| 9 | -using namespace at_npu::native; | ||
| 10 | -const static int DIMENSION_3D = 3; | ||
| 11 | -const static int DIMENSION_4D = 4; | ||
| 12 | - | ||
| 13 | -std::tuple<at::Tensor, at::Tensor, at::Tensor, at::Tensor> npu_sparse_lightning_indexer_grad_kl_loss_symint( | ||
| 14 | - const at::Tensor &query, | ||
| 15 | - const at::Tensor &key, | ||
| 16 | - const at::Tensor &query_index, | ||
| 17 | - const at::Tensor &key_index, | ||
| 18 | - const at::Tensor &weights, | ||
| 19 | - const at::Tensor &sparse_indices, | ||
| 20 | - const c10::optional<at::Tensor> &softmax_max, | ||
| 21 | - const c10::optional<at::Tensor> &softmax_sum, | ||
| 22 | - const c10::optional<at::Tensor> &query_rope, | ||
| 23 | - const c10::optional<at::Tensor> &key_rope, | ||
| 24 | - const c10::optional<std::vector<int64_t>> actual_seq_qlen, | ||
| 25 | - const c10::optional<std::vector<int64_t>> actual_seq_klen, | ||
| 26 | - double scale_value, | ||
| 27 | - c10::optional<c10::string_view> layout, | ||
| 28 | - c10::optional<int64_t> sparse_mode, | ||
| 29 | - c10::optional<int64_t> pre_tokens, | ||
| 30 | - c10::optional<int64_t> next_tokens, | ||
| 31 | - c10::optional<int64_t> cmp_ratio) | ||
| 32 | -{ | ||
| 33 | - | ||
| 34 | - const at::Tensor &softmax_max_const = softmax_max.value_or(at::Tensor()); | ||
| 35 | - const at::Tensor &softmax_sum_const = softmax_sum.value_or(at::Tensor()); | ||
| 36 | - const at::Tensor &query_rope_const = query_rope.value_or(at::Tensor()); | ||
| 37 | - const at::Tensor &key_rope_const = key_rope.value_or(at::Tensor()); | ||
| 38 | - | ||
| 39 | - auto ac_seq_qlen_tmp = actual_seq_qlen.value_or(std::vector<int64_t>{}); | ||
| 40 | - auto actual_seq_klen_tmp = actual_seq_klen.value_or(std::vector<int64_t>{}); | ||
| 41 | - c10::optional<at::IntArrayRef> actual_seq_qlen_const(ac_seq_qlen_tmp); | ||
| 42 | - c10::optional<at::IntArrayRef> actual_seq_klen_const(actual_seq_klen_tmp); | ||
| 43 | - c10::string_view layout_str = layout.value_or("BSND"); | ||
| 44 | - char *layout_ptr = const_cast<char *>(layout_str.data()); | ||
| 45 | - int64_t sparse_mode_const = sparse_mode.value_or(3); | ||
| 46 | - int64_t pre_tokens_const = pre_tokens.value_or(9223372036854775807); | ||
| 47 | - int64_t next_tokens_const = next_tokens.value_or(9223372036854775807); | ||
| 48 | - bool deterministic_const = true; | ||
| 49 | - const int64_t cmp_ratio_const = cmp_ratio.value_or(1); | ||
| 50 | - TORCH_CHECK(query.dim() == DIMENSION_3D || query.dim() == DIMENSION_4D, | ||
| 51 | - "The shapes of the input query should be 3 or 4 dimensional, but got ", | ||
| 52 | - query.dim(), "-dimensional"); | ||
| 53 | - if (query_rope_const.defined()) { | ||
| 54 | - TORCH_CHECK(query_rope_const.dim() == DIMENSION_3D || query_rope_const.dim() == DIMENSION_4D, | ||
| 55 | - "The shapes of the input query_rope should be 3 or 4 dimensional, but got ", | ||
| 56 | - query_rope_const.dim(), "-dimensional"); | ||
| 57 | - } | ||
| 58 | - TORCH_CHECK(key.dim() == DIMENSION_3D || key.dim() == DIMENSION_4D, | ||
| 59 | - "The shapes of the input key should be 3 or 4 dimensional, but got ", key.dim(), | ||
| 60 | - "-dimensional"); | ||
| 61 | - if (key_rope_const.defined()) { | ||
| 62 | - TORCH_CHECK(key_rope_const.dim() == DIMENSION_3D || key_rope_const.dim() == DIMENSION_4D, | ||
| 63 | - "The shapes of the input key_rope should be 3 or 4 dimensional, but got ", | ||
| 64 | - key_rope_const.dim(), "-dimensional"); | ||
| 65 | - } | ||
| 66 | - at::Tensor d_query_index = at::zeros(query_index.sizes(), query_index.options()); | ||
| 67 | - at::Tensor d_key_index = at::zeros(key_index.sizes(), key_index.options()); | ||
| 68 | - at::Tensor d_weights = at::zeros(weights.sizes(), weights.options()); | ||
| 69 | - at::Tensor loss = at::zeros({1}, query.options().dtype(at::kFloat)); | ||
| 70 | - | ||
| 71 | - ACLNN_CMD( | ||
| 72 | - aclnnSparseLightningIndexerGradKLLoss, query, key, query_index, key_index, weights, | ||
| 73 | - sparse_indices, softmax_max_const, softmax_sum_const, query_rope_const, key_rope_const, actual_seq_qlen_const, | ||
| 74 | - actual_seq_klen_const, scale_value, layout_ptr, sparse_mode_const, pre_tokens_const, next_tokens_const, deterministic_const, | ||
| 75 | - cmp_ratio_const, d_query_index, d_key_index, d_weights, loss); | ||
| 76 | - | ||
| 77 | - return std::make_tuple(d_query_index, d_key_index, d_weights, loss); | ||
| 78 | -} | ||
| 79 | - | ||
| 80 | - | ||
| 81 | -PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { | ||
| 82 | - m.def("npu_sparse_lightning_indexer_grad_kl_loss", &npu_sparse_lightning_indexer_grad_kl_loss_symint, "npu_sparse_lightning_indexer_grad_kl_loss"); | ||
| 83 | -} | ||
| @@ -1,69 +0,0 @@ | |||
| 1 | -import torch | ||
| 2 | - | ||
| 3 | -from mindspeed.op_builder.npu_lightning_indexer_builder import NPULightningIndexerOpBuilder | ||
| 4 | - | ||
| 5 | -__all__ = ["npu_lightning_indexer"] | ||
| 6 | - | ||
| 7 | -op_builder = NPULightningIndexerOpBuilder() | ||
| 8 | - | ||
| 9 | -class LightningIndexer(torch.autograd.Function): | ||
| 10 | - | ||
| 11 | - def forward(ctx, query, key, weights, cu_seq_lens_q, cu_seq_lens_k, layout, | ||
| 12 | - sparse_count, sparse_mode, pre_tokens, next_tokens, cmp_ratio): | ||
| 13 | - op = op_builder.load() | ||
| 14 | - | ||
| 15 | - sparse_indices, sparse_values = op.npu_lightning_indexer( | ||
| 16 | - query, | ||
| 17 | - key, | ||
| 18 | - weights, | ||
| 19 | - cu_seq_lens_q, | ||
| 20 | - cu_seq_lens_k, | ||
| 21 | - None, # BlockTable for inference | ||
| 22 | - layout, | ||
| 23 | - layout, | ||
| 24 | - sparse_count, | ||
| 25 | - sparse_mode, | ||
| 26 | - pre_tokens, | ||
| 27 | - next_tokens, | ||
| 28 | - cmp_ratio, | ||
| 29 | - True, # returnValues | ||
| 30 | - ) | ||
| 31 | - | ||
| 32 | - ctx.save_for_backward(query, key, weights, cu_seq_lens_q, cu_seq_lens_k, sparse_indices) | ||
| 33 | - ctx.layout = layout | ||
| 34 | - ctx.sparse_mode = sparse_mode | ||
| 35 | - ctx.pre_tokens = pre_tokens | ||
| 36 | - ctx.next_tokens = next_tokens | ||
| 37 | - ctx.cmp_ratio = cmp_ratio | ||
| 38 | - | ||
| 39 | - return sparse_indices, sparse_values | ||
| 40 | - | ||
| 41 | - | ||
| 42 | - def backward(ctx, _, grad_output): | ||
| 43 | - op = op_builder.load() | ||
| 44 | - query, key, weights, cu_seq_lens_q, cu_seq_lens_k, sparse_indices = ctx.saved_tensors | ||
| 45 | - query_grad, k_grad, weights_grad = op.npu_lightning_indexer_grad( | ||
| 46 | - query, | ||
| 47 | - key, | ||
| 48 | - grad_output, | ||
| 49 | - sparse_indices, | ||
| 50 | - weights, | ||
| 51 | - cu_seq_lens_q, | ||
| 52 | - cu_seq_lens_k, | ||
| 53 | - ctx.layout, | ||
| 54 | - ctx.sparse_mode, | ||
| 55 | - ctx.pre_tokens, | ||
| 56 | - ctx.next_tokens, | ||
| 57 | - ctx.cmp_ratio | ||
| 58 | - ) | ||
| 59 | - return query_grad, k_grad, weights_grad, None, None, None, None, None, None, None, None | ||
| 60 | - | ||
| 61 | - | ||
| 62 | -def npu_lightning_indexer(query, key, weights, | ||
| 63 | - layout="BSND", cu_seq_lens_q=None, cu_seq_lens_k=None, | ||
| 64 | - sparse_count=2048, sparse_mode=3, | ||
| 65 | - pre_tokens=2**63-1, next_tokens=2**63-1, cmp_ratio=1): | ||
| 66 | - cu_seq_lens_q = cu_seq_lens_k = None # not support TND | ||
| 67 | - return LightningIndexer.apply(query, key, weights, cu_seq_lens_q, cu_seq_lens_k, layout, | ||
| 68 | - sparse_count, sparse_mode, pre_tokens, next_tokens, cmp_ratio) | ||
| 69 | - | ||
| @@ -1,158 +0,0 @@ | |||
| 1 | -import torch | ||
| 2 | - | ||
| 3 | -from mindspeed.op_builder.npu_sparse_lightning_indexer_grad_kl_loss_builder import NPUSparseLIGradKlLossOpBuilder | ||
| 4 | - | ||
| 5 | -__all__ = [ | ||
| 6 | - "npu_sparse_lightning_indexer_grad_kl_loss", | ||
| 7 | - ] | ||
| 8 | - | ||
| 9 | -op_builder = NPUSparseLIGradKlLossOpBuilder() | ||
| 10 | - | ||
| 11 | -class SparseLIGradKlLoss(torch.autograd.Function): | ||
| 12 | - """ | ||
| 13 | - A custom autograd function that computes kl loss in sparse lightning indexer. | ||
| 14 | - | ||
| 15 | - This interface implements the backward functionality of npu_lightning_indexer and integrates the loss computation. | ||
| 16 | - The npu_lightning_indexer selects the top-k pairs between queries and keys in attention that exhibit the strongest | ||
| 17 | - intrinsic correlations, storing them in sparse_indices. This reduces the computational cost of attention in | ||
| 18 | - long-sequence scenarios and improves training performance. | ||
| 19 | - """ | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - def forward( | ||
| 23 | - ctx, | ||
| 24 | - query, | ||
| 25 | - key, | ||
| 26 | - query_index, | ||
| 27 | - key_index, | ||
| 28 | - weights, | ||
| 29 | - sparse_indices, | ||
| 30 | - softmax_max, | ||
| 31 | - softmax_sum, | ||
| 32 | - scale_value=1, | ||
| 33 | - query_rope=None, | ||
| 34 | - key_rope=None, | ||
| 35 | - actual_seq_qlen=None, | ||
| 36 | - actual_seq_klen=None, | ||
| 37 | - layout='BSND', | ||
| 38 | - sparse_mode=3, | ||
| 39 | - pre_tokens=2147483647, | ||
| 40 | - next_tokens=2147483647, | ||
| 41 | - cmp_ratio=1, | ||
| 42 | - ): | ||
| 43 | - """ | ||
| 44 | - Forward pass: compute the total loss by processing hidden states in chunks. | ||
| 45 | - | ||
| 46 | - Args: | ||
| 47 | - ctx: Context object used to save tensors for backward pass. | ||
| 48 | - query (Tensor): Required. Represents the Attention query. Shapes: (B, S1, N1, D), (T1, N1, D) | ||
| 49 | - key (Tensor): Required. Represents the Attention key. Shapes: (B, S2, N2, D), (T2, N2, D) | ||
| 50 | - query_index (Tensor): Required. Input query for the lightning_indexer forward pass. | ||
| 51 | - key_index (Tensor): Required. Input key for the lightning_indexer forward pass. | ||
| 52 | - weights (Tensor): Required. Weight coefficients of lightning_indexer. | ||
| 53 | - sparse_indices (Tensor): Required. Token indices of sorted key and key_index. | ||
| 54 | - softmax_max (Tensor): Required. Maximum values from Attention softmax results. | ||
| 55 | - softmax_sum (Tensor): Required. Sum values from Attention softmax results. | ||
| 56 | - scale_value (float): Required scaling coefficient. | ||
| 57 | - query_rope (Tensor, optional): RoPE information for query in MLA architecture. | ||
| 58 | - key_rope (Tensor, optional): RoPE information for key in MLA architecture. | ||
| 59 | - actual_seq_qlen (list[int], optional): Required in TND layout. Cumulative sequence lengths for query. | ||
| 60 | - actual_seq_klen (list[int], optional): Required in TND layout. Cumulative sequence lengths for key. | ||
| 61 | - layout (str, optional): Input data layout format. Supported: "BSND". Default: "BSND". | ||
| 62 | - sparse_mode (int, optional): Sparse computation mode. Default: 3. | ||
| 63 | - pre_tokens (int, optional): Number of preceding tokens for sparse Attention. Default: 65536. | ||
| 64 | - next_tokens (int, optional): Number of succeeding tokens for sparse Attention. Default: 65536. | ||
| 65 | - cmp_ratio (int, optional): Compression ratio. Default: 1. | ||
| 66 | - Returns: | ||
| 67 | - d_query_index (Tensor): Gradient of query_index. | ||
| 68 | - d_key_index (Tensor): Gradient of key_index. | ||
| 69 | - d_weights (Tensor): Gradient of weights. | ||
| 70 | - loss (Tensor): Difference between network forward output and golden value. | ||
| 71 | - """ | ||
| 72 | - op = op_builder.load() | ||
| 73 | - | ||
| 74 | - d_query_index, d_key_index, d_weights, loss = op.npu_sparse_lightning_indexer_grad_kl_loss( | ||
| 75 | - query, | ||
| 76 | - key, | ||
| 77 | - query_index, | ||
| 78 | - key_index, | ||
| 79 | - weights, | ||
| 80 | - sparse_indices, | ||
| 81 | - softmax_max, | ||
| 82 | - softmax_sum, | ||
| 83 | - query_rope, | ||
| 84 | - key_rope, | ||
| 85 | - actual_seq_qlen, | ||
| 86 | - actual_seq_klen, | ||
| 87 | - scale_value, | ||
| 88 | - layout, | ||
| 89 | - sparse_mode, | ||
| 90 | - pre_tokens, | ||
| 91 | - next_tokens, | ||
| 92 | - cmp_ratio, | ||
| 93 | - ) | ||
| 94 | - | ||
| 95 | - # Save computed gradients for use in backward pass | ||
| 96 | - ctx.save_for_backward(d_query_index, d_key_index, d_weights) | ||
| 97 | - return loss[0] | ||
| 98 | - | ||
| 99 | - | ||
| 100 | - def backward(ctx, *grad_output): | ||
| 101 | - """ | ||
| 102 | - Backward pass: propagate upstream gradients through the precomputed gradients. | ||
| 103 | - | ||
| 104 | - Args: | ||
| 105 | - ctx: Context object with saved tensors from forward pass. | ||
| 106 | - grad_output: Gradient output. | ||
| 107 | - | ||
| 108 | - Returns: | ||
| 109 | - tuple: Gradients. | ||
| 110 | - """ | ||
| 111 | - d_query_index, d_key_index, d_weights = ctx.saved_tensors | ||
| 112 | - grad_scale = grad_output[0] | ||
| 113 | - if torch.ne(grad_scale, torch.tensor(1.0, device=grad_scale.device)): | ||
| 114 | - d_query_index = d_query_index * grad_scale | ||
| 115 | - d_key_index = d_key_index * grad_scale | ||
| 116 | - d_weights = d_weights * grad_scale | ||
| 117 | - | ||
| 118 | - res_list = [None] * 13 | ||
| 119 | - return None, None, d_query_index, d_key_index, d_weights, *res_list | ||
| 120 | - | ||
| 121 | - | ||
| 122 | -def npu_sparse_lightning_indexer_grad_kl_loss( | ||
| 123 | - query, | ||
| 124 | - key, | ||
| 125 | - query_index, | ||
| 126 | - key_index, | ||
| 127 | - weights, | ||
| 128 | - topk_indices, | ||
| 129 | - softmax_max, | ||
| 130 | - softmax_sum, | ||
| 131 | - scale_value=1, | ||
| 132 | - *, | ||
| 133 | - query_rope=None, | ||
| 134 | - key_rope=None, | ||
| 135 | - actual_seq_qlen=None, | ||
| 136 | - actual_seq_klen=None, | ||
| 137 | - layout='BSND', | ||
| 138 | - sparse_mode=3, | ||
| 139 | - pre_tokens=2147483647, | ||
| 140 | - next_tokens=2147483647, | ||
| 141 | - cmp_ratio=1, | ||
| 142 | -): | ||
| 143 | - """NPU Sparse Lightning Indexer KL Divergence Loss Function""" | ||
| 144 | - query, key, query_index, key_index, weights = [x.transpose(0, 1) for x in | ||
| 145 | - [query, key, query_index, key_index, weights]] | ||
| 146 | - if len(key.shape) == 3: | ||
| 147 | - key = key.unsqueeze(2) | ||
| 148 | - topk_indices = topk_indices.unsqueeze(2) | ||
| 149 | - if query_rope is not None: | ||
| 150 | - query_rope, key_rope = [x.transpose(0, 1) for x in [query_rope, key_rope]] | ||
| 151 | - | ||
| 152 | - bsz = query.shape[0] | ||
| 153 | - sq = query.shape[1] | ||
| 154 | - loss = SparseLIGradKlLoss.apply( | ||
| 155 | - query, key, query_index, key_index, weights, topk_indices, softmax_max, softmax_sum, | ||
| 156 | - scale_value, query_rope, key_rope, actual_seq_qlen, actual_seq_klen, layout, sparse_mode, | ||
| 157 | - pre_tokens, next_tokens, cmp_ratio) | ||
| 158 | - return loss / (sq * bsz) | ||