已合并
新增融合算子mhc_pre_sinkhorn_backward #4760
kdy18482276080创建于 4月28日
新增融合算子mhc_pre_sinkhorn_backward #4760
已合并
kdy18482276080创建于 4月28日
8 个文件变更+1785-0
@@ -0,0 +1,19 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
12+if(NOT ENABLE_TEST AND NOT BENCHMARK)
13+ list(REMOVE_ITEM CURRENT_DIRS tests)
14+endif()
15+foreach(SUB_DIR ${CURRENT_DIRS})
16+ if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
17+ add_subdirectory(${SUB_DIR})
18+ endif()
19+endforeach()
@@ -0,0 +1,35 @@
1+# -----------------------------------------------------------------------------------------------------------
2+# Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+# CANN Open Software License Agreement Version 2.0 (the "License").
5+# Please refer to the License for details. You may not use this file except in compliance with the License.
6+# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+# See LICENSE in the root of the software repository for the full text of the License.
9+# -----------------------------------------------------------------------------------------------------------
10+ 
11+if (BUILD_OPEN_PROJECT)
12+ target_sources(op_host_aclnn PRIVATE
13+ mhc_pre_sinkhorn_backward_def.cpp
14+ )
15+ 
16+ if("${ASCEND_COMPUTE_UNIT}" STREQUAL "ascend950")
17+ add_ops_compile_options(
18+ OP_NAME MhcPreSinkhornBackward
19+ COMPUTE_UNIT Ascend950PR_9599
20+ OPTIONS -mllvm -cce-aicore-dcci-before-kernel-end=false
21+ )
22+ else()
23+ add_ops_compile_options(
24+ OP_NAME MhcPreSinkhornBackward
25+ OPTIONS --cce-auto-sync=off
26+ -Wno-deprecated-declarations
27+ -Werror
28+ )
29+ endif()
30+endif()
31+ 
32+if(NOT BUILD_OPS_RTY_KERNEL)
33+ add_op_to_compiled_list()
34+ add_modules_sources(OPTYPE mhc_pre_sinkhorn_backward ACLNNTYPE aclnn)
35+endif()
@@ -0,0 +1,119 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file mhc_pre_sinkhorn_backward_def.cpp
13+ * \brief MhcPreSinkhornBackward operator definition
14+ */
15+#include "register/op_def_registry.h"
16+ 
17+namespace ops {
18+class MhcPreSinkhornBackward : public OpDef {
19+public:
20+ explicit MhcPreSinkhornBackward(const char* name) : OpDef(name)
21+ {
22+ this->Input("grad_hin")
23+ .ParamType(REQUIRED)
24+ .DataType({ge::DT_BF16, ge::DT_FLOAT16})
25+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
26+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
27+ this->Input("grad_h_post")
28+ .ParamType(REQUIRED)
29+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
30+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
31+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
32+ this->Input("grad_h_res")
33+ .ParamType(REQUIRED)
34+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
35+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
36+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
37+ this->Input("x")
38+ .ParamType(REQUIRED)
39+ .DataType({ge::DT_BF16, ge::DT_FLOAT16})
40+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
41+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
42+ this->Input("phi")
43+ .ParamType(REQUIRED)
44+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
45+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
46+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
47+ this->Input("alpha")
48+ .ParamType(REQUIRED)
49+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
50+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
51+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
52+ this->Input("bias")
53+ .ParamType(REQUIRED)
54+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
55+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
56+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
57+ this->Input("h_pre")
58+ .ParamType(REQUIRED)
59+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
60+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
61+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
62+ this->Input("hc_before_norm")
63+ .ParamType(REQUIRED)
64+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
65+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
66+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
67+ this->Input("inv_rms")
68+ .ParamType(REQUIRED)
69+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
70+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
71+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
72+ this->Input("sum_out")
73+ .ParamType(REQUIRED)
74+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
75+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
76+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
77+ this->Input("norm_out")
78+ .ParamType(REQUIRED)
79+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
80+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
81+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
82+ this->Output("grad_x")
83+ .ParamType(REQUIRED)
84+ .DataType({ge::DT_BF16, ge::DT_FLOAT16})
85+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
86+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
87+ this->Output("grad_phi")
88+ .ParamType(REQUIRED)
89+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
90+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
91+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
92+ this->Output("grad_alpha")
93+ .ParamType(REQUIRED)
94+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
95+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
96+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
97+ this->Output("grad_bias")
98+ .ParamType(REQUIRED)
99+ .DataType({ge::DT_FLOAT, ge::DT_FLOAT})
100+ .Format({ge::FORMAT_ND, ge::FORMAT_ND})
101+ .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND});
102+ 
103+ OpAICoreConfig aicConfig;
104+ aicConfig.DynamicCompileStaticFlag(true)
105+ .DynamicFormatFlag(false)
106+ .DynamicRankSupportFlag(true)
107+ .DynamicShapeSupportFlag(true)
108+ .NeedCheckSupportFlag(false)
109+ .ExtendCfgInfo("aclnnSupport.value", "support_aclnn");
110+ this->AICore().AddConfig("ascend910b", aicConfig);
111+ this->AICore().AddConfig("ascend910_93", aicConfig);
112+ 
113+ this->Attr("hc_eps").AttrType(OPTIONAL).Float(1e-6f);
114+ }
115+};
116+ 
117+OP_ADD(MhcPreSinkhornBackward);
118+ 
119+}
@@ -0,0 +1,71 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file common.h
13+ * \brief
14+ */
15+#ifndef MHC_PRE_SINKHORN_BACKWARD_OP_HOST_OP_TILING_ARCH32_COMMON_H
16+#define MHC_PRE_SINKHORN_BACKWARD_OP_HOST_OP_TILING_ARCH32_COMMON_H
17+ 
18+#include "register/op_def_registry.h"
19+#include "tiling/platform/platform_ascendc.h"
20+#include "tiling/tiling_api.h"
21+#include "register/tilingdata_base.h"
22+ 
23+inline std::map<ge::DataType, uint64_t> kDataSizeMap = {
24+ {ge::DT_FLOAT, sizeof(float)},
25+ {ge::DT_INT32, sizeof(int32_t)},
26+ {ge::DT_INT64, sizeof(int64_t)}
27+};
28+ 
29+/**
30+ * if b is 0, return a
31+ */
32+template<typename T>
33+inline T DivCeil(T a, T b)
34+{
35+ if (b == 0) {
36+ return 0;
37+ }
38+ return (a + b - 1) / b;
39+}
40+ 
41+/**
42+ * if b is 0, return 0
43+ */
44+template<typename T>
45+inline T CeilAlign(T a, T b)
46+{
47+ if (b == 0) {
48+ return 0;
49+ }
50+ return DivCeil(a, b) * b;
51+}
52+ 
53+/**
54+ * if b is 0, return a
55+ */
56+template<typename T>
57+inline T DivFloor(T a, T b)
58+{
59+ return b == 0 ? a : a / b;
60+}
61+ 
62+/**
63+ * if b is 0, return 0
64+ */
65+template<typename T>
66+inline T FloorAlign(T a, T b)
67+{
68+ return b == 0 ? 0 : a / b * b;
69+}
70+ 
71+#endif // MHC_PRE_SINKHORN_BACKWARD_OP_HOST_OP_TILING_ARCH32_COMMON_H
@@ -0,0 +1,441 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file mhc_pre_sinkhorn_backward_tiling.cpp
13+ * \brief MhcPreSinkhornBackward operator tiling implementation
14+ */
15+#include "register/op_def_registry.h"
16+#include "tiling/tiling_api.h"
17+#include "tiling/platform/platform_ascendc.h"
18+#include "mhc_pre_sinkhorn_backward_tiling.h"
19+#include "log/log.h"
20+ 
21+#define CHECK_NULLPTR(ptr) \
22+ if (ptr == nullptr) { \
23+ return ge::GRAPH_FAILED; \
24+ }
25+ 
26+namespace {
27+constexpr uint8_t GRAD_HIN_IDX = 0;
28+constexpr uint8_t GRAD_H_POST_IDX = 1;
29+constexpr uint8_t GRAD_H_RES_IDX = 2;
30+constexpr uint8_t INPUT_X_IDX = 3;
31+constexpr uint8_t PHI_IDX = 4;
32+constexpr uint8_t ALPHA_IDX = 5;
33+constexpr uint8_t BIAS_IDX = 6;
34+constexpr uint8_t H_PRE_IDX = 7;
35+constexpr uint8_t HC_BEFORE_NORM_IDX = 8;
36+constexpr uint8_t INV_RMS_IDX = 9;
37+constexpr uint8_t SUM_OUT_IDX = 10;
38+constexpr uint8_t NORM_OUT_IDX = 11;
39+constexpr uint8_t GRAD_X_IDX = 0;
40+constexpr uint8_t GRAD_PHI_IDX = 1;
41+constexpr uint8_t GRAD_ALPHA_IDX = 2;
42+constexpr uint8_t GRAD_BIAS_IDX = 3;
43+constexpr float DEFAULT_EPS = 1e-6f;
44+constexpr uint8_t BATCH_SIZE_DIM_IDX = 0;
45+constexpr uint8_t SEQ_LENGTH_DIM_IDX = 1;
46+constexpr uint8_t N_DIM_IDX = 2;
47+constexpr uint8_t C_DIM_IDX = 3;
48+constexpr uint8_t ITER_COUNT_IDX = 0;
49+constexpr int64_t N_SIZE_4 = 4;
50+constexpr int64_t ALPHA_SIZE_3 = 3;
51+constexpr int64_t C_V_RATIO = 2;
52+} // namespace
53+ 
54+using namespace ge;
55+using namespace std;
56+using namespace AscendC;
57+ 
58+namespace optiling {
59+ 
60+ge::graphStatus ShapeVerify(gert::TilingContext *context, int64_t batchSize, int64_t seqLength, int64_t n, int64_t c)
61+{
62+ auto opName = context->GetNodeName();
63+ 
64+ auto gradHinShapePtr = context->GetInputShape(GRAD_HIN_IDX);
65+ OP_CHECK_IF(gradHinShapePtr == nullptr,
66+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradHin shape is nullptr"),
67+ return ge::GRAPH_FAILED);
68+ auto gradHinShape = gradHinShapePtr->GetStorageShape();
69+ OP_CHECK_IF(gradHinShape.GetDimNum() != 3,
70+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradHin must be 3D, but got %lu dims",
71+ gradHinShape.GetDimNum()),
72+ return ge::GRAPH_FAILED);
73+ OP_CHECK_IF(gradHinShape.GetDim(0) != batchSize || gradHinShape.GetDim(1) != seqLength ||
74+ gradHinShape.GetDim(2) != c,
75+ OPS_REPORT_VECTOR_INNER_ERR(
76+ opName, "ShapeVerify gradHin failed, expected (B=%ld, S=%ld, C=%ld), but got (%ld, %ld, %ld)",
77+ batchSize, seqLength, c, gradHinShape.GetDim(0), gradHinShape.GetDim(1), gradHinShape.GetDim(2)),
78+ return ge::GRAPH_FAILED);
79+ 
80+ auto gradHPostShapePtr = context->GetInputShape(GRAD_H_POST_IDX);
81+ OP_CHECK_IF(gradHPostShapePtr == nullptr,
82+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradHPost shape is nullptr"),
83+ return ge::GRAPH_FAILED);
84+ auto gradHPostShape = gradHPostShapePtr->GetStorageShape();
85+ OP_CHECK_IF(gradHPostShape.GetDimNum() != 3,
86+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradHPost must be 3D, but got %lu dims",
87+ gradHPostShape.GetDimNum()),
88+ return ge::GRAPH_FAILED);
89+ OP_CHECK_IF(
90+ gradHPostShape.GetDim(0) != batchSize || gradHPostShape.GetDim(1) != seqLength || gradHPostShape.GetDim(2) != n,
91+ OPS_REPORT_VECTOR_INNER_ERR(
92+ opName, "ShapeVerify gradHPost failed, expected (B=%ld, S=%ld, N=%ld), but got (%ld, %ld, %ld)", batchSize,
93+ seqLength, n, gradHPostShape.GetDim(0), gradHPostShape.GetDim(1), gradHPostShape.GetDim(2)),
94+ return ge::GRAPH_FAILED);
95+ 
96+ auto gradHResShapePtr = context->GetInputShape(GRAD_H_RES_IDX);
97+ OP_CHECK_IF(gradHResShapePtr == nullptr,
98+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradHRes shape is nullptr"),
99+ return ge::GRAPH_FAILED);
100+ auto gradHResShape = gradHResShapePtr->GetStorageShape();
101+ OP_CHECK_IF(gradHResShape.GetDimNum() != 4,
102+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradHRes must be 4D, but got %lu dims",
103+ gradHResShape.GetDimNum()),
104+ return ge::GRAPH_FAILED);
105+ OP_CHECK_IF(gradHResShape.GetDim(0) != batchSize || gradHResShape.GetDim(1) != seqLength ||
106+ gradHResShape.GetDim(2) != n || gradHResShape.GetDim(3) != n,
107+ OPS_REPORT_VECTOR_INNER_ERR(
108+ opName,
109+ "ShapeVerify gradHRes failed, expected (B=%ld, S=%ld, N=%ld, N=%ld), but got (%ld, %ld, %ld, %ld)",
110+ batchSize, seqLength, n, n, gradHResShape.GetDim(0), gradHResShape.GetDim(1),
111+ gradHResShape.GetDim(2), gradHResShape.GetDim(3)),
112+ return ge::GRAPH_FAILED);
113+ 
114+ auto phiShapePtr = context->GetInputShape(PHI_IDX);
115+ OP_CHECK_IF(phiShapePtr == nullptr, OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, phi shape is nullptr"),
116+ return ge::GRAPH_FAILED);
117+ auto phiShape = phiShapePtr->GetStorageShape();
118+ OP_CHECK_IF(phiShape.GetDimNum() != 2,
119+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, phi must be 2D, but got %lu dims",
120+ phiShape.GetDimNum()),
121+ return ge::GRAPH_FAILED);
122+ OP_CHECK_IF(phiShape.GetDim(0) != n * n + 2 * n || phiShape.GetDim(1) != n * c,
123+ OPS_REPORT_VECTOR_INNER_ERR(
124+ opName, "ShapeVerify phi failed, expected (2N+N^2=%ld, N*C=%ld), but got (%ld, %ld)", n * n + 2 * n,
125+ n * c, phiShape.GetDim(0), phiShape.GetDim(1)),
126+ return ge::GRAPH_FAILED);
127+ 
128+ auto alphaShapePtr = context->GetInputShape(ALPHA_IDX);
129+ OP_CHECK_IF(alphaShapePtr == nullptr,
130+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, alpha shape is nullptr"),
131+ return ge::GRAPH_FAILED);
132+ auto alphaShape = alphaShapePtr->GetStorageShape();
133+ OP_CHECK_IF(alphaShape.GetDimNum() != 1 || alphaShape.GetDim(0) != ALPHA_SIZE_3,
134+ OPS_REPORT_VECTOR_INNER_ERR(
135+ opName, "ShapeVerify alpha failed, expected (3), but got shape with %lu dims, dim0=%ld",
136+ alphaShape.GetDimNum(), alphaShape.GetDim(0)),
137+ return ge::GRAPH_FAILED);
138+ 
139+ auto biasShapePtr = context->GetInputShape(BIAS_IDX);
140+ OP_CHECK_IF(biasShapePtr == nullptr,
141+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, bias shape is nullptr"),
142+ return ge::GRAPH_FAILED);
143+ auto biasShape = biasShapePtr->GetStorageShape();
144+ OP_CHECK_IF(biasShape.GetDimNum() != 1 || biasShape.GetDim(0) != n * n + 2 * n,
145+ OPS_REPORT_VECTOR_INNER_ERR(
146+ opName, "ShapeVerify bias failed, expected (2N+N^2=%ld), but got shape with %lu dims, dim0=%ld",
147+ n * n + 2 * n, biasShape.GetDimNum(), biasShape.GetDim(0)),
148+ return ge::GRAPH_FAILED);
149+ 
150+ auto hPreShapePtr = context->GetInputShape(H_PRE_IDX);
151+ OP_CHECK_IF(hPreShapePtr == nullptr,
152+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, hPre shape is nullptr"),
153+ return ge::GRAPH_FAILED);
154+ auto hPreShape = hPreShapePtr->GetStorageShape();
155+ OP_CHECK_IF(hPreShape.GetDimNum() != 3,
156+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, hPre must be 3D, but got %lu dims",
157+ hPreShape.GetDimNum()),
158+ return ge::GRAPH_FAILED);
159+ OP_CHECK_IF(hPreShape.GetDim(0) != batchSize || hPreShape.GetDim(1) != seqLength || hPreShape.GetDim(2) != n,
160+ OPS_REPORT_VECTOR_INNER_ERR(
161+ opName, "ShapeVerify hPre failed, expected (B=%ld, S=%ld, N=%ld), but got (%ld, %ld, %ld)",
162+ batchSize, seqLength, n, hPreShape.GetDim(0), hPreShape.GetDim(1), hPreShape.GetDim(2)),
163+ return ge::GRAPH_FAILED);
164+ 
165+ auto hcBeforeNormShapePtr = context->GetInputShape(HC_BEFORE_NORM_IDX);
166+ OP_CHECK_IF(hcBeforeNormShapePtr == nullptr,
167+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, hcBeforeNorm shape is nullptr"),
168+ return ge::GRAPH_FAILED);
169+ auto hcBeforeNormShape = hcBeforeNormShapePtr->GetStorageShape();
170+ OP_CHECK_IF(hcBeforeNormShape.GetDimNum() != 3,
171+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, hcBeforeNorm must be 3D, but got %lu dims",
172+ hcBeforeNormShape.GetDimNum()),
173+ return ge::GRAPH_FAILED);
174+ OP_CHECK_IF(hcBeforeNormShape.GetDim(0) != batchSize || hcBeforeNormShape.GetDim(1) != seqLength ||
175+ hcBeforeNormShape.GetDim(2) != n * n + 2 * n,
176+ OPS_REPORT_VECTOR_INNER_ERR(
177+ opName,
178+ "ShapeVerify hcBeforeNorm failed, expected (B=%ld, S=%ld, N^2+2N=%ld), but got (%ld, %ld, %ld)",
179+ batchSize, seqLength, n * n + 2 * n, hcBeforeNormShape.GetDim(0), hcBeforeNormShape.GetDim(1),
180+ hcBeforeNormShape.GetDim(2)),
181+ return ge::GRAPH_FAILED);
182+ 
183+ auto invRmsShapePtr = context->GetInputShape(INV_RMS_IDX);
184+ OP_CHECK_IF(invRmsShapePtr == nullptr,
185+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, invRms shape is nullptr"),
186+ return ge::GRAPH_FAILED);
187+ auto invRmsShape = invRmsShapePtr->GetStorageShape();
188+ OP_CHECK_IF(invRmsShape.GetDimNum() != 3,
189+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, invRms must be 3D, but got %lu dims",
190+ invRmsShape.GetDimNum()),
191+ return ge::GRAPH_FAILED);
192+ OP_CHECK_IF(invRmsShape.GetDim(0) != batchSize || invRmsShape.GetDim(1) != seqLength || invRmsShape.GetDim(2) != 1,
193+ OPS_REPORT_VECTOR_INNER_ERR(
194+ opName, "ShapeVerify invRms failed, expected (B=%ld, S=%ld, 1), but got (%ld, %ld, %ld)", batchSize,
195+ seqLength, invRmsShape.GetDim(0), invRmsShape.GetDim(1), invRmsShape.GetDim(2)),
196+ return ge::GRAPH_FAILED);
197+ 
198+ auto sumOutShapePtr = context->GetInputShape(SUM_OUT_IDX);
199+ OP_CHECK_IF(sumOutShapePtr == nullptr,
200+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, sumOut shape is nullptr"),
201+ return ge::GRAPH_FAILED);
202+ auto sumOutShape = sumOutShapePtr->GetStorageShape();
203+ OP_CHECK_IF(sumOutShape.GetDimNum() != 4,
204+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, sumOut must be 4D, but got %lu dims",
205+ sumOutShape.GetDimNum()),
206+ return ge::GRAPH_FAILED);
207+ OP_CHECK_IF(sumOutShape.GetDim(1) != batchSize || sumOutShape.GetDim(2) != seqLength || sumOutShape.GetDim(3) != n,
208+ OPS_REPORT_VECTOR_INNER_ERR(
209+ opName,
210+ "ShapeVerify sumOut failed, expected (2*iter, B=%ld, S=%ld, N=%ld), but got (%ld, %ld, %ld, %ld)",
211+ batchSize, seqLength, n, sumOutShape.GetDim(0), sumOutShape.GetDim(1), sumOutShape.GetDim(2),
212+ sumOutShape.GetDim(3)),
213+ return ge::GRAPH_FAILED);
214+ OP_CHECK_IF(sumOutShape.GetDim(0) % 2 != 0,
215+ OPS_REPORT_VECTOR_INNER_ERR(opName,
216+ "ShapeVerify sumOut failed, dim0 must be even (2*iter_count), but got %ld",
217+ sumOutShape.GetDim(0)),
218+ return ge::GRAPH_FAILED);
219+ 
220+ auto normOutShapePtr = context->GetInputShape(NORM_OUT_IDX);
221+ OP_CHECK_IF(normOutShapePtr == nullptr,
222+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, normOut shape is nullptr"),
223+ return ge::GRAPH_FAILED);
224+ auto normOutShape = normOutShapePtr->GetStorageShape();
225+ OP_CHECK_IF(normOutShape.GetDimNum() != 5,
226+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, normOut must be 5D, but got %lu dims",
227+ normOutShape.GetDimNum()),
228+ return ge::GRAPH_FAILED);
229+ OP_CHECK_IF(normOutShape.GetDim(1) != batchSize || normOutShape.GetDim(2) != seqLength ||
230+ normOutShape.GetDim(3) != n || normOutShape.GetDim(4) != n,
231+ OPS_REPORT_VECTOR_INNER_ERR(opName,
232+ "ShapeVerify normOut failed, expected (2*iter, B=%ld, S=%ld, N=%ld, "
233+ "N=%ld), but got (%ld, %ld, %ld, %ld, %ld)",
234+ batchSize, seqLength, n, n, normOutShape.GetDim(0), normOutShape.GetDim(1),
235+ normOutShape.GetDim(2), normOutShape.GetDim(3), normOutShape.GetDim(4)),
236+ return ge::GRAPH_FAILED);
237+ OP_CHECK_IF(normOutShape.GetDim(0) != sumOutShape.GetDim(0),
238+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify normOut dim0(%ld) must equal sumOut dim0(%ld)",
239+ normOutShape.GetDim(0), sumOutShape.GetDim(0)),
240+ return ge::GRAPH_FAILED);
241+ 
242+ auto gradXShapePtr = context->GetOutputShape(GRAD_X_IDX);
243+ OP_CHECK_IF(gradXShapePtr == nullptr,
244+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradX shape is nullptr"),
245+ return ge::GRAPH_FAILED);
246+ auto gradXShape = gradXShapePtr->GetStorageShape();
247+ OP_CHECK_IF(gradXShape.GetDimNum() != 4,
248+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradX must be 4D, but got %lu dims",
249+ gradXShape.GetDimNum()),
250+ return ge::GRAPH_FAILED);
251+ OP_CHECK_IF(gradXShape.GetDim(0) != batchSize || gradXShape.GetDim(1) != seqLength || gradXShape.GetDim(2) != n ||
252+ gradXShape.GetDim(3) != c,
253+ OPS_REPORT_VECTOR_INNER_ERR(
254+ opName,
255+ "ShapeVerify gradX failed, expected (B=%ld, S=%ld, N=%ld, C=%ld), but got (%ld, %ld, %ld, %ld)",
256+ batchSize, seqLength, n, c, gradXShape.GetDim(0), gradXShape.GetDim(1), gradXShape.GetDim(2),
257+ gradXShape.GetDim(3)),
258+ return ge::GRAPH_FAILED);
259+ 
260+ auto gradPhiShapePtr = context->GetOutputShape(GRAD_PHI_IDX);
261+ OP_CHECK_IF(gradPhiShapePtr == nullptr,
262+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradPhi shape is nullptr"),
263+ return ge::GRAPH_FAILED);
264+ auto gradPhiShape = gradPhiShapePtr->GetStorageShape();
265+ OP_CHECK_IF(gradPhiShape.GetDimNum() != 2,
266+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradPhi must be 2D, but got %lu dims",
267+ gradPhiShape.GetDimNum()),
268+ return ge::GRAPH_FAILED);
269+ OP_CHECK_IF(gradPhiShape.GetDim(0) != n * n + 2 * n || gradPhiShape.GetDim(1) != n * c,
270+ OPS_REPORT_VECTOR_INNER_ERR(
271+ opName, "ShapeVerify gradPhi failed, expected (2N+N^2=%ld, N*C=%ld), but got (%ld, %ld)",
272+ n * n + 2 * n, n * c, gradPhiShape.GetDim(0), gradPhiShape.GetDim(1)),
273+ return ge::GRAPH_FAILED);
274+ 
275+ auto gradAlphaShapePtr = context->GetOutputShape(GRAD_ALPHA_IDX);
276+ OP_CHECK_IF(gradAlphaShapePtr == nullptr,
277+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradAlpha shape is nullptr"),
278+ return ge::GRAPH_FAILED);
279+ auto gradAlphaShape = gradAlphaShapePtr->GetStorageShape();
280+ OP_CHECK_IF(gradAlphaShape.GetDimNum() != 1 || gradAlphaShape.GetDim(0) != ALPHA_SIZE_3,
281+ OPS_REPORT_VECTOR_INNER_ERR(
282+ opName, "ShapeVerify gradAlpha failed, expected (3), but got shape with %lu dims, dim0=%ld",
283+ gradAlphaShape.GetDimNum(), gradAlphaShape.GetDim(0)),
284+ return ge::GRAPH_FAILED);
285+ 
286+ auto gradBiasShapePtr = context->GetOutputShape(GRAD_BIAS_IDX);
287+ OP_CHECK_IF(gradBiasShapePtr == nullptr,
288+ OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, gradBias shape is nullptr"),
289+ return ge::GRAPH_FAILED);
290+ auto gradBiasShape = gradBiasShapePtr->GetStorageShape();
291+ OP_CHECK_IF(gradBiasShape.GetDimNum() != 1 || gradBiasShape.GetDim(0) != n * n + 2 * n,
292+ OPS_REPORT_VECTOR_INNER_ERR(
293+ opName, "ShapeVerify gradBias failed, expected (2N+N^2=%ld), but got shape with %lu dims, dim0=%ld",
294+ n * n + 2 * n, gradBiasShape.GetDimNum(), gradBiasShape.GetDim(0)),
295+ return ge::GRAPH_FAILED);
296+ 
297+ OP_CHECK_IF(n != N_SIZE_4, OPS_REPORT_VECTOR_INNER_ERR(opName, "ShapeVerify failed, N must be 4, but got %ld", n),
298+ return ge::GRAPH_FAILED);
299+ 
300+ return ge::GRAPH_SUCCESS;
301+}
302+ 
303+ge::graphStatus TilingMhcPreSinkhornBackward(gert::TilingContext *context)
304+{
305+ if (context == nullptr) {
306+ OP_LOGE(context, "context is nullptr");
307+ return ge::GRAPH_FAILED;
308+ }
309+ auto platformInfoPtr = context->GetPlatformInfo();
310+ OP_CHECK_IF(platformInfoPtr == nullptr,
311+ OPS_REPORT_VECTOR_INNER_ERR(context->GetNodeName(), "platformInfoPtr is null!"),
312+ return ge::GRAPH_FAILED);
313+ auto ascendPlatformInfo = platform_ascendc::PlatformAscendC(platformInfoPtr);
314+ uint64_t ubSize;
315+ ascendPlatformInfo.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
316+ auto aicNum = ascendPlatformInfo.GetCoreNumAic();
317+ auto aivNum = ascendPlatformInfo.GetCoreNumAiv();
318+ if (aicNum == 0 || aivNum == 0) {
319+ OP_LOGE(context, "aicNum=%lu or aivNum=%lu is invalid", aicNum, aivNum);
320+ return ge::GRAPH_FAILED;
321+ }
322+ context->SetBlockDim(aicNum);
323+ const auto xShapePtr = context->GetInputShape(INPUT_X_IDX);
324+ if (xShapePtr == nullptr) {
325+ OP_LOGE(context, "input x shape is nullptr");
326+ return ge::GRAPH_FAILED;
327+ }
328+ auto xShape = xShapePtr->GetStorageShape();
329+ auto attrsPtr = context->GetAttrs();
330+ if (attrsPtr == nullptr) {
331+ OP_LOGE(context, "attrs is nullptr");
332+ return ge::GRAPH_FAILED;
333+ }
334+ auto epsPtr = attrsPtr->GetAttrPointer<float>(0);
335+ float eps = (epsPtr != nullptr) ? static_cast<float>(*epsPtr) : DEFAULT_EPS;
336+ 
337+ int64_t batchSize = xShape.GetDim(BATCH_SIZE_DIM_IDX);
338+ int64_t seqLength = xShape.GetDim(SEQ_LENGTH_DIM_IDX);
339+ int64_t n = xShape.GetDim(N_DIM_IDX);
340+ int64_t c = xShape.GetDim(C_DIM_IDX);
341+ 
342+ OP_CHECK_IF(ShapeVerify(context, batchSize, seqLength, n, c) != ge::GRAPH_SUCCESS,
343+ OPS_REPORT_VECTOR_INNER_ERR(context->GetNodeName(), "ShapeVerify failed"), return ge::GRAPH_FAILED);
344+ 
345+ int64_t c0 = 256; // must be VL
346+ int64_t c1 = c / c0;
347+ int64_t cTail = c % c0;
348+ int64_t cTailunaglin = cTail % 64 + 64 * cTail / (64 + 128);
349+ 
350+ int64_t tile = 32;
351+ const auto sumOutPtr = context->GetInputShape(10);
352+ CHECK_NULLPTR(sumOutPtr);
353+ auto sumOutShape = sumOutPtr->GetStorageShape();
354+ int64_t skIterCount = sumOutShape.GetDim(0) / C_V_RATIO;
355+ 
356+ int64_t mm1K = n * n + 2 * n;
357+ int64_t mm1M = tile * 2;
358+ int64_t mm1N = n * c;
359+ 
360+ int64_t mm2K = tile * 2;
361+ int64_t mm2M = n * n + 2 * n;
362+ int64_t mm2N = n * c;
363+ 
364+ MhcPreSinkhornBackwardTilingData tilingData;
365+ auto featureDataType = matmul_tiling::DataType::DT_FLOAT;
366+ matmul_tiling::MatmulApiTiling mm1Tiling(ascendPlatformInfo);
367+ matmul_tiling::MatmulApiTiling mm2Tiling(ascendPlatformInfo);
368+ 
369+ mm1Tiling.SetAType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, featureDataType);
370+ mm1Tiling.SetBType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, featureDataType);
371+ mm1Tiling.SetCType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, featureDataType);
372+ mm1Tiling.SetOrgShape(mm1M, mm1N, mm1K);
373+ mm1Tiling.SetShape(mm1M, mm1N, mm1K);
374+ mm1Tiling.SetBias(false);
375+ mm1Tiling.SetBufferSpace(-1, -1, -1);
376+ if (mm1Tiling.GetTiling(tilingData.mm1TilingData) == -1) {
377+ OP_LOGE(context, "mm1Tiling.GetTiling failed, M=%ld, N=%ld, K=%ld", mm1M, mm1N, mm1K);
378+ return ge::GRAPH_FAILED;
379+ }
380+ 
381+ mm2Tiling.SetAType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, featureDataType, true);
382+ mm2Tiling.SetBType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, featureDataType);
383+ mm2Tiling.SetCType(matmul_tiling::TPosition::GM, matmul_tiling::CubeFormat::ND, featureDataType);
384+ mm2Tiling.SetOrgShape(mm2M, mm2N, mm2K);
385+ mm2Tiling.SetShape(mm2M, mm2N, mm2K);
386+ mm2Tiling.SetBias(false);
387+ mm2Tiling.SetBufferSpace(-1, -1, -1);
388+ if (mm2Tiling.GetTiling(tilingData.mm2TilingData) == -1) {
389+ OP_LOGE(context, "mm2Tiling.GetTiling failed, M=%ld, N=%ld, K=%ld", mm2M, mm2N, mm2K);
390+ return ge::GRAPH_FAILED;
391+ }
392+ 
393+ tilingData.set_n(n);
394+ tilingData.set_batchSize(batchSize);
395+ tilingData.set_seqLength(seqLength);
396+ tilingData.set_c(c);
397+ tilingData.set_cTail(cTail);
398+ tilingData.set_n(n);
399+ tilingData.set_c0(c0);
400+ tilingData.set_c1(c1);
401+ tilingData.set_aivNum(aivNum);
402+ tilingData.set_ubSize(ubSize);
403+ tilingData.set_skIterCount(skIterCount);
404+ tilingData.set_eps(eps);
405+ 
406+ tilingData.set_tileSize(tile);
407+ 
408+ if (context->GetRawTilingData() == nullptr) {
409+ OP_LOGE(context, "GetRawTilingData() is nullptr");
410+ return ge::GRAPH_FAILED;
411+ }
412+ tilingData.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity());
413+ context->GetRawTilingData()->SetDataSize(tilingData.GetDataSize());
414+ 
415+ size_t xCastWorkspace = batchSize * seqLength * n * c * sizeof(float) * 2;
416+ size_t gradHat2Workspace = batchSize * seqLength * (n * n + 2 * n) * sizeof(float);
417+ 
418+ size_t systemWorkspaceSize = ascendPlatformInfo.GetLibApiWorkSpaceSize();
419+ size_t usrWorkSpaceSize = xCastWorkspace + gradHat2Workspace;
420+ size_t *currentWorkspace = context->GetWorkspaceSizes(1);
421+ if (currentWorkspace == nullptr) {
422+ OP_LOGE(context, "GetWorkspaceSizes() returned nullptr");
423+ return ge::GRAPH_FAILED;
424+ }
425+ 
426+ currentWorkspace[0] = systemWorkspaceSize + usrWorkSpaceSize;
427+ return ge::GRAPH_SUCCESS;
428+}
429+ 
430+static ge::graphStatus TilingParseForMhcPreSinkhornBackward(gert::TilingParseContext *context)
431+{
432+ return ge::GRAPH_SUCCESS;
433+}
434+ 
435+struct MhcPreSinkhornBackwardCompileInfo {
436+};
437+ 
438+IMPL_OP_OPTILING(MhcPreSinkhornBackward)
439+ .Tiling(TilingMhcPreSinkhornBackward)
440+ .TilingParse<MhcPreSinkhornBackwardCompileInfo>(TilingParseForMhcPreSinkhornBackward); // 向框架注册入口函数
441+} // namespace optiling
@@ -0,0 +1,46 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file mhc_pre_sinkhorn_backward_tiling.h
13+ * \brief MhcPreSinkhornBackward operator tiling data definition
14+ */
15+#ifndef OP_HOST_OP_TILING_ARCH32_MHC_PRE_SINKHORN_BACKWARD_TILING_H
16+#define OP_HOST_OP_TILING_ARCH32_MHC_PRE_SINKHORN_BACKWARD_TILING_H
17+#include <tiling/tiling_api.h>
18+#include "register/tilingdata_base.h"
19+#include "tiling_base/tiling_base.h"
20+#include "err/ops_err.h"
21+ 
22+namespace optiling {
23+BEGIN_TILING_DATA_DEF(MhcPreSinkhornBackwardTilingData)
24+TILING_DATA_FIELD_DEF(int64_t, batchSize)
25+TILING_DATA_FIELD_DEF(int64_t, seqLength)
26+TILING_DATA_FIELD_DEF(int64_t, c)
27+TILING_DATA_FIELD_DEF(int64_t, n)
28+TILING_DATA_FIELD_DEF(int64_t, c0) // tile for c
29+TILING_DATA_FIELD_DEF(int64_t, c1) // tile count of c
30+TILING_DATA_FIELD_DEF(int64_t, cTail) // tail of c
31+TILING_DATA_FIELD_DEF(int64_t, aivNum) // tile count of c
32+TILING_DATA_FIELD_DEF(int64_t, tileGradY)
33+TILING_DATA_FIELD_DEF(int64_t, tileHHat2)
34+TILING_DATA_FIELD_DEF(int64_t, tileSize)
35+TILING_DATA_FIELD_DEF(int64_t, skIterCount)
36+TILING_DATA_FIELD_DEF(int64_t, ubSize)
37+TILING_DATA_FIELD_DEF(float, eps)
38+ 
39+TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, mm1TilingData)
40+TILING_DATA_FIELD_DEF_STRUCT(TCubeTiling, mm2TilingData)
41+END_TILING_DATA_DEF
42+ 
43+REGISTER_TILING_DATA_CLASS(MhcPreSinkhornBackward, MhcPreSinkhornBackwardTilingData)
44+} // namespace optiling
45+ 
46+#endif // OP_HOST_OP_TILING_ARCH32_MHC_PRE_SINKHORN_BACKWARD_TILING_H
@@ -0,0 +1,1008 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file MhcPreGradKernel.h
13+ * \brief
14+ */
15+#ifndef MHC_PRE_SINKHORN_BACKWARD_OP_KERNEL_MHC_PRE_GRAD_KERNEL_H
16+#define MHC_PRE_SINKHORN_BACKWARD_OP_KERNEL_MHC_PRE_GRAD_KERNEL_H
17+ 
18+#include "kernel_operator.h"
19+#include "lib/matmul_intf.h"
20+ 
21+using namespace AscendC;
22+ 
23+namespace {
24+constexpr int32_t BYTE_SIZE_PER_BLOCK = 32;
25+constexpr int32_t ELEMENTS_SIZE_PER_BLOCK = BYTE_SIZE_PER_BLOCK / sizeof(float);
26+constexpr int32_t BYTE_SIZE_PER_REPEAT = 256;
27+constexpr int32_t ELEMENTS_SIZE_PER_REPEAT = 256 / sizeof(float);
28+constexpr int32_t REPEAT_LENTH = ELEMENTS_SIZE_PER_REPEAT;
29+constexpr int32_t BLOCK_PER_REPEAT = 8;
30+constexpr uint64_t MASK_PRE[] = {0b0000111100001111000011110000111100001111000011110000111100001111};
31+constexpr uint64_t MASK_POST[] = {0b1111000011110000111100001111000011110000111100001111000011110000};
32+constexpr uint64_t MASK_POST_SCALE[] = {0b0000000000000000000000000000000000000000000000000000000011110000};
33+constexpr int32_t PING_PONG_NUM = 2;
34+constexpr int32_t PRE_POST_NUM = 2;
35+constexpr int32_t DOUBLE_RATIO = 2;
36+ 
37+constexpr int32_t INNER_SPILT_NUM = 8;
38+ 
39+constexpr MatmulConfig MHC_PRE_GRAD_MM1_CFG = GetMDLConfig(false, false, 0, false, false, false, true);
40+constexpr MatmulConfig MHC_PRE_GRAD_MM2_CFG = GetMDLConfig(false, false, 0, false, false, false, true);
41+} // namespace
42+ 
43+template <typename TYPE_X, typename T>
44+class MhcPreGradKernel {
45+public:
46+ using A0Type = matmul::MatmulType<TPosition::GM, CubeFormat::ND, T>;
47+ using A1Type = matmul::MatmulType<TPosition::GM, CubeFormat::ND, T, true>;
48+ using BType = matmul::MatmulType<TPosition::GM, CubeFormat::ND, T>;
49+ using CType = matmul::MatmulType<TPosition::GM, CubeFormat::ND, T>;
50+ 
51+ matmul::MatmulImpl<A0Type, BType, CType, CType, MHC_PRE_GRAD_MM1_CFG> mm1_;
52+ matmul::MatmulImpl<A1Type, BType, CType, CType, MHC_PRE_GRAD_MM2_CFG> mm2_; // grad
53+ 
54+ __aicore__ inline MhcPreGradKernel() = default;
55+ 
56+ __aicore__ inline void Init(GM_ADDR x, GM_ADDR hc_fn, GM_ADDR pre, GM_ADDR grad_y, GM_ADDR grad_post,
57+ GM_ADDR grad_comb, GM_ADDR hc_scale, GM_ADDR hc_base, GM_ADDR h_hat2, GM_ADDR rsqrt,
58+ GM_ADDR sum_out, GM_ADDR norm_out, GM_ADDR grad_x, GM_ADDR grad_hc_fn,
59+ GM_ADDR grad_hc_scale, GM_ADDR grad_hc_base, GM_ADDR workspace,
60+ const MhcPreSinkhornBackwardTilingData *tilingData, TPipe *pipe)
61+ {
62+ pipe_ = pipe;
63+ blkIdx_ = GetBlockIdx();
64+ 
65+ InitTiling(tilingData);
66+ InitGM(x, hc_fn, pre, grad_y, grad_post, grad_comb, hc_scale, hc_base, h_hat2, rsqrt, sum_out, norm_out, grad_x,
67+ grad_hc_fn, grad_hc_scale, grad_hc_base, workspace);
68+ InitGradPreStageBuffer();
69+ }
70+ __aicore__ inline void Process();
71+ 
72+protected:
73+ int64_t blkIdx_, aivNum_, aicNum_;
74+ TPipe *pipe_;
75+ int64_t batchSize_, seqLength_, totalTasks_, totalTasksAligned_, BSNN, BSN;
76+ int64_t c_, n_, c0_, c1_, cTail_, c0RepeatTime_, cTailAlign_, cTailBlockStride_, c1Align_;
77+ int64_t tileCoreBS_;
78+ int64_t skIterCount_;
79+ int64_t ubSize_;
80+ int64_t mm1K_, mm1M_, mm1N_;
81+ int64_t mm2K_, mm2M_, mm2N_;
82+ int64_t tileRepeatTimes_;
83+ float eps_;
84+ event_t eventIdVToMTE3XCast;
85+ 
86+ T hcScalePre_, hcScalePost_, hcScaleRes_;
87+ 
88+ TQue<QuePosition::VECIN, 1> inputXInQueue;
89+ TQue<QuePosition::VECIN, 1> inputGradQueue;
90+ 
91+ TQue<QuePosition::VECIN, 1> SKInQueue;
92+ 
93+ TQue<QuePosition::VECOUT, 1> OutQueue;
94+ 
95+ TBuf<TPosition::VECCALC> fusedGradHPre2AndGradHPost2Buf_, gradRsqrtBuf_, gradBiasBuf_, onesBuf_, ScaleBuf_,
96+ hcBaseBuf_, tempBuf_;
97+ 
98+ LocalTensor<T> dBiasLocal_, gradRsqrtLocal_, dPrePostTempLocal_;
99+ LocalTensor<T> xCastLocal_, gradYCastLocal_, gradXCastLocal_;
100+ LocalTensor<T> scaleLocal_, dScaleLocal_;
101+ int32_t onceTask_;
102+ LocalTensor<T> gradHResLocal_;
103+ LocalTensor<T> hcBaseLocal_;
104+ LocalTensor<T> preBrcbLocal_, dRsqrtBrcbLocal_, rsqrtbrcbLocal_, tmpLocal_, hat2Scale, dhatBeforeNormLocal,
105+ gradHResTempLocal_, gradHResTempLocal2_, dhatLocal_, rsqrtTempLocal_, gradXCubeLocal_;
106+ 
107+ LocalTensor<T> hatLocal;
108+ LocalTensor<T> onesLocal_;
109+ 
110+ GlobalTensor<TYPE_X> xGlobal_, gradYGlobal_;
111+ GlobalTensor<T> preGlobal_;
112+ GlobalTensor<TYPE_X> gradXGlobal_;
113+ GlobalTensor<T> gradPreGlobal_, gradPostGlobal_;
114+ GlobalTensor<T> hcScaleGlobal_, hcBaseGlobal_;
115+ GlobalTensor<T> rsqrtGlobal_;
116+ GlobalTensor<T> gradHcScaleGlobal_, gradHcBaseGlobal_;
117+ GlobalTensor<T> h2Global_;
118+ GlobalTensor<T> skNormGlobal_, skSumGlobal_;
119+ GlobalTensor<T> gradHResGlobal_;
120+ GlobalTensor<T> gradH2Global_;
121+ GlobalTensor<T> gradWeightGlobal_;
122+ GlobalTensor<T> weightGlobal_;
123+ GlobalTensor<T> gradHcBaseWSGlobal_;
124+ GlobalTensor<T> gradHcScaleWSGlobal_;
125+ GlobalTensor<T> xWorkspaceGlobal_;
126+ GlobalTensor<T> gradXCubeGlobal_;
127+ 
128+private:
129+ __aicore__ inline void InitTiling(const MhcPreSinkhornBackwardTilingData *tilingData)
130+ {
131+ batchSize_ = tilingData->batchSize;
132+ seqLength_ = tilingData->seqLength;
133+ aivNum_ = tilingData->aivNum;
134+ aicNum_ = tilingData->aivNum / DOUBLE_RATIO;
135+ c_ = tilingData->c;
136+ n_ = tilingData->n;
137+ c0_ = tilingData->c0;
138+ c1_ = tilingData->c1;
139+ BSNN = batchSize_ * seqLength_ * n_ * n_;
140+ BSN = batchSize_ * seqLength_ * n_;
141+ c1Align_ = CeilDiv(c_, c0_);
142+ cTail_ = max((c_ - (c1Align_ - 1) * c0_), static_cast<int64_t>(0));
143+ 
144+ cTailAlign_ = AlignUp(cTail_, ELEMENTS_SIZE_PER_BLOCK);
145+ cTailBlockStride_ = c0_ / ELEMENTS_SIZE_PER_BLOCK - cTailAlign_ / ELEMENTS_SIZE_PER_BLOCK;
146+ skIterCount_ = tilingData->skIterCount;
147+ ubSize_ = tilingData->ubSize;
148+ eps_ = tilingData->eps;
149+ tileCoreBS_ = tilingData->tileSize;
150+ 
151+ c0RepeatTime_ = c0_ / ELEMENTS_SIZE_PER_REPEAT;
152+ totalTasks_ = batchSize_ * seqLength_;
153+ totalTasksAligned_ = AlignUp(totalTasks_, aivNum_ * tileCoreBS_);
154+ if ASCEND_IS_AIC {
155+ mm1K_ = n_ * n_ + PRE_POST_NUM * n_;
156+ mm1M_ = tileCoreBS_ * 2;
157+ mm1N_ = n_ * c_;
158+ 
159+ mm2K_ = batchSize_ * seqLength_;
160+ mm2M_ = n_ * n_ + PRE_POST_NUM * n_;
161+ mm2N_ = n_ * c_;
162+ }
163+ }
164+ 
165+ __aicore__ inline void InitGM(GM_ADDR x, GM_ADDR hc_fn, GM_ADDR pre, GM_ADDR grad_y, GM_ADDR grad_post,
166+ GM_ADDR grad_comb, GM_ADDR hc_scale, GM_ADDR hc_base, GM_ADDR h_hat2, GM_ADDR rsqrt,
167+ GM_ADDR sum_out, GM_ADDR norm_out, GM_ADDR grad_x, GM_ADDR grad_hc_fn,
168+ GM_ADDR grad_hc_scale, GM_ADDR grad_hc_base, GM_ADDR workspace)
169+ {
170+ // cube 和 vector 共用
171+ xGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ TYPE_X *>(x)); // fp32
172+ gradWeightGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(grad_hc_fn)); // fp32
173+ int64_t workspaceOffset = 0;
174+ gradH2Global_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(workspace) + workspaceOffset); // fp32
175+ workspaceOffset += batchSize_ * seqLength_ * (n_ * PRE_POST_NUM + n_ * n_);
176+ xWorkspaceGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(workspace) + workspaceOffset); // fp32
177+ workspaceOffset += batchSize_ * seqLength_ * (n_ * c_);
178+ gradXCubeGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(workspace) + workspaceOffset);
179+ 
180+ workspaceOffset += batchSize_ * seqLength_ * (n_ * c_);
181+ weightGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(hc_fn)); // fp32
182+ gradXGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ TYPE_X *>(grad_x)); // fp32
183+ 
184+ if ASCEND_IS_AIV {
185+ preGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(pre)); // fp32
186+ gradYGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ TYPE_X *>(grad_y)); // fp32
187+ 
188+ gradPostGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(grad_post)); // fp32
189+ 
190+ hcScaleGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(hc_scale)); // fp32
191+ hcBaseGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(hc_base)); // fp32
192+ 
193+ rsqrtGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(rsqrt)); // fp32
194+ gradHcScaleGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(grad_hc_scale));
195+ gradHcBaseGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(grad_hc_base));
196+ gradHcScaleWSGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(workspace) + workspaceOffset);
197+ workspaceOffset += aivNum_ * (n_ * 2 + n_ * n_);
198+ gradHcBaseWSGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(workspace) + workspaceOffset);
199+ workspaceOffset += aivNum_ * (n_ * 2 + n_ * n_);
200+ 
201+ h2Global_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(h_hat2)); // fp32
202+ 
203+ skNormGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(norm_out)); // fp32
204+ skSumGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(sum_out)); // fp32
205+ 
206+ gradHResGlobal_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(grad_comb)); // fp32
207+ if (blkIdx_ == aivNum_ - 1) {
208+ InitOutput<T>(gradHcBaseGlobal_, (n_ * n_ + PRE_POST_NUM * n_), 0);
209+ InitOutput<T>(gradHcScaleGlobal_, 3, 0);
210+ }
211+ for (int64_t taskOffset = blkIdx_ * tileCoreBS_; taskOffset < n_ * c_;
212+ taskOffset += aivNum_ * tileCoreBS_) {
213+ int32_t tileTaskCount =
214+ min(static_cast<int32_t>(tileCoreBS_), static_cast<int32_t>(n_ * c_ - taskOffset));
215+ InitOutput<T>(gradWeightGlobal_[taskOffset * (n_ * n_ + PRE_POST_NUM * n_)],
216+ (n_ * n_ + PRE_POST_NUM * n_) * tileTaskCount, 0);
217+ }
218+ SyncAll<true>();
219+ }
220+ }
221+ 
222+ __aicore__ inline void InitGradPreStageBuffer()
223+ {
224+ if ASCEND_IS_AIV {
225+ pipe_->InitBuffer(fusedGradHPre2AndGradHPost2Buf_, tileCoreBS_ * n_ * 2 * sizeof(float));
226+ pipe_->InitBuffer(gradRsqrtBuf_, tileCoreBS_ * (n_ * n_ + 2 * n_) * sizeof(float) * 2);
227+ pipe_->InitBuffer(gradBiasBuf_, 2 * tileCoreBS_ * (n_ * n_ + 2 * n_) * sizeof(float));
228+ pipe_->InitBuffer(onesBuf_, tileCoreBS_ * n_ * 2 * sizeof(float)); // pre: tileCoreBS_ * n
229+ pipe_->InitBuffer(hcBaseBuf_, (n_ * n_ + 2 * n_) * sizeof(float));
230+ pipe_->InitBuffer(ScaleBuf_, BYTE_SIZE_PER_BLOCK * 2); // pre: tileCoreBS_ * n
231+ pipe_->InitBuffer(inputXInQueue, 2, tileCoreBS_ * n_ * c0_ * sizeof(float) / 4);
232+ pipe_->InitBuffer(inputGradQueue, 1, tileCoreBS_ * n_ * ELEMENTS_SIZE_PER_BLOCK * sizeof(float));
233+ pipe_->InitBuffer(SKInQueue, 2,
234+ tileCoreBS_ * (n_ * ELEMENTS_SIZE_PER_BLOCK + ELEMENTS_SIZE_PER_BLOCK * 2) *
235+ sizeof(float));
236+ pipe_->InitBuffer(OutQueue, 2, tileCoreBS_ * n_ * c0_ * sizeof(float) / 8);
237+ auto ubSizeRemain =
238+ CeilDiv(tileCoreBS_, ELEMENTS_SIZE_PER_BLOCK) * ELEMENTS_SIZE_PER_REPEAT * sizeof(float) +
239+ CeilDiv(tileCoreBS_ * n_, ELEMENTS_SIZE_PER_BLOCK) * ELEMENTS_SIZE_PER_REPEAT * sizeof(float) +
240+ onceTask_ * n_ * c0_ * sizeof(float) * 2 + onceTask_ * c0_ * sizeof(float);
241+ 
242+ pipe_->InitBuffer(tempBuf_, ubSizeRemain); // grad_y: n_ * c0_
243+ 
244+ eventIdVToMTE3XCast = static_cast<event_t>(GetTPipePtr()->AllocEventID<HardEvent::V_MTE3>());
245+ dPrePostTempLocal_ = fusedGradHPre2AndGradHPost2Buf_.Get<T>();
246+ hcBaseLocal_ = hcBaseBuf_.Get<T>();
247+ onesLocal_ = onesBuf_.Get<T>();
248+ scaleLocal_ = ScaleBuf_.Get<T>();
249+ gradRsqrtLocal_ = gradRsqrtBuf_.Get<T>();
250+ dBiasLocal_ = gradBiasBuf_.Get<T>();
251+ dScaleLocal_ = dBiasLocal_[tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_)];
252+ 
253+ onceTask_ = tileCoreBS_ / INNER_SPILT_NUM;
254+ 
255+ int32_t offset = 0;
256+ int32_t brcbAlign = CeilDiv(tileCoreBS_ * n_, ELEMENTS_SIZE_PER_BLOCK);
257+ preBrcbLocal_ = tempBuf_.GetWithOffset<T>(brcbAlign * ELEMENTS_SIZE_PER_REPEAT, offset);
258+ offset += brcbAlign * ELEMENTS_SIZE_PER_REPEAT * sizeof(float);
259+ brcbAlign = CeilDiv(tileCoreBS_, ELEMENTS_SIZE_PER_BLOCK);
260+ dRsqrtBrcbLocal_ = tempBuf_.GetWithOffset<T>(brcbAlign * ELEMENTS_SIZE_PER_REPEAT, offset);
261+ offset += brcbAlign * ELEMENTS_SIZE_PER_REPEAT * sizeof(float);
262+ gradYCastLocal_ = tempBuf_.GetWithOffset<T>(onceTask_ * c0_, offset);
263+ offset += onceTask_ * c0_ * sizeof(float);
264+ gradXCastLocal_ = tempBuf_.GetWithOffset<T>(onceTask_ * n_ * c0_, offset);
265+ offset += onceTask_ * n_ * c0_ * sizeof(float);
266+ xCastLocal_ = tempBuf_.GetWithOffset<T>(onceTask_ * n_ * c0_, offset);
267+ offset = 0;
268+ // ComputeGradPre
269+ tmpLocal_ = gradXCastLocal_;
270+ 
271+ // SinkhornGrad
272+ gradHResTempLocal_ =
273+ tempBuf_.GetWithOffset<T>(tileCoreBS_ * n_ * ELEMENTS_SIZE_PER_BLOCK * 2, offset); // int64大小
274+ offset += tileCoreBS_ * n_ * ELEMENTS_SIZE_PER_BLOCK * 2 * sizeof(float);
275+ gradHResTempLocal2_ = tempBuf_.GetWithOffset<T>(tileCoreBS_ * n_ * ELEMENTS_SIZE_PER_BLOCK, offset);
276+ offset = 0;
277+ 
278+ dhatLocal_ = tempBuf_.GetWithOffset<T>(tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_), offset);
279+ offset += tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_) * sizeof(float);
280+ brcbAlign = CeilDiv(tileCoreBS_, ELEMENTS_SIZE_PER_BLOCK);
281+ rsqrtbrcbLocal_ =
282+ tempBuf_.GetWithOffset<T>(brcbAlign * ELEMENTS_SIZE_PER_BLOCK * (n_ * n_ + PRE_POST_NUM * n_), offset);
283+ offset += brcbAlign * ELEMENTS_SIZE_PER_BLOCK * (n_ * n_ + PRE_POST_NUM * n_) * sizeof(float);
284+ hat2Scale = tempBuf_.GetWithOffset<T>(tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_), offset);
285+ offset += tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_) * sizeof(float);
286+ dhatBeforeNormLocal = tempBuf_.GetWithOffset<T>(tileCoreBS_ * (2 * n_), offset);
287+ offset += tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_) * sizeof(float);
288+ hatLocal = tempBuf_.GetWithOffset<T>(tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_), offset);
289+ offset += tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_) * sizeof(float);
290+ rsqrtTempLocal_ = tempBuf_.GetWithOffset<T>(tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_), offset);
291+ 
292+ Duplicate(onesLocal_, 1.f, tileCoreBS_ * n_ * 2);
293+ Duplicate(dScaleLocal_, 0.f, tileCoreBS_ * (2 * n_ + n_ * n_));
294+ Duplicate(dBiasLocal_, 0.f, tileCoreBS_ * (2 * n_ + n_ * n_));
295+ }
296+ }
297+ 
298+ __aicore__ inline void ComputeGradPre(const int32_t taskOffset, const int32_t tileTaskCount, const int32_t innerId);
299+ 
300+ __aicore__ inline void ComputeGradHHat2(const int32_t taskOffset, const int32_t tileTaskCount);
301+ __aicore__ inline void SinkhornGrad(const int32_t taskOffset, const int32_t tileTaskCount);
302+ 
303+ __aicore__ inline void ComputeGradX1(const int32_t taskOffset, const int32_t tileTaskCount, const int32_t innerId);
304+ __aicore__ inline void GetHcScaleAndHcBase();
305+ __aicore__ inline void ProcessMatmul1(const int32_t taskOffset, const int32_t mm1M);
306+ __aicore__ inline void ProcessMatmul2(const int32_t taskOffset, const int32_t mm2K);
307+ __aicore__ inline void ComputeScaleBias();
308+};
309+ 
310+template <typename TYPE_X, typename T>
311+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::GetHcScaleAndHcBase()
312+{
313+ // HcBase 常驻 UB
314+ hcScalePre_ = hcScaleGlobal_.GetValue(0);
315+ hcScalePost_ = hcScaleGlobal_.GetValue(1);
316+ hcScaleRes_ = hcScaleGlobal_.GetValue(2);
317+ event_t eventIDSToV = static_cast<event_t>(GetTPipePtr()->FetchEventID(HardEvent::S_V));
318+ SetFlag<HardEvent::S_V>(eventIDSToV);
319+ WaitFlag<HardEvent::S_V>(eventIDSToV);
320+ Duplicate(scaleLocal_[8], hcScaleRes_, 8);
321+ Duplicate(scaleLocal_, hcScalePost_, 8);
322+ PipeBarrier<PIPE_V>();
323+ 
324+ Duplicate(scaleLocal_, hcScalePre_, 4);
325+ 
326+ DataCopyPad(hcBaseLocal_, hcBaseGlobal_,
327+ {static_cast<uint16_t>(1), static_cast<uint32_t>((n_ * n_ + PRE_POST_NUM * n_) * sizeof(T)), 0, 0, 0},
328+ {false, 0, 0, 0});
329+}
330+ 
331+template <typename TYPE_X, typename T>
332+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::SinkhornGrad(const int32_t taskOffset, const int32_t tileTaskCount)
333+{
334+ gradHResLocal_ = inputGradQueue.AllocTensor<T>();
335+ 
336+ DataCopyPad(gradHResLocal_, gradHResGlobal_[taskOffset * n_ * n_],
337+ {static_cast<uint16_t>(tileTaskCount * n_), static_cast<uint32_t>(n_ * sizeof(T)), 0, 0, 0},
338+ {true, 0, static_cast<uint8_t>(ELEMENTS_SIZE_PER_BLOCK - n_), 0});
339+ inputGradQueue.EnQue(gradHResLocal_);
340+ inputGradQueue.DeQue();
341+ 
342+ int64_t iterRowNormOffset = (skIterCount_ - 1) * 2 * BSNN + taskOffset * n_ * n_;
343+ int64_t iterColNormOffset = ((skIterCount_ - 1) * 2 + 1) * BSNN + taskOffset * n_ * n_;
344+ 
345+ int64_t iterRowSumOffset = (skIterCount_ - 1) * 2 * BSN + taskOffset * n_;
346+ int64_t iterColSumOffset = ((skIterCount_ - 1) * 2 + 1) * BSN + taskOffset * n_;
347+ int32_t brcbAlign = CeilDiv(tileTaskCount * n_, ELEMENTS_SIZE_PER_BLOCK);
348+ for (int32_t iter = skIterCount_ - 1; iter > 0; iter--) {
349+ auto skRowNormLocal_ = SKInQueue.AllocTensor<T>();
350+ auto skRowSumLocal_ = skRowNormLocal_[tileCoreBS_ * n_ * ELEMENTS_SIZE_PER_BLOCK];
351+ auto skColSumLocal_ =
352+ skRowNormLocal_[tileCoreBS_ * n_ * ELEMENTS_SIZE_PER_BLOCK + tileCoreBS_ * ELEMENTS_SIZE_PER_BLOCK];
353+ 
354+ DataCopyPad(skRowNormLocal_, skNormGlobal_[iterRowNormOffset],
355+ {static_cast<uint16_t>(tileTaskCount * n_), static_cast<uint32_t>(n_ * sizeof(T)), 0, 0, 0},
356+ {true, 0, static_cast<uint8_t>(ELEMENTS_SIZE_PER_BLOCK - n_), 0});
357+ DataCopyPad(skColSumLocal_, skSumGlobal_[iterColSumOffset],
358+ {static_cast<uint16_t>(tileTaskCount), static_cast<uint32_t>(n_ * sizeof(T)), 0, 0, 0},
359+ {true, 0, static_cast<uint8_t>(ELEMENTS_SIZE_PER_BLOCK - n_), 0});
360+ DataCopyPad(skRowSumLocal_, skSumGlobal_[iterRowSumOffset],
361+ {static_cast<uint16_t>(1), static_cast<uint32_t>(tileTaskCount * n_ * sizeof(T)), 0, 0, 0},
362+ {false, 0, static_cast<uint8_t>(ELEMENTS_SIZE_PER_BLOCK - n_), 0});
363+ SKInQueue.EnQue(skRowNormLocal_);
364+ SKInQueue.DeQue();
365+ PipeBarrier<PIPE_V>();
366+ 
367+ Adds(skColSumLocal_, skColSumLocal_, eps_, tileTaskCount * ELEMENTS_SIZE_PER_BLOCK);
368+ PipeBarrier<PIPE_V>();
369+ 
370+ 
371+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
372+ Div(gradHResTempLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK],
373+ gradHResLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK], skColSumLocal_, ELEMENTS_SIZE_PER_REPEAT,
374+ tileRepeatTimes_,
375+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), 1, static_cast<uint8_t>(n_ * 8),
376+ static_cast<uint8_t>(n_ * 8), 8});
377+ }
378+ 
379+ PipeBarrier<PIPE_V>();
380+ // xg * normed
381+ Mul(gradHResTempLocal2_, gradHResLocal_, skRowNormLocal_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
382+ PipeBarrier<PIPE_V>();
383+ // (sum + eps) ** 2
384+ Mul(skColSumLocal_, skColSumLocal_, skColSumLocal_, tileTaskCount * ELEMENTS_SIZE_PER_BLOCK);
385+ // bs * n *8 reduce bs * 1 *8 out : value other other other value other other other
386+ // sum(xg * normed )
387+ 
388+ for (int32_t loopIdN = 1; loopIdN < n_; loopIdN += 1) {
389+ Add(gradHResTempLocal2_, gradHResTempLocal2_, gradHResTempLocal2_[loopIdN * ELEMENTS_SIZE_PER_BLOCK],
390+ ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
391+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), static_cast<uint8_t>(n_),
392+ static_cast<uint8_t>(n_ * 8), static_cast<uint8_t>(n_ * 8), static_cast<uint8_t>(n_ * 8)});
393+ PipeBarrier<PIPE_V>();
394+ }
395+ PipeBarrier<PIPE_V>();
396+ // sum(xg * normed )/ (sum + eps) ** 2
397+ Div(gradHResTempLocal2_, gradHResTempLocal2_, skColSumLocal_, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
398+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), 1, static_cast<uint8_t>(n_ * 8),
399+ static_cast<uint8_t>(n_ * 8), 8});
400+ PipeBarrier<PIPE_V>();
401+ 
402+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
403+ Sub(gradHResLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK],
404+ gradHResTempLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK], gradHResTempLocal2_, ELEMENTS_SIZE_PER_REPEAT,
405+ tileRepeatTimes_,
406+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), static_cast<uint8_t>(n_),
407+ static_cast<uint8_t>(n_ * 8), static_cast<uint8_t>(n_ * 8), static_cast<uint8_t>(n_ * 8)});
408+ }
409+ PipeBarrier<PIPE_V>();
410+ 
411+ Adds(skRowSumLocal_, skRowSumLocal_, eps_, tileTaskCount * n_);
412+ PipeBarrier<PIPE_V>();
413+ 
414+ Brcb(gradHResTempLocal2_, skRowSumLocal_, brcbAlign, {1, 8});
415+ 
416+ PipeBarrier<PIPE_V>();
417+ Div(gradHResTempLocal_, gradHResLocal_, gradHResTempLocal2_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
418+ PipeBarrier<PIPE_V>();
419+ 
420+ Mul(skRowNormLocal_, skRowNormLocal_, gradHResTempLocal_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
421+ PipeBarrier<PIPE_V>();
422+ 
423+ AscendC::BlockReduceSum(skRowNormLocal_, skRowNormLocal_,
424+ static_cast<int32_t>(static_cast<int64_t>(tileRepeatTimes_) * n_), MASK_PRE, 1, 1, 8);
425+ 
426+ PipeBarrier<PIPE_V>();
427+ Brcb(gradHResTempLocal2_, skRowNormLocal_, brcbAlign, {1, 8});
428+ SKInQueue.FreeTensor(skRowNormLocal_);
429+ 
430+ PipeBarrier<PIPE_V>();
431+ Sub(gradHResLocal_, gradHResTempLocal_, gradHResTempLocal2_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
432+ PipeBarrier<PIPE_V>();
433+ 
434+ iterRowNormOffset = iterRowNormOffset - 2 * BSNN;
435+ iterRowSumOffset = iterRowSumOffset - 2 * BSN;
436+ iterColSumOffset = iterColSumOffset - 2 * BSN;
437+ }
438+ auto skRowNormLocal_ = SKInQueue.AllocTensor<T>();
439+ auto skColSumLocal_ = skRowNormLocal_[tileCoreBS_ * n_ * ELEMENTS_SIZE_PER_BLOCK];
440+ DataCopyPad(skRowNormLocal_, skNormGlobal_[iterRowNormOffset],
441+ {static_cast<uint16_t>(tileTaskCount * n_), static_cast<uint32_t>(n_ * sizeof(T)), 0, 0, 0},
442+ {true, 0, static_cast<uint8_t>(ELEMENTS_SIZE_PER_BLOCK - n_), 0});
443+ DataCopyPad(skColSumLocal_, skSumGlobal_[iterColSumOffset],
444+ {static_cast<uint16_t>(tileTaskCount), static_cast<uint32_t>(n_ * sizeof(T)), 0, 0, 0},
445+ {true, 0, static_cast<uint8_t>(ELEMENTS_SIZE_PER_BLOCK - n_), 0});
446+ PipeBarrier<PIPE_V>();
447+ SKInQueue.EnQue(skRowNormLocal_);
448+ 
449+ SKInQueue.DeQue();
450+ Adds(skColSumLocal_, skColSumLocal_, eps_, tileTaskCount * ELEMENTS_SIZE_PER_BLOCK);
451+ PipeBarrier<PIPE_V>();
452+ 
453+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
454+ Div(gradHResTempLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK], gradHResLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK],
455+ skColSumLocal_, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
456+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), 1, static_cast<uint8_t>(n_ * 8),
457+ static_cast<uint8_t>(n_ * 8), 8});
458+ }
459+ 
460+ PipeBarrier<PIPE_V>();
461+ // xg * normed
462+ Mul(gradHResTempLocal2_, gradHResLocal_, skRowNormLocal_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
463+ PipeBarrier<PIPE_V>();
464+ // (sum + eps) ** 2
465+ Mul(skColSumLocal_, skColSumLocal_, skColSumLocal_, tileTaskCount * ELEMENTS_SIZE_PER_BLOCK);
466+ // bs * n *8 reduce bs * 1 *8 out : value other other other value other other other
467+ // sum(xg * normed )
468+ for (int32_t loopIdN = 1; loopIdN < n_; loopIdN += 1) {
469+ Add(gradHResTempLocal2_, gradHResTempLocal2_, gradHResTempLocal2_[loopIdN * ELEMENTS_SIZE_PER_BLOCK],
470+ ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
471+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), static_cast<uint8_t>(n_ * 8),
472+ static_cast<uint8_t>(n_ * 8), static_cast<uint8_t>(n_ * 8)});
473+ PipeBarrier<PIPE_V>();
474+ }
475+ 
476+ PipeBarrier<PIPE_V>();
477+ // sum(xg * normed )/ (sum + eps) ** 2
478+ Div(gradHResTempLocal2_, gradHResTempLocal2_, skColSumLocal_, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
479+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), 1, static_cast<uint8_t>(n_ * 8),
480+ static_cast<uint8_t>(n_ * 8), 8});
481+ PipeBarrier<PIPE_V>();
482+ 
483+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
484+ Sub(gradHResLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK], gradHResTempLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK],
485+ gradHResTempLocal2_, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
486+ {static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), static_cast<uint8_t>(n_), static_cast<uint8_t>(n_ * 8),
487+ static_cast<uint8_t>(n_ * 8), static_cast<uint8_t>(n_ * 8)});
488+ }
489+ // SOFTMAX
490+ PipeBarrier<PIPE_V>();
491+ Mul(gradHResTempLocal_, gradHResLocal_, skRowNormLocal_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
492+ SKInQueue.FreeTensor(skRowNormLocal_);
493+ 
494+ PipeBarrier<PIPE_V>();
495+ AscendC::BlockReduceSum(gradHResTempLocal_, gradHResTempLocal_,
496+ static_cast<int32_t>(static_cast<int64_t>(tileRepeatTimes_) * n_), MASK_PRE, 1, 1, 8);
497+ PipeBarrier<PIPE_V>();
498+ Brcb(gradHResTempLocal2_, gradHResTempLocal_, brcbAlign, {1, 8});
499+ PipeBarrier<PIPE_V>();
500+ Sub(gradHResLocal_, gradHResLocal_, gradHResTempLocal2_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
501+ PipeBarrier<PIPE_V>();
502+ Mul(gradHResLocal_, gradHResLocal_, skRowNormLocal_, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
503+ PipeBarrier<PIPE_V>();
504+ // 4v 4o 4v 40 -> 4v 4v
505+ Cast(gradHResTempLocal_.template ReinterpretCast<int64_t>(), gradHResLocal_.template ReinterpretCast<int32_t>(),
506+ RoundMode::CAST_NONE, tileTaskCount * n_ * ELEMENTS_SIZE_PER_BLOCK);
507+ PipeBarrier<PIPE_V>();
508+ 
509+ Copy(gradHResLocal_, gradHResTempLocal_, ELEMENTS_SIZE_PER_REPEAT, n_ * tileRepeatTimes_, {1, 2, 8, 16});
510+ PipeBarrier<PIPE_V>();
511+ 
512+ Cast(gradHResLocal_.template ReinterpretCast<int32_t>(), gradHResLocal_.template ReinterpretCast<int64_t>(),
513+ RoundMode::CAST_NONE, tileTaskCount * n_ * n_);
514+ 
515+ PipeBarrier<PIPE_V>();
516+}
517+ 
518+template <typename TYPE_X, typename T>
519+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::ComputeGradHHat2(const int32_t taskOffset,
520+ const int32_t tileTaskCount)
521+{
522+ SinkhornGrad(taskOffset, tileTaskCount);
523+ for (int32_t loopIdN = 0; loopIdN < 2; loopIdN += 1) {
524+ Copy(dhatLocal_[ELEMENTS_SIZE_PER_BLOCK + loopIdN * ELEMENTS_SIZE_PER_BLOCK],
525+ gradHResLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
526+ {static_cast<uint16_t>(3), static_cast<uint16_t>(2), static_cast<uint16_t>((2 + 1) * 8),
527+ static_cast<uint16_t>(2 * 8)});
528+ }
529+ 
530+ inputGradQueue.FreeTensor(gradHResLocal_);
531+ 
532+ auto gradHPostLocal_ = inputXInQueue.AllocTensor<T>();
533+ auto hat2LocalTemp = gradHPostLocal_[tileCoreBS_ * ELEMENTS_SIZE_PER_BLOCK];
534+ auto rsqrtLocal_ =
535+ gradHPostLocal_[tileCoreBS_ * ELEMENTS_SIZE_PER_BLOCK + tileCoreBS_ * (n_ * n_ + PRE_POST_NUM * n_)];
536+ 
537+ DataCopyPad(gradHPostLocal_, gradPostGlobal_[taskOffset * n_],
538+ {static_cast<uint16_t>(tileTaskCount), static_cast<uint32_t>(n_ * sizeof(T)), 0, 0, 0},
539+ {true, static_cast<uint8_t>(ELEMENTS_SIZE_PER_BLOCK - n_), 0, 0});
540+ 
541+ DataCopyPad(hat2LocalTemp, h2Global_[taskOffset * (n_ * n_ + PRE_POST_NUM * n_)],
542+ {static_cast<uint16_t>(tileTaskCount), static_cast<uint32_t>((n_ * n_ + PRE_POST_NUM * n_) * sizeof(T)),
543+ 0, 0, 0},
544+ {false, 0, 0, 0});
545+ DataCopyPad(rsqrtLocal_, rsqrtGlobal_[taskOffset],
546+ {static_cast<uint16_t>(1), static_cast<uint32_t>(tileTaskCount * sizeof(T)), 0, 0, 0},
547+ {false, 0, 0, 0});
548+ inputXInQueue.EnQue(gradHPostLocal_);
549+ inputXInQueue.DeQue();
550+ 
551+ PipeBarrier<PIPE_V>();
552+ 
553+ // GRAD_PRE*1_POST*2
554+ Axpy(dPrePostTempLocal_, gradHPostLocal_, float(2), tileTaskCount * 2 * n_);
555+ 
556+ const uint32_t srcShape[2] = {static_cast<uint32_t>(tileTaskCount * n_), 1}; // 源数据shape
557+ const uint32_t dstShape[2] = {static_cast<uint32_t>(tileTaskCount * n_),
558+ static_cast<uint32_t>(n_ * n_ + PRE_POST_NUM * n_)}; // broadcast数据shape
559+ PipeBarrier<PIPE_V>();
560+ 
561+ AscendC::Broadcast<float, 2, 1>(rsqrtbrcbLocal_, rsqrtLocal_, dstShape, srcShape);
562+ 
563+ PipeBarrier<PIPE_V>();
564+ 
565+ // GRAD_HAT_RSQRT
566+ Mul(hat2Scale, scaleLocal_, hat2LocalTemp, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {3, 0, 3, 24, 0, 24});
567+ Mul(hat2Scale[8], scaleLocal_[8], hat2LocalTemp[8], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
568+ {3, 0, 3, 24, 0, 24});
569+ Mul(hat2Scale[16], scaleLocal_[8], hat2LocalTemp[16], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
570+ {3, 0, 3, 24, 0, 24});
571+ PipeBarrier<PIPE_V>();
572+ 
573+ Add(hatLocal, hcBaseLocal_, hat2Scale, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {3, 0, 3, 24, 0, 24});
574+ Add(hatLocal[8], hcBaseLocal_[8], hat2Scale[8], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {3, 0, 3, 24, 0, 24});
575+ Add(hatLocal[16], hcBaseLocal_[16], hat2Scale[16], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
576+ {3, 0, 3, 24, 0, 24});
577+ PipeBarrier<PIPE_V>();
578+ 
579+ Mul(hatLocal, hatLocal, rsqrtbrcbLocal_, (n_ * n_ + PRE_POST_NUM * n_) * tileTaskCount);
580+ PipeBarrier<PIPE_V>();
581+ 
582+ auto hatLocalTemp = dhatBeforeNormLocal;
583+ 
584+ // 正向pre post sigmoid
585+ Muls(hatLocalTemp, hatLocal, float(-1), ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {1, 3, 8, 24});
586+ PipeBarrier<PIPE_V>();
587+ 
588+ Exp(hatLocal, hatLocalTemp, (2 * n_) * tileTaskCount);
589+ PipeBarrier<PIPE_V>();
590+ 
591+ Adds(hatLocal, hatLocal, float(1), (2 * n_) * tileTaskCount);
592+ PipeBarrier<PIPE_V>();
593+ Div(hatLocal, onesLocal_, hatLocal, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {1, 0, 1, 8, 0, 8});
594+ 
595+ // 正向pre post sigmoidgrad x_sigmoid * (1 - x_sigmoid) * grad_output
596+ PipeBarrier<PIPE_V>();
597+ Mul(hatLocalTemp, hatLocal, hatLocal, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {1, 1, 1, 8, 8, 8});
598+ PipeBarrier<PIPE_V>();
599+ Sub(hatLocalTemp, hatLocal, hatLocalTemp, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {1, 1, 1, 8, 8, 8});
600+ PipeBarrier<PIPE_V>();
601+ 
602+ Mul(dhatLocal_, dPrePostTempLocal_, hatLocalTemp, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {3, 1, 1, 24, 8, 8});
603+ PipeBarrier<PIPE_V>();
604+ 
605+ // 不同批次dbias累加
606+ PipeBarrier<PIPE_V>();
607+ Add(dBiasLocal_, dhatLocal_, dBiasLocal_, (n_ * n_ + PRE_POST_NUM * n_) * tileTaskCount);
608+ // dRsqrt = dhatLocal_ * beforenorm
609+ Mul(gradRsqrtLocal_, dhatLocal_, hat2Scale, (n_ * n_ + PRE_POST_NUM * n_) * tileTaskCount);
610+ PipeBarrier<PIPE_V>();
611+ // rmsnormgrad (- (rsqrt ** 3)) * grad_rsqrt / float(NC)
612+ WholeReduceSum(rsqrtTempLocal_, gradRsqrtLocal_, (n_ * n_ + PRE_POST_NUM * n_), tileTaskCount, 1, 1,
613+ 3); // dst repeat stride srcblkstride srcrepstride (n*2 +2n)/8
614+ PipeBarrier<PIPE_V>();
615+ Mul(rsqrtTempLocal_, rsqrtTempLocal_, rsqrtLocal_, tileTaskCount);
616+ PipeBarrier<PIPE_V>();
617+ 
618+ Mul(rsqrtLocal_, rsqrtLocal_, rsqrtLocal_, tileTaskCount);
619+ PipeBarrier<PIPE_V>();
620+ Mul(rsqrtTempLocal_, rsqrtTempLocal_, rsqrtLocal_, tileTaskCount);
621+ inputXInQueue.FreeTensor(gradHPostLocal_);
622+ 
623+ PipeBarrier<PIPE_V>();
624+ 
625+ PipeBarrier<PIPE_V>();
626+ Muls(gradRsqrtLocal_, rsqrtTempLocal_, float(-1) / (n_ * c_), tileTaskCount);
627+ 
628+ // dscale = dhatLocal_ * hat pre * rsqrt
629+ // dhat2 = dhatLocal_ * rsqrt * scale
630+ Mul(dhatBeforeNormLocal, rsqrtbrcbLocal_, dhatLocal_, (n_ * n_ + PRE_POST_NUM * n_) * tileTaskCount);
631+ PipeBarrier<PIPE_V>();
632+ // 不同批次dScale累加
633+ MulAddDst(dScaleLocal_, hat2LocalTemp, dhatBeforeNormLocal, (n_ * n_ + PRE_POST_NUM * n_) * tileTaskCount);
634+ auto dhat2Local = OutQueue.AllocTensor<float>();
635+ 
636+ Mul(dhat2Local, scaleLocal_, dhatBeforeNormLocal, ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_, {3, 0, 3, 24, 0, 24});
637+ Mul(dhat2Local[8], scaleLocal_[8], dhatBeforeNormLocal[8], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
638+ {3, 0, 3, 24, 0, 24});
639+ Mul(dhat2Local[16], scaleLocal_[8], dhatBeforeNormLocal[16], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
640+ {3, 0, 3, 24, 0, 24});
641+ OutQueue.EnQue(dhat2Local);
642+ dhat2Local = OutQueue.DeQue<T>();
643+ DataCopyPad(gradH2Global_[taskOffset * (n_ * n_ + PRE_POST_NUM * n_)], dhat2Local,
644+ {static_cast<uint16_t>(1),
645+ static_cast<uint32_t>(tileTaskCount * (n_ * n_ + PRE_POST_NUM * n_) * sizeof(float)), 0, 0, 0});
646+ OutQueue.FreeTensor(dhat2Local);
647+}
648+ 
649+template <typename TYPE_X, typename T>
650+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::ComputeGradX1(const int32_t taskOffset, const int32_t tileTaskCount,
651+ const int32_t innerId)
652+{
653+ PipeBarrier<PIPE_V>();
654+ 
655+ for (int32_t loopIdC = 0; loopIdC < c1Align_; loopIdC += 1) {
656+ int64_t copyLen = c0_;
657+ bool isPad = false;
658+ uint8_t padLen = 0;
659+ int64_t ubAlignC = c0_;
660+ if (loopIdC == c1_) {
661+ isPad = true;
662+ copyLen = cTail_;
663+ ubAlignC = cTailAlign_;
664+ padLen = static_cast<uint8_t>(cTailAlign_ - cTail_);
665+ }
666+ auto xLocal_ = inputXInQueue.AllocTensor<TYPE_X>();
667+ auto gradXCubeLocal_ = xLocal_.template ReinterpretCast<float>()[onceTask_ * n_ * c0_ / 2];
668+ auto gradYLocal_ = xLocal_[onceTask_ * n_ * c0_ * 3];
669+ 
670+ DataCopyPad(gradYLocal_, gradYGlobal_[taskOffset * c_ + c0_ * loopIdC],
671+ {static_cast<uint16_t>(tileTaskCount), static_cast<uint32_t>(copyLen * sizeof(TYPE_X)),
672+ static_cast<uint32_t>((c_ - copyLen) * sizeof(TYPE_X)), 0, 0},
673+ {isPad, 0, padLen, 0});
674+ DataCopyPad(xLocal_, xGlobal_[taskOffset * n_ * c_ + c0_ * loopIdC],
675+ {static_cast<uint16_t>(n_ * tileTaskCount), static_cast<uint32_t>(copyLen * sizeof(TYPE_X)),
676+ static_cast<uint32_t>((c_ - copyLen) * sizeof(TYPE_X)), 0, 0},
677+ {isPad, 0, padLen, 0});
678+ DataCopyPad(gradXCubeLocal_, gradXCubeGlobal_[taskOffset * n_ * c_ + c0_ * loopIdC],
679+ {static_cast<uint16_t>(n_ * tileTaskCount), static_cast<uint32_t>(copyLen * sizeof(T)),
680+ static_cast<uint32_t>((c_ - copyLen) * sizeof(T)), 0, 0},
681+ {isPad, 0, padLen, 0});
682+ 
683+ inputXInQueue.EnQue(xLocal_);
684+ inputXInQueue.DeQue();
685+ PipeBarrier<PIPE_V>();
686+ Cast(gradYCastLocal_, gradYLocal_, RoundMode::CAST_NONE, ubAlignC * tileTaskCount);
687+ 
688+ Cast(xCastLocal_, xLocal_, RoundMode::CAST_NONE, ubAlignC * n_ * tileTaskCount);
689+ 
690+ PipeBarrier<PIPE_V>();
691+ uint8_t blkStride1 = static_cast<uint8_t>(ubAlignC / ELEMENTS_SIZE_PER_BLOCK);
692+ uint8_t blkStride2 = static_cast<uint8_t>(n_ * ubAlignC / ELEMENTS_SIZE_PER_BLOCK);
693+ 
694+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
695+ for (int32_t loopOffsetC0 = 0; loopOffsetC0 < copyLen; loopOffsetC0 += ELEMENTS_SIZE_PER_REPEAT) {
696+ uint64_t mask =
697+ min(static_cast<uint64_t>(ELEMENTS_SIZE_PER_REPEAT), static_cast<uint64_t>(copyLen - loopOffsetC0));
698+ Mul(gradXCastLocal_[loopIdN * ubAlignC + loopOffsetC0], gradYCastLocal_[loopOffsetC0],
699+ preBrcbLocal_[loopIdN * ELEMENTS_SIZE_PER_BLOCK + innerId * n_ * ELEMENTS_SIZE_PER_BLOCK], mask,
700+ tileTaskCount, {1, 1, 0, blkStride2, blkStride1, static_cast<uint8_t>(n_)});
701+ }
702+ }
703+ 
704+ PipeBarrier<PIPE_V>();
705+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
706+ for (int32_t loopOffsetC0 = 0; loopOffsetC0 < copyLen; loopOffsetC0 += ELEMENTS_SIZE_PER_REPEAT) {
707+ uint64_t mask =
708+ min(static_cast<uint64_t>(ELEMENTS_SIZE_PER_REPEAT), static_cast<uint64_t>(copyLen - loopOffsetC0));
709+ MulAddDst(gradXCastLocal_[loopIdN * ubAlignC + loopOffsetC0],
710+ xCastLocal_[loopIdN * ubAlignC + loopOffsetC0],
711+ dRsqrtBrcbLocal_[innerId * ELEMENTS_SIZE_PER_BLOCK], mask, tileTaskCount,
712+ {1, 1, 0, blkStride2, blkStride2, 1});
713+ }
714+ }
715+ 
716+ Add(gradXCastLocal_, gradXCubeLocal_, gradXCastLocal_, ubAlignC * n_ * tileTaskCount);
717+ PipeBarrier<PIPE_V>();
718+ inputXInQueue.FreeTensor(xLocal_);
719+ 
720+ auto gradXLocalOut = OutQueue.AllocTensor<TYPE_X>();
721+ 
722+ Cast(gradXLocalOut, gradXCastLocal_, RoundMode::CAST_RINT, ubAlignC * n_ * tileTaskCount);
723+ OutQueue.EnQue(gradXLocalOut);
724+ gradXLocalOut = OutQueue.DeQue<TYPE_X>();
725+ 
726+ DataCopyPad(gradXGlobal_[taskOffset * n_ * c_ + c0_ * loopIdC], gradXLocalOut,
727+ {static_cast<uint16_t>(n_ * tileTaskCount), static_cast<uint32_t>(copyLen * sizeof(TYPE_X)), 0,
728+ static_cast<uint32_t>((c_ - copyLen) * sizeof(TYPE_X)), 0});
729+ OutQueue.FreeTensor(gradXLocalOut);
730+ 
731+ PipeBarrier<PIPE_V>();
732+ }
733+}
734+template <typename TYPE_X, typename T>
735+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::ComputeGradPre(const int32_t taskOffset,
736+ const int32_t tileTaskCount, const int32_t innerId)
737+{
738+ Duplicate(tmpLocal_, 0.f, tileTaskCount * n_ * ELEMENTS_SIZE_PER_REPEAT);
739+ for (int32_t loopIdC = 0; loopIdC < c1Align_; loopIdC += 1) {
740+ int64_t copyLen = c0_;
741+ bool isPad = false;
742+ uint8_t padLen = 0;
743+ int64_t ubAlignC = c0_;
744+ if (loopIdC == c1_) {
745+ isPad = false;
746+ copyLen = cTail_;
747+ ubAlignC = cTailAlign_;
748+ padLen = static_cast<uint8_t>(cTailAlign_ - cTail_);
749+ }
750+ auto xLocal_ = inputXInQueue.AllocTensor<TYPE_X>();
751+ 
752+ auto gradYLocal_ = xLocal_[onceTask_ * n_ * c0_];
753+ 
754+ DataCopyPad(gradYLocal_, gradYGlobal_[taskOffset * c_ + c0_ * loopIdC],
755+ {static_cast<uint16_t>(tileTaskCount), static_cast<uint32_t>(copyLen * sizeof(TYPE_X)),
756+ static_cast<uint32_t>((c_ - copyLen) * sizeof(TYPE_X)), 0, 0},
757+ {isPad, 0, padLen, 0});
758+ 
759+ DataCopyPad(xLocal_, xGlobal_[taskOffset * n_ * c_ + c0_ * loopIdC],
760+ {static_cast<uint16_t>(n_ * tileTaskCount), static_cast<uint32_t>(copyLen * sizeof(TYPE_X)),
761+ static_cast<uint32_t>((c_ - copyLen) * sizeof(TYPE_X)), 0, 0},
762+ {isPad, 0, padLen, 0});
763+ 
764+ inputXInQueue.EnQue(xLocal_);
765+ inputXInQueue.DeQue();
766+ PipeBarrier<PIPE_V>();
767+ auto xCastOutLocal = OutQueue.AllocTensor<float>();
768+ 
769+ Cast(xCastOutLocal, xLocal_, RoundMode::CAST_NONE, ubAlignC * n_ * tileTaskCount);
770+ 
771+ Cast(gradYCastLocal_, gradYLocal_, RoundMode::CAST_NONE, ubAlignC * tileTaskCount);
772+ inputXInQueue.FreeTensor(xLocal_);
773+ 
774+ PipeBarrier<PIPE_V>();
775+ uint8_t blkStride = static_cast<uint8_t>(ubAlignC / ELEMENTS_SIZE_PER_BLOCK);
776+ 
777+ uint8_t blkStride3 = static_cast<uint8_t>(n_ * ubAlignC / ELEMENTS_SIZE_PER_BLOCK);
778+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
779+ for (int32_t loopOffsetC0 = 0; loopOffsetC0 < copyLen; loopOffsetC0 += ELEMENTS_SIZE_PER_REPEAT) {
780+ uint64_t mask =
781+ min(static_cast<uint64_t>(ELEMENTS_SIZE_PER_REPEAT), static_cast<uint64_t>(copyLen - loopOffsetC0));
782+ Mul(xCastLocal_[loopIdN * ubAlignC + loopOffsetC0], xCastOutLocal[loopIdN * ubAlignC + loopOffsetC0],
783+ gradYCastLocal_[loopOffsetC0], mask, tileTaskCount, {1, 1, 1, blkStride3, blkStride3, blkStride});
784+ }
785+ }
786+ OutQueue.EnQue(xCastOutLocal);
787+ xCastOutLocal = OutQueue.DeQue<T>();
788+ 
789+ DataCopyPad(xWorkspaceGlobal_[taskOffset * n_ * c_ + c0_ * loopIdC], xCastOutLocal,
790+ {static_cast<uint16_t>(n_ * tileTaskCount), static_cast<uint32_t>(copyLen * sizeof(float)), 0,
791+ static_cast<uint32_t>((c_ - copyLen) * sizeof(float)), 0});
792+ OutQueue.FreeTensor(xCastOutLocal);
793+ 
794+ PipeBarrier<PIPE_V>();
795+ 
796+ int64_t reduceLen = ubAlignC;
797+ if (ubAlignC != c0_) {
798+ Add(xCastLocal_[64], xCastLocal_[64], xCastLocal_[128 + 64], ELEMENTS_SIZE_PER_REPEAT, tileTaskCount * n_,
799+ {1, 1, 1, blkStride, blkStride, blkStride});
800+ Add(xCastLocal_, xCastLocal_, xCastLocal_[128], ELEMENTS_SIZE_PER_REPEAT, tileTaskCount * n_,
801+ {1, 1, 1, blkStride, blkStride, blkStride});
802+ PipeBarrier<PIPE_V>();
803+ 
804+ Add(xCastLocal_, xCastLocal_, xCastLocal_[64], ELEMENTS_SIZE_PER_REPEAT, tileTaskCount * n_,
805+ {1, 1, 1, blkStride, blkStride, blkStride});
806+ } else {
807+ if (cTail_ - (128 + 64) > 0) {
808+ uint64_t mask = min(static_cast<uint64_t>(cTail_ - (128 + 64)), static_cast<uint64_t>(REPEAT_LENTH));
809+ Add(xCastLocal_[64], xCastLocal_[64], xCastLocal_[128 + 64], mask, tileTaskCount * n_,
810+ {1, 1, 1, blkStride, blkStride, blkStride});
811+ }
812+ if (cTail_ - (128) > 0) {
813+ uint64_t mask = min(static_cast<uint64_t>(cTail_ - (128)), static_cast<uint64_t>(REPEAT_LENTH));
814+ Add(xCastLocal_, xCastLocal_, xCastLocal_[128], mask, tileTaskCount * n_,
815+ {1, 1, 1, blkStride, blkStride, blkStride});
816+ }
817+ PipeBarrier<PIPE_V>();
818+ if (cTail_ - (64) > 0) {
819+ uint64_t mask = min(static_cast<uint64_t>(cTail_ - (64)), static_cast<uint64_t>(REPEAT_LENTH));
820+ Add(xCastLocal_, xCastLocal_, xCastLocal_[64], mask, tileTaskCount * n_,
821+ {1, 1, 1, blkStride, blkStride, blkStride});
822+ }
823+ }
824+ PipeBarrier<PIPE_V>();
825+ 
826+ uint64_t mask = min(static_cast<uint64_t>(cTail_), static_cast<uint64_t>(REPEAT_LENTH));
827+ PipeBarrier<PIPE_V>();
828+ 
829+ Add(tmpLocal_, tmpLocal_, xCastLocal_, mask, tileTaskCount * n_, {1, 1, 1, 8, 8, blkStride});
830+ PipeBarrier<PIPE_V>();
831+ }
832+ PipeBarrier<PIPE_V>();
833+ PipeBarrier<PIPE_V>();
834+ 
835+ for (int32_t loopIdN = 0; loopIdN < n_; loopIdN += 1) {
836+ WholeReduceSum(dPrePostTempLocal_[loopIdN + innerId * n_ * 2], tmpLocal_[loopIdN * REPEAT_LENTH], REPEAT_LENTH,
837+ tileTaskCount, n_ * 2, 1, n_ * 8); // dst repeat stride srcblkstride srcrepstride
838+ // dst 间隔 8 2 *n
839+ }
840+}
841+ 
842+template <typename TYPE_X, typename T>
843+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::ComputeScaleBias()
844+{
845+ // 调换顺序
846+ Add(dScaleLocal_[8], dScaleLocal_[8], dScaleLocal_[16], ELEMENTS_SIZE_PER_REPEAT, tileRepeatTimes_,
847+ {3, 3, 3, 24, 24, 24});
848+ // 和tileTaskCount相关
849+ 
850+ PipeBarrier<PIPE_V>();
851+ // tileCoreBS_ 必须8 *2的幂次
852+ for (int32_t bsCount = tileCoreBS_ / 2; bsCount > 0; bsCount = bsCount / 2) {
853+ Add(dBiasLocal_, dBiasLocal_, dBiasLocal_[bsCount * (n_ * n_ + PRE_POST_NUM * n_)],
854+ bsCount * (n_ * n_ + PRE_POST_NUM * n_));
855+ Add(dScaleLocal_, dScaleLocal_, dScaleLocal_[bsCount * (n_ * n_ + PRE_POST_NUM * n_)],
856+ bsCount * (n_ * n_ + PRE_POST_NUM * n_));
857+ PipeBarrier<PIPE_V>();
858+ }
859+ SetFlag<HardEvent::V_MTE3>(eventIdVToMTE3XCast);
860+ WaitFlag<HardEvent::V_MTE3>(eventIdVToMTE3XCast);
861+ 
862+ SetAtomicAdd<T>();
863+ 
864+ LocalTensor<T> dscaleOut = tempBuf_.GetWithOffset<T>(ELEMENTS_SIZE_PER_BLOCK, 0);
865+ 
866+ DataCopyPad(gradHcBaseGlobal_, dBiasLocal_,
867+ {static_cast<uint16_t>(1), static_cast<uint32_t>((n_ * n_ + PRE_POST_NUM * n_) * sizeof(T)), 0, 0, 0});
868+ PipeBarrier<PIPE_V>();
869+ 
870+ WholeReduceSum(dscaleOut, dScaleLocal_, 4, 1, 1, 3, 8); // dst repeat stride srcblkstride srcrepstride
871+ WholeReduceSum(dscaleOut[1], dScaleLocal_, MASK_POST_SCALE, 1, 1, 3,
872+ 8); // dst repeat stride srcblkstride srcrepstride
873+ WholeReduceSum(dscaleOut[2], dScaleLocal_[8], 8, 1, 1, 3, 8); // dst repeat stride srcblkstride srcrepstride
874+ WholeReduceSum(dscaleOut[3], dScaleLocal_, 8, 1, 1, 3, 8); // dst repeat stride srcblkstride srcrepstride
875+ 
876+ SetFlag<HardEvent::V_MTE3>(eventIdVToMTE3XCast);
877+ WaitFlag<HardEvent::V_MTE3>(eventIdVToMTE3XCast);
878+ DataCopyPad(gradHcScaleGlobal_, dscaleOut,
879+ {static_cast<uint16_t>(1), static_cast<uint32_t>((3) * sizeof(T)), 0, 0, 0});
880+ SetAtomicNone();
881+}
882+/**
883+ 约束:
884+ c: c >= 64 && c % 64 == 0
885+ n: n == 4
886+*/
887+ 
888+template <typename TYPE_X, typename T>
889+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::Process()
890+{
891+ if ASCEND_IS_AIV {
892+ GetHcScaleAndHcBase();
893+ 
894+ int8_t ping = 0;
895+ 
896+ for (int32_t taskOffset = blkIdx_ * tileCoreBS_; taskOffset < totalTasksAligned_;
897+ taskOffset += aivNum_ * tileCoreBS_) {
898+ int32_t tileTaskCount =
899+ min(static_cast<int32_t>(tileCoreBS_), static_cast<int32_t>(totalTasks_ - taskOffset));
900+ tileRepeatTimes_ = CeilDiv(tileTaskCount * 2 * n_, ELEMENTS_SIZE_PER_REPEAT);
901+ if (tileTaskCount > 0) {
902+ int32_t innerId = 0;
903+ Duplicate(dPrePostTempLocal_, 0.f, tileCoreBS_ * n_ * 2);
904+ for (int32_t taskOffsetInner = 0; taskOffsetInner < tileTaskCount; taskOffsetInner += onceTask_) {
905+ int32_t tileTaskCountInner =
906+ min(static_cast<int32_t>(onceTask_), static_cast<int32_t>(tileTaskCount - taskOffsetInner));
907+ 
908+ ComputeGradPre(taskOffset + taskOffsetInner, tileTaskCountInner, taskOffsetInner);
909+ innerId++;
910+ }
911+ ComputeGradHHat2(taskOffset, tileTaskCount);
912+ }
913+ CrossCoreSetFlag<0x2, PIPE_MTE3>(0);
914+ CrossCoreWaitFlag<0x2>(1);
915+ ping = (ping + 1) % 10;
916+ if (tileTaskCount > 0) {
917+ int32_t innerId = 0;
918+ for (int32_t taskOffsetInner = 0; taskOffsetInner < tileTaskCount; taskOffsetInner += onceTask_) {
919+ int32_t tileTaskCountInner =
920+ min(static_cast<int32_t>(onceTask_), static_cast<int32_t>(tileTaskCount - taskOffsetInner));
921+ int32_t brcbAlign = CeilDiv(tileTaskCount * n_, ELEMENTS_SIZE_PER_BLOCK);
922+ int32_t offset = 0;
923+ auto preLocal_ = inputGradQueue.AllocTensor<T>();
924+ DataCopyPad(
925+ preLocal_, preGlobal_[taskOffset * n_],
926+ {static_cast<uint16_t>(1), static_cast<uint32_t>(tileTaskCount * n_ * sizeof(T)), 0, 0, 0},
927+ {false, 0, 0, 0});
928+ inputGradQueue.EnQue(preLocal_);
929+ inputGradQueue.DeQue();
930+ 
931+ const uint32_t srcShape[2] = {static_cast<uint32_t>(tileTaskCount * n_), 1}; // 源数据shape
932+ const uint32_t dstShape[2] = {static_cast<uint32_t>(tileTaskCount * n_),
933+ ELEMENTS_SIZE_PER_BLOCK}; // broadcast数据shape
934+ preBrcbLocal_ = tempBuf_.GetWithOffset<T>(brcbAlign * 8, offset);
935+ Brcb(preBrcbLocal_, preLocal_, brcbAlign, {static_cast<uint8_t>(1), static_cast<uint8_t>(8)});
936+ inputGradQueue.FreeTensor(preLocal_);
937+ 
938+ offset += brcbAlign * 8 * sizeof(float);
939+ brcbAlign = CeilDiv(tileTaskCount, ELEMENTS_SIZE_PER_BLOCK);
940+ 
941+ offset += brcbAlign * 8 * sizeof(float);
942+ 
943+ const uint32_t srcRsqrtShape[2] = {static_cast<uint32_t>(tileTaskCount), 1}; // 源数据shape
944+ const uint32_t dstRsqrtShape[2] = {static_cast<uint32_t>(tileTaskCount),
945+ ELEMENTS_SIZE_PER_BLOCK}; // broadcast数据shape
946+ Brcb(dRsqrtBrcbLocal_, gradRsqrtLocal_, brcbAlign,
947+ {static_cast<uint8_t>(1), static_cast<uint8_t>(8)});
948+ ComputeGradX1(taskOffset + taskOffsetInner, tileTaskCountInner, taskOffsetInner);
949+ innerId++;
950+ }
951+ }
952+ }
953+ 
954+ tileRepeatTimes_ = CeilDiv(tileCoreBS_ * 2 * n_, ELEMENTS_SIZE_PER_REPEAT);
955+ ComputeScaleBias();
956+ GetTPipePtr()->ReleaseEventID<HardEvent::V_MTE3>(eventIdVToMTE3XCast);
957+ }
958+ 
959+ if ASCEND_IS_AIC {
960+ int8_t ping = 0;
961+ for (int32_t taskOffset = blkIdx_ * 2 * tileCoreBS_; taskOffset < totalTasksAligned_;
962+ taskOffset += aicNum_ * 2 * tileCoreBS_) {
963+ int32_t tileTaskCount =
964+ min(static_cast<int32_t>(2 * tileCoreBS_), static_cast<int32_t>(totalTasks_ - taskOffset));
965+ CrossCoreWaitFlag<0x2>(0);
966+ 
967+ if (tileTaskCount > 0) {
968+ ProcessMatmul1(taskOffset, tileTaskCount);
969+ ProcessMatmul2(taskOffset, tileTaskCount);
970+ }
971+ AscendC::CrossCoreSetFlag<0x2, PIPE_FIX>(1);
972+ 
973+ ping = (ping + 1) % 10;
974+ }
975+ }
976+}
977+ 
978+template <typename TYPE_X, typename T>
979+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::ProcessMatmul1(const int32_t taskOffset, const int32_t mm1M)
980+{
981+ if (mm1M <= 0)
982+ return;
983+ 
984+ mm1_.SetTensorA(gradH2Global_[taskOffset * mm1K_]);
985+ mm1_.SetTensorB(weightGlobal_);
986+ mm1_.SetHF32(true, 1);
987+ mm1_.SetOrgShape(mm1M, mm1N_, mm1K_);
988+ mm1_.SetSingleShape(mm1M, mm1N_, mm1K_);
989+ mm1_.template IterateAll<false>(gradXCubeGlobal_[taskOffset * (n_ * c_)]);
990+ mm1_.End();
991+}
992+ 
993+template <typename TYPE_X, typename T>
994+__aicore__ inline void MhcPreGradKernel<TYPE_X, T>::ProcessMatmul2(const int32_t taskOffset, const int32_t mm2K)
995+{
996+ if (mm2K <= 0)
997+ return;
998+ 
999+ mm2_.SetTensorA(gradH2Global_[taskOffset * mm2M_], true);
1000+ mm2_.SetTensorB(xWorkspaceGlobal_[taskOffset * mm2N_]);
1001+ mm2_.SetHF32(true, 1);
1002+ mm2_.SetOrgShape(mm2M_, mm2N_, mm2K);
1003+ mm2_.SetSingleShape(mm2M_, mm2N_, mm2K);
1004+ mm2_.template IterateAll<false>(gradWeightGlobal_, 1);
1005+ mm2_.End();
1006+}
1007+ 
1008+#endif // MHC_PRE_SINKHORN_BACKWARD_OP_KERNEL_MHC_PRE_GRAD_KERNEL_H
@@ -0,0 +1,46 @@
1+/**
2+ * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file mhc_pre_sinkhorn_backward.cpp
13+ * \brief
14+ */
15+#include "kernel_operator.h"
16+#include "lib/matmul_intf.h"
17+#include "mhc_pre_grad_kernel.h"
18+ 
19+using namespace AscendC;
20+ 
21+extern "C" __global__ __aicore__ void
22+mhc_pre_sinkhorn_backward(GM_ADDR grad_hin, GM_ADDR grad_h_post, GM_ADDR grad_h_res, GM_ADDR x, GM_ADDR phi,
23+ GM_ADDR alpha, GM_ADDR bias, GM_ADDR h_pre, GM_ADDR hc_before_norm, GM_ADDR inv_rms,
24+ GM_ADDR sum_out, GM_ADDR norm_out, GM_ADDR grad_x, GM_ADDR grad_phi, GM_ADDR grad_alpha,
25+ GM_ADDR grad_bias, GM_ADDR workspace, GM_ADDR tiling)
26+{
27+ GET_TILING_DATA(tilingData, tiling);
28+ GM_ADDR usrWorkspace = GetUserWorkspace(workspace);
29+ if (usrWorkspace == nullptr) {
30+ return;
31+ }
32+ KERNEL_TASK_TYPE(0, KERNEL_TYPE_MIX_AIC_1_2);
33+ 
34+ TPipe pipe;
35+ 
36+ MhcPreGradKernel<DTYPE_X, DTYPE_PHI> op;
37+ 
38+ op.mm1_.SetSubBlockIdx(0);
39+ op.mm1_.Init(&tilingData.mm1TilingData, &pipe);
40+ op.mm2_.SetSubBlockIdx(0);
41+ op.mm2_.Init(&tilingData.mm2TilingData, &pipe);
42+ 
43+ op.Init(x, phi, h_pre, grad_hin, grad_h_post, grad_h_res, alpha, bias, hc_before_norm, inv_rms, sum_out, norm_out,
44+ grad_x, grad_phi, grad_alpha, grad_bias, usrWorkspace, &tilingData, &pipe);
45+ op.Process();
46+}