已合并
refactor: Remove SFA/SFAG/SLI #3591
JialiZheng创建于 6月30日
refactor: Remove SFA/SFAG/SLI #3591
已合并
JialiZheng创建于 6月30日
共 10 个文件变更+31-764
@@ -24,4 +24,34 @@ from .matmul_add_builder import MatmulAddOpBuilder
24from .groupmatmul_add_builder import GroupMatmulAddOpBuilder24from .groupmatmul_add_builder import GroupMatmulAddOpBuilder
25from .fused_ema_adamw_builder import FusedEmaAdamWOpBuilder25from .fused_ema_adamw_builder import FusedEmaAdamWOpBuilder
26from .smart_swap_builder import SmartSwapBuilder26from .smart_swap_builder import SmartSwapBuilder
27-from .npu_sparse_lightning_indexer_grad_kl_loss_builder import NPUSparseLIGradKlLossOpBuilder27+ 
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 NPUSparseAttnSharedKVOpBuilder(MindSpeedOpBuilder):
8- OP_NAME = "npu_sparse_attn_shared_kv"
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(NPUSparseAttnSharedKVOpBuilder, self).__init__(self.OP_NAME)
15- 
16- def sources(self):
17- return ['ops/csrc/cann/npu_sparse_attn_shared_kv.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-#include <torch/extension.h>
2-#include <torch_npu/csrc/framework/utils/RandomOpAdapter.h>
3-#include <torch_npu/csrc/framework/utils/OpAdapter.h>
4-#include <torch_npu/csrc/core/npu/NPUFormat.h>
5-#include <torch_npu/csrc/include/ops.h>
6- 
7-#include "inc/aclnn_common.h"
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,110 +0,0 @@
1-#include <torch/extension.h>
2-#include <torch_npu/csrc/framework/utils/RandomOpAdapter.h>
3-#include <torch_npu/csrc/framework/utils/OpAdapter.h>
4-#include <torch_npu/csrc/core/npu/NPUFormat.h>
5-#include <torch_npu/csrc/include/ops.h>
6- 
7-#include "inc/aclnn_common.h"
8- 
9-at::Tensor npu_sparse_attn_shared_kv_metadata(const c10::optional<at::Tensor> &cuSeqLensQ,
10- const c10::optional<at::Tensor> &sequsedOriKv, const c10::optional<at::Tensor> &sequsedCmpKv,
11- const c10::optional<at::Tensor> &sequsedQ, const c10::optional<at::Tensor> &sequsedKv, int64_t numHeadsQ,
12- int64_t numHeadsKv, int64_t headDim, int64_t batchSize, int64_t maxSeqLenQ, int64_t maxSeqLenKv, int64_t oriTopk,
13- int64_t cmpTopk, int64_t cmpRatio, int64_t oriMaskMode, int64_t cmpMaskMode, int64_t oriWinLeft,
14- int64_t oriWinRight, const c10::optional<std::string> layoutQ, const c10::optional<std::string> layoutKv,
15- bool hasOriKv, bool hasCmpKv) {
16- char *layoutQPtr = const_cast<char *>(layoutQ.value_or("SBH").c_str());
17- char *layoutKvPtr = const_cast<char *>(layoutKv.value_or("SBH").c_str());
18- at::Tensor metadata = at::empty(1024, at::TensorOptions(torch_npu::utils::get_npu_device_type()).dtype(at::kInt));
19- ACLNN_CMD(aclnnSparseAttnSharedkvMetadata, cuSeqLensQ, sequsedOriKv, sequsedCmpKv, sequsedQ, sequsedKv, numHeadsQ,
20- numHeadsKv, headDim, batchSize, maxSeqLenQ, maxSeqLenKv, oriTopk, cmpTopk, cmpRatio, oriMaskMode, cmpMaskMode,
21- oriWinLeft, oriWinRight, layoutQPtr, layoutKvPtr, hasOriKv, hasCmpKv, metadata);
22- return metadata;
23-}
24- 
25-std::tuple<at::Tensor, at::Tensor> npu_sparse_attn_shared_kv(const at::Tensor &query,
26- const c10::optional<at::Tensor> &oriKv, const c10::optional<at::Tensor> &cmpKv,
27- const c10::optional<at::Tensor> &oriSparseIndices, const c10::optional<at::Tensor> &cmpSparseIndices,
28- const c10::optional<at::Tensor> &oriBlockTable, const c10::optional<at::Tensor> &cmpBlockTable,
29- const c10::optional<at::Tensor> &cuSeqLensQ, const c10::optional<at::Tensor> &cuSeqLensOriKv,
30- const c10::optional<at::Tensor> &cuSeqLensCmpKv, const c10::optional<at::Tensor> &sequsedQ,
31- const c10::optional<at::Tensor> &sequsedKv, const c10::optional<at::Tensor> &sinks,
32- const c10::optional<at::Tensor> &metadata, double softmaxScale, int64_t cmpRatio, int64_t oriMaskMode,
33- int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight, const c10::optional<std::string> layoutQ,
34- const c10::optional<std::string> layoutKv, bool returnSoftmaxLse) {
35- std::string layoutq = layoutQ.value_or("SBH");
36- std::string layoutkv = layoutKv.value_or("SBH");
37- char *layoutQPtr = const_cast<char *>(layoutq.c_str());
38- char *layoutKvPtr = const_cast<char *>(layoutkv.c_str());
39- 
40- at::Tensor attnOutput = at::empty(query.sizes(), query.options());
41- at::Tensor softmaxLseOut;
42- if (returnSoftmaxLse) {
43- std::vector<int64_t> lse_sizes(query.sizes().begin(), query.sizes().end());
44- lse_sizes.back() = 1;
45- softmaxLseOut = at::empty(lse_sizes, query.options().dtype(c10::ScalarType::Float));
46- } else {
47- softmaxLseOut = at::Tensor();
48- }
49- int64_t ori_kv_stride = 0;
50- int64_t cmp_kv_stride = 0;
51- if (oriKv.has_value()) {
52- const at::Tensor &tmp_kv = *oriKv;
53- ori_kv_stride = tmp_kv.stride(0);
54- }
55- if (cmpKv.has_value()) {
56- const at::Tensor &tmp_kv = *cmpKv;
57- cmp_kv_stride = tmp_kv.stride(0);
58- }
59- ACLNN_CMD(aclnnSparseAttnSharedkv, query, oriKv, cmpKv, oriSparseIndices, cmpSparseIndices, oriBlockTable,
60- cmpBlockTable, cuSeqLensQ, cuSeqLensOriKv, cuSeqLensCmpKv, sequsedQ, sequsedKv, sinks, metadata, softmaxScale,
61- cmpRatio, oriMaskMode, cmpMaskMode, ori_kv_stride, cmp_kv_stride, oriWinLeft, oriWinRight, layoutQPtr,
62- layoutKvPtr, returnSoftmaxLse, attnOutput, softmaxLseOut);
63- return std::make_tuple(attnOutput, softmaxLseOut);
64-}
65- 
66-std::tuple<at::Tensor, at::Tensor, const c10::optional<at::Tensor>, at::Tensor> npu_sparse_attn_shared_kv_grad(
67- const at::Tensor &query, const at::Tensor &oriKv, const c10::optional<const at::Tensor> &cmpKvOptional,
68- const c10::optional<const at::Tensor> &dOutOptional, const c10::optional<const at::Tensor> &outOptional,
69- const c10::optional<const at::Tensor> &lseOptional, const c10::optional<const at::Tensor> &oriSparseIndicesOptional,
70- const c10::optional<const at::Tensor> &cmpSparseIndicesOptional,
71- const c10::optional<const at::Tensor> &cuSeqlensQOptional,
72- const c10::optional<const at::Tensor> &cuSeqlensOriKvOptional,
73- const c10::optional<const at::Tensor> &cuSeqlensCmpKvOptional, const at::Tensor &sinks, double scaleValue,
74- int64_t cmpRatio, int64_t oriMaskMode, int64_t cmpMaskMode, int64_t oriWinLeft, int64_t oriWinRight,
75- const c10::optional<std::string> layout) {
76- const at::Tensor &cmpKv = cmpKvOptional.value_or(at::Tensor());
77- const at::Tensor &dOut = dOutOptional.value_or(at::Tensor());
78- const at::Tensor &out = outOptional.value_or(at::Tensor());
79- const at::Tensor &lse = lseOptional.value_or(at::Tensor());
80- const at::Tensor &oriSparseIndices = oriSparseIndicesOptional.value_or(at::Tensor());
81- const at::Tensor &cmpSparseIndices = cmpSparseIndicesOptional.value_or(at::Tensor());
82- const at::Tensor &cuSeqlensQ = cuSeqlensQOptional.value_or(at::Tensor());
83- const at::Tensor &cuSeqlensOriKv = cuSeqlensOriKvOptional.value_or(at::Tensor());
84- const at::Tensor &cuSeqlensCmpKv = cuSeqlensCmpKvOptional.value_or(at::Tensor());
85- 
86- std::string layoutValue = layout.value_or("SBH");
87- char *layoutPtr = const_cast<char *>(layoutValue.c_str());
88- at::Tensor dQuery = at::empty(query.sizes(), query.options());
89- at::Tensor dOriKv = at::empty(oriKv.sizes(), oriKv.options());
90- at::Tensor dSinks = at::empty(sinks.sizes(), sinks.options());
91- 
92- at::Tensor dCmpKv;
93- if (cmpRatio > 1) {
94- dCmpKv = at::empty(cmpKv.sizes(), cmpKv.options());
95- } else {
96- dCmpKv = at::Tensor();
97- }
98- 
99- ACLNN_CMD(aclnnSparseAttnSharedkvGrad, query, oriKv, cmpKv, dOut, out, lse, oriSparseIndices, cmpSparseIndices,
100- cuSeqlensQ, cuSeqlensOriKv, cuSeqlensCmpKv, sinks, scaleValue, cmpRatio, oriMaskMode, cmpMaskMode, oriWinLeft,
101- oriWinRight, layoutPtr, dQuery, dOriKv, dCmpKv, dSinks);
102- return std::make_tuple(dQuery, dOriKv, dCmpKv, dSinks);
103-}
104- 
105-PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
106- m.def("npu_sparse_attn_shared_kv_metadata", &npu_sparse_attn_shared_kv_metadata,
107- "npu_sparse_attn_shared_kv metadata");
108- m.def("npu_sparse_attn_shared_kv", &npu_sparse_attn_shared_kv, "npu_sparse_attn_shared_kv forward");
109- m.def("npu_sparse_attn_shared_kv_grad", &npu_sparse_attn_shared_kv_grad, "npu_sparse_attn_shared_kv_grad backward");
110-}
@@ -1,83 +0,0 @@
1-#include <torch/extension.h>
2-#include <torch_npu/csrc/framework/utils/RandomOpAdapter.h>
3-#include <torch_npu/csrc/framework/utils/OpAdapter.h>
4-#include <torch_npu/csrc/core/npu/NPUFormat.h>
5-#include <torch_npu/csrc/include/ops.h>
6- 
7-#include "inc/aclnn_common.h"
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- @staticmethod
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- @staticmethod
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,126 +0,0 @@
1-import torch
2- 
3-from mindspeed.op_builder.npu_sparse_attn_shared_kv_builder import NPUSparseAttnSharedKVOpBuilder
4- 
5-op_builder = NPUSparseAttnSharedKVOpBuilder()
6- 
7- 
8-class SparseAttnSharedKV(torch.autograd.Function):
9- @staticmethod
10- def forward(ctx, query, ori_kv, cmp_kv, cu_seq_lens_q, cu_seq_lens_ori_kv, cu_seq_lens_cmp_kv, ori_sparse_indices,
11- cmp_sparse_indices, sinks, softmax_scale, cmp_ratio, ori_mask_mode, cmp_mask_mode, ori_win_left,
12- ori_win_right, num_heads_q, num_heads_kv, head_dim, batch_size, max_seq_len_q, max_seq_len_kv, topk,
13- layout_q, layout_kv):
14- op = op_builder.load()
15- 
16- metadata = op.npu_sparse_attn_shared_kv_metadata(
17- cu_seq_lens_q if cu_seq_lens_q is not None else torch.tensor([]).npu(),
18- torch.tensor([]).npu(), # sequsedOriKv for inference
19- torch.tensor([]).npu(), # sequsedCmpKv for inference
20- torch.tensor([]).npu(), # sequsedQ for inference
21- torch.tensor([]).npu(), # sequsedKv for inference
22- num_heads_q,
23- num_heads_kv,
24- head_dim,
25- batch_size,
26- max_seq_len_q,
27- max_seq_len_kv,
28- topk, # oriTopk not support now
29- topk,
30- cmp_ratio,
31- ori_mask_mode,
32- cmp_mask_mode,
33- ori_win_left,
34- ori_win_right,
35- layout_q,
36- layout_kv,
37- ori_kv is not None, # hasOriKv
38- cmp_kv is not None, # hasCmpKv
39- )
40- 
41- result, softmax_lse = op.npu_sparse_attn_shared_kv(
42- query,
43- ori_kv,
44- cmp_kv,
45- ori_sparse_indices,
46- cmp_sparse_indices,
47- None, # oriBlockTable for inference
48- None, # cmpBlockTable for inference
49- cu_seq_lens_q,
50- cu_seq_lens_ori_kv,
51- cu_seq_lens_cmp_kv,
52- None, # sequsedQ for inference
53- None, # sequsedKv for inference
54- sinks,
55- metadata,
56- softmax_scale,
57- cmp_ratio,
58- ori_mask_mode,
59- cmp_mask_mode,
60- ori_win_left,
61- ori_win_right,
62- layout_q,
63- layout_kv,
64- True, # returnSoftmaxLse
65- )
66- 
67- ctx.save_for_backward(query, ori_kv, cmp_kv, result, softmax_lse, ori_sparse_indices, cmp_sparse_indices,
68- cu_seq_lens_q, cu_seq_lens_ori_kv, cu_seq_lens_cmp_kv, sinks)
69- ctx.softmax_scale = softmax_scale
70- ctx.cmp_ratio = cmp_ratio
71- ctx.ori_mask_mode = ori_mask_mode
72- ctx.cmp_mask_mode = cmp_mask_mode
73- ctx.ori_win_left = ori_win_left
74- ctx.ori_win_right = ori_win_right
75- ctx.layout_q = layout_q
76- return result
77- 
78- @staticmethod
79- def backward(ctx, grad_output):
80- op = op_builder.load()
81- query, ori_kv, cmp_kv, result, softmax_lse, ori_sparse_indices, cmp_sparse_indices, cu_seq_lens_q, \
82- cu_seq_lens_ori_kv, cu_seq_lens_cmp_kv, sinks = ctx.saved_tensors
83- query_grad, ori_kv_grad, cmp_kv_grad, sinks_grad = op.npu_sparse_attn_shared_kv_grad(
84- query,
85- ori_kv,
86- cmp_kv,
87- grad_output,
88- result,
89- softmax_lse,
90- ori_sparse_indices,
91- cmp_sparse_indices,
92- cu_seq_lens_q,
93- cu_seq_lens_ori_kv,
94- cu_seq_lens_cmp_kv,
95- sinks,
96- ctx.softmax_scale,
97- ctx.cmp_ratio,
98- ctx.ori_mask_mode,
99- ctx.cmp_mask_mode,
100- ctx.ori_win_left,
101- ctx.ori_win_right,
102- ctx.layout_q
103- )
104- return query_grad, ori_kv_grad, cmp_kv_grad, None, None, None, None, None, sinks_grad, None, None, None, None, \
105- None, None, None, None, None, None, None, None, None, None, None
106- 
107- 
108-def npu_sparse_attn_shared_kv(query, ori_kv, cmp_kv, cmp_sparse_indices, sinks, softmax_scale, cmp_ratio,
109- ori_mask_mode=4, cmp_mask_mode=3, ori_win_left=127, ori_win_right=0):
110- cu_seq_lens_q = cu_seq_lens_ori_kv = cu_seq_lens_cmp_kv = None # not support TND
111- ori_sparse_indices = None # ori kv use band mode
112- max_seq_len_q, batch_size, num_heads_q, head_dim = query.size()
113- num_heads_kv = 1
114- max_seq_len_kv = ori_kv.size(0)
115- topk = 0 if cmp_ratio != 4 else cmp_sparse_indices.size(-1)
116- layout_q = layout_kv = 'BSND'
117- query = query.permute(1, 0, 2, 3).contiguous() # [S, B, N, D] --> [B, S, N, D]
118- ori_kv = ori_kv.permute(1, 0, 2).unsqueeze(2).contiguous() # [S, B, D] --> [B, S, 1, D]
119- cmp_kv = cmp_kv if cmp_kv is None else cmp_kv.permute(1, 0, 2).unsqueeze(2).contiguous() # [S, B, D] --> [B, S, 1, D]
120- cmp_sparse_indices = None if cmp_ratio != 4 else cmp_sparse_indices.unsqueeze(2).contiguous() # [B, S, K] --> [B, S, 1, K]
121- output = SparseAttnSharedKV.apply(query, ori_kv, cmp_kv, cu_seq_lens_q, cu_seq_lens_ori_kv, cu_seq_lens_cmp_kv,
122- ori_sparse_indices, cmp_sparse_indices, sinks, softmax_scale, cmp_ratio,
123- ori_mask_mode, cmp_mask_mode, ori_win_left, ori_win_right, num_heads_q,
124- num_heads_kv, head_dim, batch_size, max_seq_len_q, max_seq_len_kv, topk, layout_q,
125- layout_kv)
126- return output.transpose(0, 1).contiguous()
@@ -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- @staticmethod
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- @staticmethod
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)