已合并
fix: AcosGrad、AsinGrad 共性问题及 tiling 函数整改 #5268
fix: AcosGrad、AsinGrad 共性问题及 tiling 函数整改 #5268
已合并
wangweidong创建于 8月31日
共 7 个文件变更+216-257
@@ -76,14 +76,10 @@ static ge::graphStatus GetShapeInfo(gert::TilingContext* context, uint64_t& tota
76 OP_CHECK_NULL_WITH_CONTEXT(context, outputZ);76 OP_CHECK_NULL_WITH_CONTEXT(context, outputZ);
77 auto zShape = EnsureNotScalar(outputZ->GetStorageShape());77 auto zShape = EnsureNotScalar(outputZ->GetStorageShape());
78 78 
79- OP_CHECK_IF(79+ OP_CHECK_IF(yShape.GetShapeSize() != dyShape.GetShapeSize() || yShape.GetShapeSize() != zShape.GetShapeSize(),
80- yShape.GetShapeSize() != dyShape.GetShapeSize() ||80+ OP_LOGE(context, "AcosGrad: shape size mismatch: y=%ld, dy=%ld, z=%ld", yShape.GetShapeSize(),
81- yShape.GetShapeSize() != zShape.GetShapeSize(),81+ dyShape.GetShapeSize(), zShape.GetShapeSize()),
82- OP_LOGE(82+ return ge::GRAPH_FAILED);
83- context,
84- "AcosGrad: shape size mismatch: y=%ld, dy=%ld, z=%ld",
85- yShape.GetShapeSize(), dyShape.GetShapeSize(), zShape.GetShapeSize()),
86- return ge::GRAPH_FAILED);
87 83 
88 totalLength = static_cast<uint64_t>(yShape.GetShapeSize());84 totalLength = static_cast<uint64_t>(yShape.GetShapeSize());
89 85 
@@ -106,8 +102,7 @@ static ge::graphStatus SetWorkspaceSize(gert::TilingContext* context)
106 return ge::GRAPH_SUCCESS;102 return ge::GRAPH_SUCCESS;
107}103}
108 104 
109-static void CalcTilingParams(uint64_t totalLength, uint32_t availCoreNum,105+static void CalcBlockParams(uint64_t totalLength, uint32_t availCoreNum, uint32_t& blockFormer, uint32_t& blockNum)
110- ge::DataType dataType, AcosGradTilingData* tiling)
111{106{
112 uint32_t coreNum = static_cast<uint32_t>(107 uint32_t coreNum = static_cast<uint32_t>(
113 CeilDiv(static_cast<int64_t>(totalLength), static_cast<int64_t>(ELEM_ALIGN)));108 CeilDiv(static_cast<int64_t>(totalLength), static_cast<int64_t>(ELEM_ALIGN)));
@@ -120,31 +115,38 @@ static void CalcTilingParams(uint64_t totalLength, uint32_t availCoreNum,
120 115 
121 uint32_t blockFormerRaw = static_cast<uint32_t>(116 uint32_t blockFormerRaw = static_cast<uint32_t>(
122 CeilDiv(static_cast<int64_t>(totalLength), static_cast<int64_t>(coreNum)));117 CeilDiv(static_cast<int64_t>(totalLength), static_cast<int64_t>(coreNum)));
123- uint32_t blockFormer = static_cast<uint32_t>(118+ blockFormer = static_cast<uint32_t>(
124 CeilDiv(static_cast<int64_t>(blockFormerRaw), static_cast<int64_t>(ELEM_ALIGN)) * ELEM_ALIGN);119 CeilDiv(static_cast<int64_t>(blockFormerRaw), static_cast<int64_t>(ELEM_ALIGN)) * ELEM_ALIGN);
125 if (blockFormer < ELEM_ALIGN) {120 if (blockFormer < ELEM_ALIGN) {
126 blockFormer = ELEM_ALIGN;121 blockFormer = ELEM_ALIGN;
127 }122 }
128 123 
129- uint32_t blockNum = static_cast<uint32_t>(124+ blockNum = static_cast<uint32_t>(CeilDiv(static_cast<int64_t>(totalLength), static_cast<int64_t>(blockFormer)));
130- CeilDiv(static_cast<int64_t>(totalLength), static_cast<int64_t>(blockFormer)));
131 if (blockNum < 1U) {125 if (blockNum < 1U) {
132 blockNum = 1U;126 blockNum = 1U;
133 }127 }
128+}
129+ 
130+static void CalcTilingParams(uint64_t totalLength, uint32_t availCoreNum, ge::DataType dataType,
131+ AcosGradTilingData* tiling)
132+{
133+ uint32_t blockFormer = 0U;
134+ uint32_t blockNum = 0U;
135+ CalcBlockParams(totalLength, availCoreNum, blockFormer, blockNum);
134 136 
135 uint32_t bytesPerElem;137 uint32_t bytesPerElem;
136 uint32_t alignFactor;138 uint32_t alignFactor;
137 139 
138 if (dataType == ge::DT_FLOAT) {140 if (dataType == ge::DT_FLOAT) {
139 bytesPerElem = 32U;141 bytesPerElem = 32U;
140- alignFactor = 64U;142+ alignFactor = 64U;
141 } else {143 } else {
142 bytesPerElem = 28U;144 bytesPerElem = 28U;
143- alignFactor = 128U;145+ alignFactor = 128U;
144 }146 }
145 147 
146 uint32_t ubFormerRaw = static_cast<uint32_t>(UB_SIZE_BYTES / bytesPerElem);148 uint32_t ubFormerRaw = static_cast<uint32_t>(UB_SIZE_BYTES / bytesPerElem);
147- uint32_t ubFormer = static_cast<uint32_t>(149+ uint32_t ubFormer = static_cast<uint32_t>(
148 FloorDiv(static_cast<int64_t>(ubFormerRaw), static_cast<int64_t>(alignFactor)) * alignFactor);150 FloorDiv(static_cast<int64_t>(ubFormerRaw), static_cast<int64_t>(alignFactor)) * alignFactor);
149 if (ubFormer < alignFactor) {151 if (ubFormer < alignFactor) {
150 ubFormer = alignFactor;152 ubFormer = alignFactor;
@@ -154,9 +156,7 @@ static void CalcTilingParams(uint64_t totalLength, uint32_t availCoreNum,
154 }156 }
155 157 
156 uint64_t tailBlockStart = static_cast<uint64_t>(blockNum - 1) * blockFormer;158 uint64_t tailBlockStart = static_cast<uint64_t>(blockNum - 1) * blockFormer;
157- uint32_t tailBlockLen = (tailBlockStart < totalLength)159+ uint32_t tailBlockLen = (tailBlockStart < totalLength) ? static_cast<uint32_t>(totalLength - tailBlockStart) : 0U;
158- ? static_cast<uint32_t>(totalLength - tailBlockStart)
159- : 0U;
160 160 
161 uint32_t ubLoopOfFormerBlock = (ubFormer > 0U) ? (blockFormer / ubFormer) : 0U;161 uint32_t ubLoopOfFormerBlock = (ubFormer > 0U) ? (blockFormer / ubFormer) : 0U;
162 uint32_t ubTailOfFormerBlock = (ubFormer > 0U) ? (blockFormer % ubFormer) : blockFormer;162 uint32_t ubTailOfFormerBlock = (ubFormer > 0U) ? (blockFormer % ubFormer) : blockFormer;
@@ -164,70 +164,62 @@ static void CalcTilingParams(uint64_t totalLength, uint32_t availCoreNum,
164 uint32_t ubLoopOfTailBlock = (ubFormer > 0U && tailBlockLen > 0U) ? (tailBlockLen / ubFormer) : 0U;164 uint32_t ubLoopOfTailBlock = (ubFormer > 0U && tailBlockLen > 0U) ? (tailBlockLen / ubFormer) : 0U;
165 uint32_t ubTailOfTailBlock = (ubFormer > 0U && tailBlockLen > 0U) ? (tailBlockLen % ubFormer) : tailBlockLen;165 uint32_t ubTailOfTailBlock = (ubFormer > 0U && tailBlockLen > 0U) ? (tailBlockLen % ubFormer) : tailBlockLen;
166 166 
167- tiling->totalLength = totalLength;167+ tiling->totalLength = totalLength;
168- tiling->blockFormer = blockFormer;168+ tiling->blockFormer = blockFormer;
169- tiling->blockNum = blockNum;169+ tiling->blockNum = blockNum;
170- tiling->ubFormer = ubFormer;170+ tiling->ubFormer = ubFormer;
171- tiling->ubLoopOfFormerBlock = ubLoopOfFormerBlock;171+ tiling->ubLoopOfFormerBlock = ubLoopOfFormerBlock;
172- tiling->ubTailOfFormerBlock = ubTailOfFormerBlock;172+ tiling->ubTailOfFormerBlock = ubTailOfFormerBlock;
173- tiling->ubLoopOfTailBlock = ubLoopOfTailBlock;173+ tiling->ubLoopOfTailBlock = ubLoopOfTailBlock;
174- tiling->ubTailOfTailBlock = ubTailOfTailBlock;174+ tiling->ubTailOfTailBlock = ubTailOfTailBlock;
175}175}
176 176 
177static ge::graphStatus AcosGradTilingFunc(gert::TilingContext* context)177static ge::graphStatus AcosGradTilingFunc(gert::TilingContext* context)
178{178{
179 OP_LOGI(context->GetNodeName(), "Enter AcosGradTilingFunc");179 OP_LOGI(context->GetNodeName(), "Enter AcosGradTilingFunc");
180- uint64_t ubSize = 0UL;180+ uint64_t ubSize = 0UL;
181 uint32_t coreNum = 0U;181 uint32_t coreNum = 0U;
182- OP_CHECK_IF(182+ OP_CHECK_IF(GetPlatformInfo(context, ubSize, coreNum) != ge::GRAPH_SUCCESS,
183- GetPlatformInfo(context, ubSize, coreNum) != ge::GRAPH_SUCCESS,183+ OP_LOGE(context, "AcosGrad: GetPlatformInfo error"), return ge::GRAPH_FAILED);
184- OP_LOGE(context, "AcosGrad: GetPlatformInfo error"),
185- return ge::GRAPH_FAILED);
186 184 
187 OP_LOGI(context, "[AcosGrad Tiling] coreNum=%u, ubSize=%lu", coreNum, ubSize);185 OP_LOGI(context, "[AcosGrad Tiling] coreNum=%u, ubSize=%lu", coreNum, ubSize);
188 186 
189 uint64_t totalLength = 0UL;187 uint64_t totalLength = 0UL;
190 ge::DataType dataType;188 ge::DataType dataType;
191- OP_CHECK_IF(189+ OP_CHECK_IF(GetShapeInfo(context, totalLength, dataType) != ge::GRAPH_SUCCESS,
192- GetShapeInfo(context, totalLength, dataType) != ge::GRAPH_SUCCESS,190+ OP_LOGE(context, "AcosGrad: GetShapeInfo error"), return ge::GRAPH_FAILED);
193- OP_LOGE(context, "AcosGrad: GetShapeInfo error"),
194- return ge::GRAPH_FAILED);
195 191 
196- OP_CHECK_IF(192+ OP_CHECK_IF(SetWorkspaceSize(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "AcosGrad: SetWorkspaceSize error"),
197- SetWorkspaceSize(context) != ge::GRAPH_SUCCESS,193+ return ge::GRAPH_FAILED);
198- OP_LOGE(context, "AcosGrad: SetWorkspaceSize error"),
199- return ge::GRAPH_FAILED);
200 194 
201 if (totalLength == 0UL) {195 if (totalLength == 0UL) {
202 AcosGradTilingData* tiling = context->GetTilingData<AcosGradTilingData>();196 AcosGradTilingData* tiling = context->GetTilingData<AcosGradTilingData>();
203 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);197 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
204 memset_s(tiling, sizeof(AcosGradTilingData), 0, sizeof(AcosGradTilingData));198 memset_s(tiling, sizeof(AcosGradTilingData), 0, sizeof(AcosGradTilingData));
205 context->SetBlockDim(1U);199 context->SetBlockDim(1U);
206- uint32_t dTypeX = static_cast<uint32_t>(dataType);200+ uint64_t useDoubleBuffer = 0;
207- ASCENDC_TPL_SEL_PARAM(context, dTypeX);201+ ASCENDC_TPL_SEL_PARAM(context, useDoubleBuffer);
208 return ge::GRAPH_SUCCESS;202 return ge::GRAPH_SUCCESS;
209 }203 }
210 204 
211 AcosGradTilingData* tiling = context->GetTilingData<AcosGradTilingData>();205 AcosGradTilingData* tiling = context->GetTilingData<AcosGradTilingData>();
212 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);206 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
213- OP_CHECK_IF(207+ OP_CHECK_IF(memset_s(tiling, sizeof(AcosGradTilingData), 0, sizeof(AcosGradTilingData)) != EOK,
214- memset_s(tiling, sizeof(AcosGradTilingData), 0, sizeof(AcosGradTilingData)) != EOK,208+ OP_LOGE(context, "AcosGrad: memset_s tiling data error"), return ge::GRAPH_FAILED);
215- OP_LOGE(context, "AcosGrad: memset_s tiling data error"),
216- return ge::GRAPH_FAILED);
217 209 
218 CalcTilingParams(totalLength, coreNum, dataType, tiling);210 CalcTilingParams(totalLength, coreNum, dataType, tiling);
219 211 
220 context->SetBlockDim(tiling->blockNum);212 context->SetBlockDim(tiling->blockNum);
221 213 
222- OP_LOGI(context,214+ uint64_t useDoubleBuffer = (totalLength > 1024UL) ? 1 : 0;
223- "[AcosGrad Tiling] totalLength=%lu, blockFormer=%u, blockNum=%u, ubFormer=%u, "215+ ASCENDC_TPL_SEL_PARAM(context, useDoubleBuffer);
224- "ubLoopFormer=%u, ubTailFormer=%u, ubLoopTail=%u, ubTailTail=%u",216+ 
225- tiling->totalLength, tiling->blockFormer, tiling->blockNum, tiling->ubFormer,217+ OP_LOGI(context,
226- tiling->ubLoopOfFormerBlock, tiling->ubTailOfFormerBlock,218+ "[AcosGrad Tiling] totalLength=%lu, blockFormer=%u, blockNum=%u, ubFormer=%u, "
227- tiling->ubLoopOfTailBlock, tiling->ubTailOfTailBlock);219+ "ubLoopFormer=%u, ubTailFormer=%u, ubLoopTail=%u, ubTailTail=%u",
220+ tiling->totalLength, tiling->blockFormer, tiling->blockNum, tiling->ubFormer, tiling->ubLoopOfFormerBlock,
221+ tiling->ubTailOfFormerBlock, tiling->ubLoopOfTailBlock, tiling->ubTailOfTailBlock);
228 222 
229- uint32_t dTypeX = static_cast<uint32_t>(dataType);
230- ASCENDC_TPL_SEL_PARAM(context, dTypeX);
231 return ge::GRAPH_SUCCESS;223 return ge::GRAPH_SUCCESS;
232}224}
233 225 
@@ -20,17 +20,29 @@
20 * z : output gradient (canndev OUTPUT(z))20 * z : output gradient (canndev OUTPUT(z))
21 *21 *
22 * Formula: z = -dy / sqrt(1 - y^2)22 * Formula: z = -dy / sqrt(1 - y^2)
23+ *
24+ * def 驱动 dtype 模式:dtype 由 def.cpp DataType 列表驱动,
25+ * 构建系统注入 DTYPE_Y 编译宏(按输入名 y 大写),kernel 直接使用 DTYPE_Y 获取实际类型。
23 */26 */
24 27 
25#include "arch35/acos_grad.h"28#include "arch35/acos_grad.h"
26 29 
27-template <typename D_T>30+#ifdef __CCE_KT_TEST__
28-__global__ __aicore__ void acos_grad(GM_ADDR y, GM_ADDR dy, GM_ADDR z,31+extern "C" __global__ __aicore__ void acos_grad(GM_ADDR y, GM_ADDR dy, GM_ADDR z, GM_ADDR workspace, GM_ADDR tiling)
29- GM_ADDR workspace, GM_ADDR tiling)
30{32{
31- REGISTER_TILING_DEFAULT(AcosGradTilingData);
32 GET_TILING_DATA_WITH_STRUCT(AcosGradTilingData, tilingData, tiling);33 GET_TILING_DATA_WITH_STRUCT(AcosGradTilingData, tilingData, tiling);
33- NsAcosGrad::KernelAcosGrad<D_T> op;34+ NsAcosGrad::KernelAcosGrad<DTYPE_Y> op;
34 op.Init(y, dy, z, &tilingData);35 op.Init(y, dy, z, &tilingData);
35 op.Process();36 op.Process();
36}37}
38+#else
39+template <int BUFFER_MODE>
40+__global__ __aicore__ void acos_grad(GM_ADDR y, GM_ADDR dy, GM_ADDR z, GM_ADDR workspace, GM_ADDR tiling)
41+{
42+ REGISTER_TILING_DEFAULT(AcosGradTilingData);
43+ GET_TILING_DATA_WITH_STRUCT(AcosGradTilingData, tilingData, tiling);
44+ NsAcosGrad::KernelAcosGrad<DTYPE_Y> op;
45+ op.Init(y, dy, z, &tilingData);
46+ op.Process();
47+}
48+#endif
@@ -12,28 +12,24 @@
12 12 
13/*!13/*!
14 * \file acos_grad_tiling_key.h14 * \file acos_grad_tiling_key.h
15- * \brief AcosGrad TilingKey definition15+ * \brief AcosGrad TilingKey template parameter definition
16+ *
17+ * Template parameters:
18+ * - BUFFER_MODE: Buffer mode (0=single buffer, 1=double buffer)
19+ *
20+ * dtype 由 def.cpp 的 DataType({DT_FLOAT16, DT_FLOAT, DT_BF16}) 驱动,
21+ * 构建系统通过 DTYPE_Y 宏注入实际类型,TilingKey 不再重复编码 dtype。
16 */22 */
17 23 
18#ifndef ACOS_GRAD_TILING_KEY_H24#ifndef ACOS_GRAD_TILING_KEY_H
19#define ACOS_GRAD_TILING_KEY_H25#define ACOS_GRAD_TILING_KEY_H
20 26 
27+#ifndef __CCE_KT_TEST__
21#include "ascendc/host_api/tiling/template_argument.h"28#include "ascendc/host_api/tiling/template_argument.h"
22 29 
23-ASCENDC_TPL_ARGS_DECL(AcosGrad,30+ASCENDC_TPL_ARGS_DECL(AcosGrad, ASCENDC_TPL_UINT_DECL(BUFFER_MODE, 8, ASCENDC_TPL_UI_LIST, 0, 1));
24- ASCENDC_TPL_DATATYPE_DECL(D_T, C_DT_FLOAT, C_DT_FLOAT16, C_DT_BF16, ASCENDC_TPL_INPUT(0)),
25-);
26 31 
27-ASCENDC_TPL_SEL(32+ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0, 1)), );
28- ASCENDC_TPL_ARGS_SEL(33+#endif
29- ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT)
30- ),
31- ASCENDC_TPL_ARGS_SEL(
32- ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT16)
33- ),
34- ASCENDC_TPL_ARGS_SEL(
35- ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_BF16)
36- ),
37-);
38 34 
39#endif // ACOS_GRAD_TILING_KEY_H35#endif // ACOS_GRAD_TILING_KEY_H
@@ -16,9 +16,9 @@
16 *16 *
17 * Covers paths in op_host/arch35/acos_grad_tiling_arch35.cpp:17 * Covers paths in op_host/arch35/acos_grad_tiling_arch35.cpp:
18 * 1) Dtype paths:18 * 1) Dtype paths:
19- * FP32 -> TilingKey 019+ * FP32 -> TilingKey 0 (small) or 1 (large, double buffer)
20- * FP16 -> TilingKey 120+ * FP16 -> TilingKey 0 (small) or 1 (large, double buffer)
21- * BF16 -> TilingKey 2721+ * BF16 -> TilingKey 0 (small) or 1 (large, double buffer)
22 * 2) Multi-core path (large shape) vs single-core path (small shape).22 * 2) Multi-core path (large shape) vs single-core path (small shape).
23 * 3) Non-aligned tail.23 * 3) Non-aligned tail.
24 * 4) Empty tensor (totalLength==0) → early-return, all fields zero.24 * 4) Empty tensor (totalLength==0) → early-return, all fields zero.
@@ -57,37 +57,30 @@ using namespace std;
57 57 
58class AcosGradTilingTest : public testing::Test {58class AcosGradTilingTest : public testing::Test {
59protected:59protected:
60- static void SetUpTestCase()60+ static void SetUpTestCase() { std::cout << "AcosGradTilingTest SetUp" << std::endl; }
61- {
62- std::cout << "AcosGradTilingTest SetUp" << std::endl;
63- }
64 61 
65- static void TearDownTestCase()62+ static void TearDownTestCase() { std::cout << "AcosGradTilingTest TearDown" << std::endl; }
66- {
67- std::cout << "AcosGradTilingTest TearDown" << std::endl;
68- }
69};63};
70 64 
71// ===========================================================================65// ===========================================================================
72// 1) FP32 multi-core aligned — 8192 elem {1,64,2,64}66// 1) FP32 multi-core aligned — 8192 elem {1,64,2,64}
73// coreNum=16, blockFormer=512, blockNum=16, ubFormer=512,67// coreNum=16, blockFormer=512, blockNum=16, ubFormer=512,
74// ubLoopFormer=1, ubTailFormer=0, ubLoopTail=1, ubTailTail=068// ubLoopFormer=1, ubTailFormer=0, ubLoopTail=1, ubTailTail=0
75-// TilingKey 0 = FP3269+// TilingKey 1 = double buffer (totalLength=8192 > 1024)
76// ===========================================================================70// ===========================================================================
77TEST_F(AcosGradTilingTest, test_tiling_fp32_multi_core_001)71TEST_F(AcosGradTilingTest, test_tiling_fp32_multi_core_001)
78{72{
79 optiling::AcosGradCompileInfo compileInfo;73 optiling::AcosGradCompileInfo compileInfo;
80- gert::TilingContextPara tilingContextPara(74+ gert::TilingContextPara tilingContextPara("AcosGrad",
81- "AcosGrad",75+ {
82- {76+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND}, // y
83- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND}, // y77+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND}, // dy
84- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND}, // dy78+ },
85- },79+ {
86- {80+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND}, // z
87- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND}, // z81+ },
88- },82+ &compileInfo);
89- &compileInfo);83+ uint64_t expectTilingKey = 1;
90- uint64_t expectTilingKey = 0;
91 string expectTilingData = "8192 68719477248 4294967808 4294967296 0 ";84 string expectTilingData = "8192 68719477248 4294967808 4294967296 0 ";
92 std::vector<size_t> expectWorkspaces = {0};85 std::vector<size_t> expectWorkspaces = {0};
93 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);86 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
@@ -95,21 +88,20 @@ TEST_F(AcosGradTilingTest, test_tiling_fp32_multi_core_001)
95 88 
96// ===========================================================================89// ===========================================================================
97// 2) FP16 multi-core aligned — same shape as above90// 2) FP16 multi-core aligned — same shape as above
98-// TilingKey 1 = FP1691+// TilingKey 1 = double buffer (totalLength=8192 > 1024)
99// ===========================================================================92// ===========================================================================
100TEST_F(AcosGradTilingTest, test_tiling_fp16_multi_core_002)93TEST_F(AcosGradTilingTest, test_tiling_fp16_multi_core_002)
101{94{
102 optiling::AcosGradCompileInfo compileInfo;95 optiling::AcosGradCompileInfo compileInfo;
103- gert::TilingContextPara tilingContextPara(96+ gert::TilingContextPara tilingContextPara("AcosGrad",
104- "AcosGrad",97+ {
105- {98+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},
106- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},99+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},
107- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},100+ },
108- },101+ {
109- {102+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},
110- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},103+ },
111- },104+ &compileInfo);
112- &compileInfo);
113 uint64_t expectTilingKey = 1;105 uint64_t expectTilingKey = 1;
114 string expectTilingData = "8192 68719477248 4294967808 4294967296 0 ";106 string expectTilingData = "8192 68719477248 4294967808 4294967296 0 ";
115 std::vector<size_t> expectWorkspaces = {0};107 std::vector<size_t> expectWorkspaces = {0};
@@ -118,22 +110,21 @@ TEST_F(AcosGradTilingTest, test_tiling_fp16_multi_core_002)
118 110 
119// ===========================================================================111// ===========================================================================
120// 3) BF16 multi-core aligned112// 3) BF16 multi-core aligned
121-// TilingKey 27 = BF16113+// TilingKey 1 = double buffer (totalLength=8192 > 1024)
122// ===========================================================================114// ===========================================================================
123TEST_F(AcosGradTilingTest, test_tiling_bf16_multi_core_003)115TEST_F(AcosGradTilingTest, test_tiling_bf16_multi_core_003)
124{116{
125 optiling::AcosGradCompileInfo compileInfo;117 optiling::AcosGradCompileInfo compileInfo;
126- gert::TilingContextPara tilingContextPara(118+ gert::TilingContextPara tilingContextPara("AcosGrad",
127- "AcosGrad",119+ {
128- {120+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},
129- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},121+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},
130- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},122+ },
131- },123+ {
132- {124+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},
133- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},125+ },
134- },126+ &compileInfo);
135- &compileInfo);127+ uint64_t expectTilingKey = 1;
136- uint64_t expectTilingKey = 27;
137 string expectTilingData = "8192 68719477248 4294967808 4294967296 0 ";128 string expectTilingData = "8192 68719477248 4294967808 4294967296 0 ";
138 std::vector<size_t> expectWorkspaces = {0};129 std::vector<size_t> expectWorkspaces = {0};
139 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);130 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
@@ -147,16 +138,15 @@ TEST_F(AcosGradTilingTest, test_tiling_bf16_multi_core_003)
147TEST_F(AcosGradTilingTest, test_tiling_fp32_small_tail_004)138TEST_F(AcosGradTilingTest, test_tiling_fp32_small_tail_004)
148{139{
149 optiling::AcosGradCompileInfo compileInfo;140 optiling::AcosGradCompileInfo compileInfo;
150- gert::TilingContextPara tilingContextPara(141+ gert::TilingContextPara tilingContextPara("AcosGrad",
151- "AcosGrad",142+ {
152- {143+ {{{7}, {7}}, ge::DT_FLOAT, ge::FORMAT_ND},
153- {{{7}, {7}}, ge::DT_FLOAT, ge::FORMAT_ND},144+ {{{7}, {7}}, ge::DT_FLOAT, ge::FORMAT_ND},
154- {{{7}, {7}}, ge::DT_FLOAT, ge::FORMAT_ND},145+ },
155- },146+ {
156- {147+ {{{7}, {7}}, ge::DT_FLOAT, ge::FORMAT_ND},
157- {{{7}, {7}}, ge::DT_FLOAT, ge::FORMAT_ND},148+ },
158- },149+ &compileInfo);
159- &compileInfo);
160 uint64_t expectTilingKey = 0;150 uint64_t expectTilingKey = 0;
161 string expectTilingData = "7 4294967808 4294967808 0 7 ";151 string expectTilingData = "7 4294967808 4294967808 0 7 ";
162 std::vector<size_t> expectWorkspaces = {0};152 std::vector<size_t> expectWorkspaces = {0};
@@ -171,17 +161,16 @@ TEST_F(AcosGradTilingTest, test_tiling_fp32_small_tail_004)
171TEST_F(AcosGradTilingTest, test_tiling_fp16_unalign_005)161TEST_F(AcosGradTilingTest, test_tiling_fp16_unalign_005)
172{162{
173 optiling::AcosGradCompileInfo compileInfo;163 optiling::AcosGradCompileInfo compileInfo;
174- gert::TilingContextPara tilingContextPara(164+ gert::TilingContextPara tilingContextPara("AcosGrad",
175- "AcosGrad",165+ {
176- {166+ {{{17}, {17}}, ge::DT_FLOAT16, ge::FORMAT_ND},
177- {{{17}, {17}}, ge::DT_FLOAT16, ge::FORMAT_ND},167+ {{{17}, {17}}, ge::DT_FLOAT16, ge::FORMAT_ND},
178- {{{17}, {17}}, ge::DT_FLOAT16, ge::FORMAT_ND},168+ },
179- },169+ {
180- {170+ {{{17}, {17}}, ge::DT_FLOAT16, ge::FORMAT_ND},
181- {{{17}, {17}}, ge::DT_FLOAT16, ge::FORMAT_ND},171+ },
182- },172+ &compileInfo);
183- &compileInfo);173+ uint64_t expectTilingKey = 0;
184- uint64_t expectTilingKey = 1;
185 string expectTilingData = "17 4294967808 4294967808 0 17 ";174 string expectTilingData = "17 4294967808 4294967808 0 17 ";
186 std::vector<size_t> expectWorkspaces = {0};175 std::vector<size_t> expectWorkspaces = {0};
187 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);176 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
@@ -195,17 +184,16 @@ TEST_F(AcosGradTilingTest, test_tiling_fp16_unalign_005)
195TEST_F(AcosGradTilingTest, test_tiling_fp32_large_multi_core_006)184TEST_F(AcosGradTilingTest, test_tiling_fp32_large_multi_core_006)
196{185{
197 optiling::AcosGradCompileInfo compileInfo;186 optiling::AcosGradCompileInfo compileInfo;
198- gert::TilingContextPara tilingContextPara(187+ gert::TilingContextPara tilingContextPara("AcosGrad",
199- "AcosGrad",188+ {
200- {189+ {{{416910}, {416910}}, ge::DT_FLOAT, ge::FORMAT_ND},
201- {{{416910}, {416910}}, ge::DT_FLOAT, ge::FORMAT_ND},190+ {{{416910}, {416910}}, ge::DT_FLOAT, ge::FORMAT_ND},
202- {{{416910}, {416910}}, ge::DT_FLOAT, ge::FORMAT_ND},191+ },
203- },192+ {
204- {193+ {{{416910}, {416910}}, ge::DT_FLOAT, ge::FORMAT_ND},
205- {{{416910}, {416910}}, ge::DT_FLOAT, ge::FORMAT_ND},194+ },
206- },195+ &compileInfo);
207- &compileInfo);196+ uint64_t expectTilingKey = 1;
208- uint64_t expectTilingKey = 0;
209 string expectTilingData = "416910 270582946304 4294973184 768 4238 ";197 string expectTilingData = "416910 270582946304 4294973184 768 4238 ";
210 std::vector<size_t> expectWorkspaces = {0};198 std::vector<size_t> expectWorkspaces = {0};
211 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);199 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
@@ -224,16 +212,15 @@ TEST_F(AcosGradTilingTest, test_tiling_fp32_large_multi_core_006)
224TEST_F(AcosGradTilingTest, test_tiling_fp16_2d_large_multi_core_007)212TEST_F(AcosGradTilingTest, test_tiling_fp16_2d_large_multi_core_007)
225{213{
226 optiling::AcosGradCompileInfo compileInfo;214 optiling::AcosGradCompileInfo compileInfo;
227- gert::TilingContextPara tilingContextPara(215+ gert::TilingContextPara tilingContextPara("AcosGrad",
228- "AcosGrad",216+ {
229- {217+ {{{60882, 23}, {60882, 23}}, ge::DT_FLOAT16, ge::FORMAT_ND},
230- {{{60882, 23}, {60882, 23}}, ge::DT_FLOAT16, ge::FORMAT_ND},218+ {{{60882, 23}, {60882, 23}}, ge::DT_FLOAT16, ge::FORMAT_ND},
231- {{{60882, 23}, {60882, 23}}, ge::DT_FLOAT16, ge::FORMAT_ND},219+ },
232- },220+ {
233- {221+ {{{60882, 23}, {60882, 23}}, ge::DT_FLOAT16, ge::FORMAT_ND},
234- {{{60882, 23}, {60882, 23}}, ge::DT_FLOAT16, ge::FORMAT_ND},222+ },
235- },223+ &compileInfo);
236- &compileInfo);
237 uint64_t expectTilingKey = 1;224 uint64_t expectTilingKey = 1;
238 // totalLength=1400286, blockFormer=22016, blockNum=64225 // totalLength=1400286, blockFormer=22016, blockNum=64
239 // [1] = 22016 | (64<<32) = 22016 + 274877906944 = 274877928960226 // [1] = 22016 | (64<<32) = 22016 + 274877906944 = 274877928960
@@ -251,16 +238,15 @@ TEST_F(AcosGradTilingTest, test_tiling_fp16_2d_large_multi_core_007)
251TEST_F(AcosGradTilingTest, test_tiling_empty_fp32_008)238TEST_F(AcosGradTilingTest, test_tiling_empty_fp32_008)
252{239{
253 optiling::AcosGradCompileInfo compileInfo;240 optiling::AcosGradCompileInfo compileInfo;
254- gert::TilingContextPara tilingContextPara(241+ gert::TilingContextPara tilingContextPara("AcosGrad",
255- "AcosGrad",242+ {
256- {243+ {{{0}, {0}}, ge::DT_FLOAT, ge::FORMAT_ND},
257- {{{0}, {0}}, ge::DT_FLOAT, ge::FORMAT_ND},244+ {{{0}, {0}}, ge::DT_FLOAT, ge::FORMAT_ND},
258- {{{0}, {0}}, ge::DT_FLOAT, ge::FORMAT_ND},245+ },
259- },246+ {
260- {247+ {{{0}, {0}}, ge::DT_FLOAT, ge::FORMAT_ND},
261- {{{0}, {0}}, ge::DT_FLOAT, ge::FORMAT_ND},248+ },
262- },249+ &compileInfo);
263- &compileInfo);
264 uint64_t expectTilingKey = 0;250 uint64_t expectTilingKey = 0;
265 string expectTilingData = "0 0 0 0 0 ";251 string expectTilingData = "0 0 0 0 0 ";
266 std::vector<size_t> expectWorkspaces = {0};252 std::vector<size_t> expectWorkspaces = {0};
@@ -273,17 +259,16 @@ TEST_F(AcosGradTilingTest, test_tiling_empty_fp32_008)
273TEST_F(AcosGradTilingTest, test_tiling_empty_bf16_009)259TEST_F(AcosGradTilingTest, test_tiling_empty_bf16_009)
274{260{
275 optiling::AcosGradCompileInfo compileInfo;261 optiling::AcosGradCompileInfo compileInfo;
276- gert::TilingContextPara tilingContextPara(262+ gert::TilingContextPara tilingContextPara("AcosGrad",
277- "AcosGrad",263+ {
278- {264+ {{{0, 8}, {0, 8}}, ge::DT_BF16, ge::FORMAT_ND},
279- {{{0, 8}, {0, 8}}, ge::DT_BF16, ge::FORMAT_ND},265+ {{{0, 8}, {0, 8}}, ge::DT_BF16, ge::FORMAT_ND},
280- {{{0, 8}, {0, 8}}, ge::DT_BF16, ge::FORMAT_ND},266+ },
281- },267+ {
282- {268+ {{{0, 8}, {0, 8}}, ge::DT_BF16, ge::FORMAT_ND},
283- {{{0, 8}, {0, 8}}, ge::DT_BF16, ge::FORMAT_ND},269+ },
284- },270+ &compileInfo);
285- &compileInfo);271+ uint64_t expectTilingKey = 0;
286- uint64_t expectTilingKey = 27;
287 string expectTilingData = "0 0 0 0 0 ";272 string expectTilingData = "0 0 0 0 0 ";
288 std::vector<size_t> expectWorkspaces = {0};273 std::vector<size_t> expectWorkspaces = {0};
289 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);274 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
@@ -295,16 +280,15 @@ TEST_F(AcosGradTilingTest, test_tiling_empty_bf16_009)
295TEST_F(AcosGradTilingTest, test_tiling_fail_shape_mismatch_010)280TEST_F(AcosGradTilingTest, test_tiling_fail_shape_mismatch_010)
296{281{
297 optiling::AcosGradCompileInfo compileInfo;282 optiling::AcosGradCompileInfo compileInfo;
298- gert::TilingContextPara tilingContextPara(283+ gert::TilingContextPara tilingContextPara("AcosGrad",
299- "AcosGrad",284+ {
300- {285+ {{{16}, {16}}, ge::DT_FLOAT, ge::FORMAT_ND},
301- {{{16}, {16}}, ge::DT_FLOAT, ge::FORMAT_ND},286+ {{{8}, {8}}, ge::DT_FLOAT, ge::FORMAT_ND},
302- {{{8}, {8}}, ge::DT_FLOAT, ge::FORMAT_ND},287+ },
303- },288+ {
304- {289+ {{{16}, {16}}, ge::DT_FLOAT, ge::FORMAT_ND},
305- {{{16}, {16}}, ge::DT_FLOAT, ge::FORMAT_ND},290+ },
306- },291+ &compileInfo);
307- &compileInfo);
308 uint64_t expectTilingKey = 0;292 uint64_t expectTilingKey = 0;
309 string expectTilingData = "";293 string expectTilingData = "";
310 std::vector<size_t> expectWorkspaces = {0};294 std::vector<size_t> expectWorkspaces = {0};
@@ -317,16 +301,15 @@ TEST_F(AcosGradTilingTest, test_tiling_fail_shape_mismatch_010)
317TEST_F(AcosGradTilingTest, test_tiling_fail_unsupported_dtype_011)301TEST_F(AcosGradTilingTest, test_tiling_fail_unsupported_dtype_011)
318{302{
319 optiling::AcosGradCompileInfo compileInfo;303 optiling::AcosGradCompileInfo compileInfo;
320- gert::TilingContextPara tilingContextPara(304+ gert::TilingContextPara tilingContextPara("AcosGrad",
321- "AcosGrad",305+ {
322- {306+ {{{8}, {8}}, ge::DT_DOUBLE, ge::FORMAT_ND},
323- {{{8}, {8}}, ge::DT_DOUBLE, ge::FORMAT_ND},307+ {{{8}, {8}}, ge::DT_DOUBLE, ge::FORMAT_ND},
324- {{{8}, {8}}, ge::DT_DOUBLE, ge::FORMAT_ND},308+ },
325- },309+ {
326- {310+ {{{8}, {8}}, ge::DT_DOUBLE, ge::FORMAT_ND},
327- {{{8}, {8}}, ge::DT_DOUBLE, ge::FORMAT_ND},311+ },
328- },312+ &compileInfo);
329- &compileInfo);
330 uint64_t expectTilingKey = 0;313 uint64_t expectTilingKey = 0;
331 string expectTilingData = "";314 string expectTilingData = "";
332 std::vector<size_t> expectWorkspaces = {0};315 std::vector<size_t> expectWorkspaces = {0};
@@ -29,10 +29,10 @@
29 29 
30namespace optiling {30namespace optiling {
31 31 
32-using Ops::Base::CeilDiv;
33using Ops::Base::CeilAlign;32using Ops::Base::CeilAlign;
34-using Ops::Base::FloorDiv;33+using Ops::Base::CeilDiv;
35using Ops::Base::FloorAlign;34using Ops::Base::FloorAlign;
35+using Ops::Base::FloorDiv;
36using Ops::Base::GetUbBlockSize;36using Ops::Base::GetUbBlockSize;
37 37 
38constexpr uint32_t WS_SYS_SIZE = 0U;38constexpr uint32_t WS_SYS_SIZE = 0U;
@@ -41,7 +41,8 @@ constexpr int64_t MIN_SPLIT_THRESHOLD = 1024;
41 41 
42static const gert::Shape g_vec_1_shape = {1};42static const gert::Shape g_vec_1_shape = {1};
43 43 
44-static inline const gert::Shape EnsureNotScalar(const gert::Shape& in_shape) {44+static inline const gert::Shape EnsureNotScalar(const gert::Shape& in_shape)
45+{
45 if (in_shape.GetDimNum() == 0) {46 if (in_shape.GetDimNum() == 0) {
46 return g_vec_1_shape;47 return g_vec_1_shape;
47 }48 }
@@ -79,12 +80,10 @@ static ge::graphStatus GetShapeAttrsInfo(gert::TilingContext* context, int64_t&
79 auto zShape = EnsureNotScalar(outZ->GetStorageShape());80 auto zShape = EnsureNotScalar(outZ->GetStorageShape());
80 81 
81 // Shape validation: y, dy, z must have same shape82 // Shape validation: y, dy, z must have same shape
82- OP_CHECK_IF(83+ OP_CHECK_IF(yShape.GetShapeSize() != dyShape.GetShapeSize() || yShape.GetShapeSize() != zShape.GetShapeSize(),
83- yShape.GetShapeSize() != dyShape.GetShapeSize() ||84+ OP_LOGE(context, "AsinGrad: shape size mismatch: y=%ld, dy=%ld, z=%ld", yShape.GetShapeSize(),
84- yShape.GetShapeSize() != zShape.GetShapeSize(),85+ dyShape.GetShapeSize(), zShape.GetShapeSize()),
85- OP_LOGE(context, "AsinGrad: shape size mismatch: y=%ld, dy=%ld, z=%ld",86+ return ge::GRAPH_FAILED);
86- yShape.GetShapeSize(), dyShape.GetShapeSize(), zShape.GetShapeSize()),
87- return ge::GRAPH_FAILED);
88 87 
89 totalNum = yShape.GetShapeSize();88 totalNum = yShape.GetShapeSize();
90 89 
@@ -135,22 +134,16 @@ static ge::graphStatus AsinGradTilingFunc(gert::TilingContext* context)
135 // 1. Get platform info134 // 1. Get platform info
136 uint64_t ubSize;135 uint64_t ubSize;
137 int64_t coreNum;136 int64_t coreNum;
138- OP_CHECK_IF(137+ OP_CHECK_IF(GetPlatformInfo(context, ubSize, coreNum) != ge::GRAPH_SUCCESS,
139- GetPlatformInfo(context, ubSize, coreNum) != ge::GRAPH_SUCCESS,138+ OP_LOGE(context, "GetPlatformInfo error"), return ge::GRAPH_FAILED);
140- OP_LOGE(context, "GetPlatformInfo error"),
141- return ge::GRAPH_FAILED);
142 // 2. Get shape and attr info139 // 2. Get shape and attr info
143 int64_t totalNum;140 int64_t totalNum;
144 ge::DataType dataType;141 ge::DataType dataType;
145- OP_CHECK_IF(142+ OP_CHECK_IF(GetShapeAttrsInfo(context, totalNum, dataType) != ge::GRAPH_SUCCESS,
146- GetShapeAttrsInfo(context, totalNum, dataType) != ge::GRAPH_SUCCESS,143+ OP_LOGE(context, "GetShapeAttrsInfo error"), return ge::GRAPH_FAILED);
147- OP_LOGE(context, "GetShapeAttrsInfo error"),
148- return ge::GRAPH_FAILED);
149 // 3. Get workspace size144 // 3. Get workspace size
150- OP_CHECK_IF(145+ OP_CHECK_IF(GetWorkspaceSize(context) != ge::GRAPH_SUCCESS, OP_LOGE(context, "GetWorkspaceSize error"),
151- GetWorkspaceSize(context) != ge::GRAPH_SUCCESS,146+ return ge::GRAPH_FAILED);
152- OP_LOGE(context, "GetWorkspaceSize error"),
153- return ge::GRAPH_FAILED);
154 // Handle empty tensor147 // Handle empty tensor
155 if (totalNum == 0) {148 if (totalNum == 0) {
156 context->SetBlockDim(0);149 context->SetBlockDim(0);
@@ -158,18 +151,15 @@ static ge::graphStatus AsinGradTilingFunc(gert::TilingContext* context)
158 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);151 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
159 memset_s(tiling, sizeof(AsinGradTilingData), 0, sizeof(AsinGradTilingData));152 memset_s(tiling, sizeof(AsinGradTilingData), 0, sizeof(AsinGradTilingData));
160 // Still need to set TilingKey for empty case153 // Still need to set TilingKey for empty case
161- uint32_t dType = static_cast<uint32_t>(dataType);
162 uint64_t useDoubleBuffer = 0;154 uint64_t useDoubleBuffer = 0;
163- ASCENDC_TPL_SEL_PARAM(context, dType, useDoubleBuffer);155+ ASCENDC_TPL_SEL_PARAM(context, useDoubleBuffer);
164 return ge::GRAPH_SUCCESS;156 return ge::GRAPH_SUCCESS;
165 }157 }
166 // 4. Set tiling data158 // 4. Set tiling data
167 AsinGradTilingData* tiling = context->GetTilingData<AsinGradTilingData>();159 AsinGradTilingData* tiling = context->GetTilingData<AsinGradTilingData>();
168 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);160 OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
169- OP_CHECK_IF(161+ OP_CHECK_IF(memset_s(tiling, sizeof(AsinGradTilingData), 0, sizeof(AsinGradTilingData)) != EOK,
170- memset_s(tiling, sizeof(AsinGradTilingData), 0, sizeof(AsinGradTilingData)) != EOK,162+ OP_LOGE(context, "set tiling data error"), return ge::GRAPH_FAILED);
171- OP_LOGE(context, "set tiling data error"),
172- return ge::GRAPH_FAILED);
173 163 
174 int64_t ubBlockSize = GetUbBlockSize(context);164 int64_t ubBlockSize = GetUbBlockSize(context);
175 tiling->totalNum = totalNum;165 tiling->totalNum = totalNum;
@@ -181,9 +171,8 @@ static ge::graphStatus AsinGradTilingFunc(gert::TilingContext* context)
181 tiling->ubFactor = CalcUbFactor(dataType, ubSize, ubBlockSize, useDoubleBuffer);171 tiling->ubFactor = CalcUbFactor(dataType, ubSize, ubBlockSize, useDoubleBuffer);
182 172 
183 context->SetBlockDim(usedCoreNum);173 context->SetBlockDim(usedCoreNum);
184- // 5. Set TilingKey174+ // 5. Set TilingKey — dtype 由 def.cpp 驱动,TilingKey 只编码 BUFFER_MODE
185- uint32_t dType = static_cast<uint32_t>(dataType);175+ ASCENDC_TPL_SEL_PARAM(context, useDoubleBuffer);
186- ASCENDC_TPL_SEL_PARAM(context, dType, useDoubleBuffer);
187 return ge::GRAPH_SUCCESS;176 return ge::GRAPH_SUCCESS;
188}177}
189 178 
@@ -15,8 +15,10 @@
15 * \brief AsinGrad TilingKey template parameter definition15 * \brief AsinGrad TilingKey template parameter definition
16 *16 *
17 * Template parameters:17 * Template parameters:
18- * - D_T: Input data type (C_DT_FLOAT16, C_DT_FLOAT, C_DT_BF16)
19 * - BUFFER_MODE: Buffer mode (0=single buffer, 1=double buffer)18 * - BUFFER_MODE: Buffer mode (0=single buffer, 1=double buffer)
19+ *
20+ * dtype 由 def.cpp 的 DataType({DT_FLOAT16, DT_FLOAT, DT_BF16}) 驱动,
21+ * 构建系统通过 DTYPE_Y 宏注入实际类型,TilingKey 不再重复编码 dtype。
20 */22 */
21 23 
22#ifndef __ASIN_GRAD_TILING_KEY_H__24#ifndef __ASIN_GRAD_TILING_KEY_H__
@@ -25,25 +27,9 @@
25#ifndef __CCE_KT_TEST__27#ifndef __CCE_KT_TEST__
26#include "ascendc/host_api/tiling/template_argument.h"28#include "ascendc/host_api/tiling/template_argument.h"
27 29 
28-ASCENDC_TPL_ARGS_DECL(AsinGrad,30+ASCENDC_TPL_ARGS_DECL(AsinGrad, ASCENDC_TPL_UINT_DECL(BUFFER_MODE, 8, ASCENDC_TPL_UI_LIST, 0, 1));
29- ASCENDC_TPL_DATATYPE_DECL(D_T, C_DT_FLOAT16, C_DT_FLOAT, C_DT_BF16, ASCENDC_TPL_INPUT(0)),
30- ASCENDC_TPL_UINT_DECL(BUFFER_MODE, 8, ASCENDC_TPL_UI_LIST, 0, 1)
31-);
32 31 
33-ASCENDC_TPL_SEL(32+ASCENDC_TPL_SEL(ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0, 1)), );
34- ASCENDC_TPL_ARGS_SEL(
35- ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT16),
36- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0, 1)
37- ),
38- ASCENDC_TPL_ARGS_SEL(
39- ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_FLOAT),
40- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0, 1)
41- ),
42- ASCENDC_TPL_ARGS_SEL(
43- ASCENDC_TPL_DATATYPE_SEL(D_T, C_DT_BF16),
44- ASCENDC_TPL_UINT_SEL(BUFFER_MODE, ASCENDC_TPL_UI_LIST, 0, 1)
45- ),
46-);
47#endif33#endif
48 34 
49#endif35#endif
@@ -18,15 +18,16 @@
18 * 输出:z (grad input)18 * 输出:z (grad input)
19 *19 *
20 * Template parameters (matching asin_grad_tiling_key.h ASCENDC_TPL_ARGS_DECL):20 * Template parameters (matching asin_grad_tiling_key.h ASCENDC_TPL_ARGS_DECL):
21- * - D_T: Data type, defined by ASCENDC_TPL_DATATYPE_DECL
22 * - BUFFER_MODE: Buffer mode (0=single, 1=double), defined by ASCENDC_TPL_UINT_DECL21 * - BUFFER_MODE: Buffer mode (0=single, 1=double), defined by ASCENDC_TPL_UINT_DECL
22+ *
23+ * dtype 由 def.cpp 驱动,构建系统通过 DTYPE_Y 宏注入实际存储类型,
24+ * kernel 中直接使用 DTYPE_Y 作为 StorageT,无需 TilingKey 编码 dtype。
23 */25 */
24 26 
25#include "arch35/asin_grad.h"27#include "arch35/asin_grad.h"
26 28 
27#ifdef __CCE_KT_TEST__29#ifdef __CCE_KT_TEST__
28-extern "C" __global__ __aicore__ void asin_grad(30+extern "C" __global__ __aicore__ void asin_grad(GM_ADDR y, GM_ADDR dy, GM_ADDR z, GM_ADDR workspace, GM_ADDR tiling)
29- GM_ADDR y, GM_ADDR dy, GM_ADDR z, GM_ADDR workspace, GM_ADDR tiling)
30{31{
31 GET_TILING_DATA_WITH_STRUCT(AsinGradTilingData, tilingData, tiling);32 GET_TILING_DATA_WITH_STRUCT(AsinGradTilingData, tilingData, tiling);
32 NsAsinGrad::AsinGrad<DTYPE_Y, float, 0> op;33 NsAsinGrad::AsinGrad<DTYPE_Y, float, 0> op;
@@ -34,13 +35,13 @@ extern "C" __global__ __aicore__ void asin_grad(
34 op.Process();35 op.Process();
35}36}
36#else37#else
37-template <typename D_T, int BUFFER_MODE>38+template <int BUFFER_MODE>
38__global__ __aicore__ void asin_grad(GM_ADDR y, GM_ADDR dy, GM_ADDR z, GM_ADDR workspace, GM_ADDR tiling)39__global__ __aicore__ void asin_grad(GM_ADDR y, GM_ADDR dy, GM_ADDR z, GM_ADDR workspace, GM_ADDR tiling)
39{40{
40 REGISTER_TILING_DEFAULT(AsinGradTilingData);41 REGISTER_TILING_DEFAULT(AsinGradTilingData);
41 GET_TILING_DATA_WITH_STRUCT(AsinGradTilingData, tilingData, tiling);42 GET_TILING_DATA_WITH_STRUCT(AsinGradTilingData, tilingData, tiling);
42 43 
43- NsAsinGrad::AsinGrad<D_T, float, BUFFER_MODE> op;44+ NsAsinGrad::AsinGrad<DTYPE_Y, float, BUFFER_MODE> op;
44 op.Init(y, dy, z, &tilingData);45 op.Init(y, dy, z, &tilingData);
45 op.Process();46 op.Process();
46}47}