已合并
CleanCode整改——SoftmaxV2、GroupNormV2、LpNormV2 #4031
zhuzixian-lr创建于 4月21日
CleanCode整改——SoftmaxV2、GroupNormV2、LpNormV2 #4031
已合并
zhuzixian-lr创建于 4月21日
7 个文件变更+167-149
@@ -19,11 +19,11 @@
19#include "register/tilingdata_base.h"19#include "register/tilingdata_base.h"
20#include "op_host/tiling_base.h"20#include "op_host/tiling_base.h"
21#include "register/op_impl_registry.h"21#include "register/op_impl_registry.h"
22+#include <string>
22#include <vector>23#include <vector>
23#include <exe_graph/runtime/tiling_context.h>24#include <exe_graph/runtime/tiling_context.h>
24 25 
25using namespace Ops::NN::Optiling;26using namespace Ops::NN::Optiling;
26-using namespace std;
27namespace optiling27namespace optiling
28{28{
29// ar小尾轴29// ar小尾轴
@@ -268,7 +268,7 @@ protected:
268 int64_t yDtypeSize_{0};268 int64_t yDtypeSize_{0};
269 269 
270 int64_t xShapeSize_;270 int64_t xShapeSize_;
271- vector<int64_t> xShape_;271+ std::vector<int64_t> xShape_;
272 272 
273 int64_t a1_{DIM_NUM_ONE};273 int64_t a1_{DIM_NUM_ONE};
274 int64_t r_{DIM_NUM_ONE};274 int64_t r_{DIM_NUM_ONE};
@@ -30,20 +30,18 @@ static const int32_t INDEX_EPSILON = 2;
30static const int32_t INDEX_X = 0;30static const int32_t INDEX_X = 0;
31static const int32_t BYTES_FOR_ALIGN = 1024;31static const int32_t BYTES_FOR_ALIGN = 1024;
32static const int32_t FLOAT32_BYTES = 4;32static const int32_t FLOAT32_BYTES = 4;
33-static const uint64_t INPUT_IDX_X = 0;33+static const int64_t INPUT_IDX_X = 0;
34-static const uint64_t INPUT_IDX_GAMMA = 1;34+static const int64_t INPUT_IDX_GAMMA = 1;
35-static const uint64_t INPUT_IDX_BETA = 2;35+static const int64_t INPUT_IDX_BETA = 2;
36-static const uint64_t PROCESSSIZE = 8192;36+static const int64_t PROCESSSIZE = 8192;
37-static const uint64_t BLOCK_SIZE = 32U;37+static const int64_t RESERVED_WORKSPACE_SIZE_950 = 16L * 1024L * 1024L;
38-static const uint64_t VECTOR_LENGTH = 256U;38+static const int64_t FOUR_BUFFER = 4;
39-static const uint64_t RESERVED_WORKSPACE_SIZE_950 = 16UL * 1024UL * 1024UL;39+static const int64_t BUFFER_NUM = 2;
40-static const uint64_t FOUR_BUFFER = 4;40+static const int64_t DOUBLE_BUFFER = 2;
41-static const uint64_t BUFFER_NUM = 2;41+static const int64_t DICHOTOMY_ADD_COEFF = 2;
42-static const uint64_t DOUBLE_BUFFER = 2;42+static const int64_t ULONG_BIT_LEN = 64;
43-static const uint64_t DICHOTOMY_ADD_COEFF = 2;43+static const int64_t MAX_CHANNEL_SIZE = 4096;
44-static const uint64_t ULONG_BIT_LEN = 64;44+static const int64_t MAX_NUM_PER_CORE = 2048;
45-static const uint64_t MAX_CHANNEL_SIZE = 4096;
46-static const uint64_t MAX_NUM_PER_CORE = 2048;
47static const float DEFAULT_EPS = 1e-5;45static const float DEFAULT_EPS = 1e-5;
48 46 
49inline std::unique_ptr<nlohmann::json> GetCompileInfoJson(gert::TilingParseContext* context) {47inline std::unique_ptr<nlohmann::json> GetCompileInfoJson(gert::TilingParseContext* context) {
@@ -55,13 +53,13 @@ inline std::unique_ptr<nlohmann::json> GetCompileInfoJson(gert::TilingParseConte
55}53}
56 54 
57struct WelfordTilingInitResult {55struct WelfordTilingInitResult {
58- uint64_t loopNum{0};56+ int64_t loopNum{0};
59- uint64_t loopTail{0};57+ int64_t loopTail{0};
60- uint64_t processSize{0};58+ int64_t processSize{0};
61- uint64_t innerLoopNum{0};59+ int64_t innerLoopNum{0};
62- uint64_t innerLoopTail{0};60+ int64_t innerLoopTail{0};
63- uint64_t hwNum{0};61+ int64_t hwNum{0};
64- uint64_t hwNumAlign{0};62+ int64_t hwNumAlign{0};
65 bool checkResult{false};63 bool checkResult{false};
66};64};
67 65 
@@ -97,9 +95,9 @@ inline static int64_t RoundUp(int64_t a, int64_t b)
97static bool isMixType(const gert::TilingContext *context)95static bool isMixType(const gert::TilingContext *context)
98{96{
99 auto xDtype = context->GetInputDesc(INPUT_IDX_X)->GetDataType();97 auto xDtype = context->GetInputDesc(INPUT_IDX_X)->GetDataType();
100- uint64_t xDtypeSize = ge::GetSizeByDataType(xDtype);98+ int64_t xDtypeSize = static_cast<int64_t>(ge::GetSizeByDataType(xDtype));
101 auto gammaDesc = context->GetInputDesc(INPUT_IDX_GAMMA);99 auto gammaDesc = context->GetInputDesc(INPUT_IDX_GAMMA);
102- uint64_t gammaDtypeSize = ge::GetSizeByDataType(gammaDesc->GetDataType());100+ int64_t gammaDtypeSize = static_cast<int64_t>(ge::GetSizeByDataType(gammaDesc->GetDataType()));
103 if (gammaDtypeSize == xDtypeSize) {101 if (gammaDtypeSize == xDtypeSize) {
104 return false;102 return false;
105 }103 }
@@ -108,8 +106,8 @@ static bool isMixType(const gert::TilingContext *context)
108 106 
109static ge::graphStatus CheckInputXShape(const gert::TilingContext *context, const gert::Shape &xShape)107static ge::graphStatus CheckInputXShape(const gert::TilingContext *context, const gert::Shape &xShape)
110{108{
111- uint64_t xDims = xShape.GetDimNum();109+ size_t xDims = xShape.GetDimNum();
112- for (uint64_t i = 0; i < xDims; i++) {110+ for (size_t i = 0; i < xDims; i++) {
113 int64_t curDim = xShape.GetDim(i);111 int64_t curDim = xShape.GetDim(i);
114 OP_CHECK_IF((curDim <= 0),112 OP_CHECK_IF((curDim <= 0),
115 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context->GetNodeName(), "x",113 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context->GetNodeName(), "x",
@@ -142,12 +140,12 @@ static ge::graphStatus CheckInputParams(const gert::TilingContext *context)
142 auto inputX = context->GetInputTensor(INPUT_IDX_X);140 auto inputX = context->GetInputTensor(INPUT_IDX_X);
143 OP_CHECK_NULL_WITH_CONTEXT(context, inputX);141 OP_CHECK_NULL_WITH_CONTEXT(context, inputX);
144 auto xDtype = context->GetInputDesc(INPUT_IDX_X)->GetDataType();142 auto xDtype = context->GetInputDesc(INPUT_IDX_X)->GetDataType();
145- uint64_t xDtypeSize = ge::GetSizeByDataType(xDtype);143+ int64_t xDtypeSize = static_cast<int64_t>(ge::GetSizeByDataType(xDtype));
146 OP_CHECK_IF((xDtypeSize <= 0),144 OP_CHECK_IF((xDtypeSize <= 0),
147- OP_LOGE(context->GetNodeName(), "xDtypeSize is invalid %lu, please check.", xDtypeSize),145+ OP_LOGE(context->GetNodeName(), "xDtypeSize is invalid %ld, please check.", xDtypeSize),
148 return ge::GRAPH_FAILED);146 return ge::GRAPH_FAILED);
149 auto xShape = inputX->GetStorageShape();147 auto xShape = inputX->GetStorageShape();
150- uint64_t channel = xShape.GetDim(DIM_1);148+ int64_t channel = xShape.GetDim(DIM_1);
151 if (CheckInputXShape(context, xShape) != ge::GRAPH_SUCCESS) {149 if (CheckInputXShape(context, xShape) != ge::GRAPH_SUCCESS) {
152 return ge::GRAPH_FAILED;150 return ge::GRAPH_FAILED;
153 }151 }
@@ -171,7 +169,7 @@ static ge::graphStatus CheckInputParams(const gert::TilingContext *context)
171 auto betaShapePtr = context->GetInputShape(INPUT_IDX_BETA);169 auto betaShapePtr = context->GetInputShape(INPUT_IDX_BETA);
172 OP_CHECK_NULL_WITH_CONTEXT(context, betaShapePtr);170 OP_CHECK_NULL_WITH_CONTEXT(context, betaShapePtr);
173 auto betaShape = betaShapePtr->GetStorageShape();171 auto betaShape = betaShapePtr->GetStorageShape();
174- uint64_t betaSizes = betaShape.GetDim(DIM_0);172+ int64_t betaSizes = betaShape.GetDim(DIM_0);
175 OP_CHECK_IF((betaShape.GetDimNum() != 1 || betaSizes != channel),173 OP_CHECK_IF((betaShape.GetDimNum() != 1 || betaSizes != channel),
176 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context->GetNodeName(), "beta",174 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(context->GetNodeName(), "beta",
177 Ops::Base::ToString(betaShape).c_str(),175 Ops::Base::ToString(betaShape).c_str(),
@@ -182,11 +180,11 @@ static ge::graphStatus CheckInputParams(const gert::TilingContext *context)
182 auto gammaDtypePtr = context->GetInputDesc(INPUT_IDX_GAMMA);180 auto gammaDtypePtr = context->GetInputDesc(INPUT_IDX_GAMMA);
183 OP_CHECK_NULL_WITH_CONTEXT(context, gammaDtypePtr);181 OP_CHECK_NULL_WITH_CONTEXT(context, gammaDtypePtr);
184 auto gammaDtype = gammaDtypePtr->GetDataType();182 auto gammaDtype = gammaDtypePtr->GetDataType();
185- uint64_t gammaDtypeSize = ge::GetSizeByDataType(gammaDtype);183+ int64_t gammaDtypeSize = static_cast<int64_t>(ge::GetSizeByDataType(gammaDtype));
186 auto betaDtypePtr = context->GetInputDesc(INPUT_IDX_BETA);184 auto betaDtypePtr = context->GetInputDesc(INPUT_IDX_BETA);
187 OP_CHECK_NULL_WITH_CONTEXT(context, betaDtypePtr);185 OP_CHECK_NULL_WITH_CONTEXT(context, betaDtypePtr);
188 auto betaDtype = betaDtypePtr->GetDataType();186 auto betaDtype = betaDtypePtr->GetDataType();
189- uint64_t betaDtypeSize = ge::GetSizeByDataType(betaDtype);187+ int64_t betaDtypeSize = static_cast<int64_t>(ge::GetSizeByDataType(betaDtype));
190 OP_CHECK_IF((gammaDtypeSize < 0 || gammaDtypeSize != betaDtypeSize),188 OP_CHECK_IF((gammaDtypeSize < 0 || gammaDtypeSize != betaDtypeSize),
191 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "gamma",189 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "gamma",
192 (ge::TypeUtils::DataTypeToSerialString(gammaDtype)).c_str(),190 (ge::TypeUtils::DataTypeToSerialString(gammaDtype)).c_str(),
@@ -203,7 +201,7 @@ static ge::graphStatus CheckAttrParams(const gert::TilingContext *context)
203{201{
204 auto inputX = context->GetInputTensor(INPUT_IDX_X);202 auto inputX = context->GetInputTensor(INPUT_IDX_X);
205 auto xShape = inputX->GetStorageShape();203 auto xShape = inputX->GetStorageShape();
206- uint64_t channel = xShape.GetDim(DIM_1);204+ int64_t channel = xShape.GetDim(DIM_1);
207 // check num_groups205 // check num_groups
208 auto attrs = context->GetAttrs();206 auto attrs = context->GetAttrs();
209 OP_CHECK_NULL_WITH_CONTEXT(context, attrs);207 OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
@@ -221,48 +219,48 @@ static ge::graphStatus CheckAttrParams(const gert::TilingContext *context)
221 return ge::GRAPH_SUCCESS;219 return ge::GRAPH_SUCCESS;
222}220}
223 221 
224-static uint64_t GetOptionalInputTensorSize(const gert::TilingContext *context, uint64_t index,222+static int64_t GetOptionalInputTensorSize(const gert::TilingContext *context, int64_t index,
225- uint64_t specifiedValue = 0)223+ int64_t specifiedValue = 0)
226{224{
227 auto tensorDesc = context->GetInputDesc(index);225 auto tensorDesc = context->GetInputDesc(index);
228 if (tensorDesc == nullptr) {226 if (tensorDesc == nullptr) {
229 return 0;227 return 0;
230 }228 }
231- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());229+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
232- uint32_t blockSize = BLOCK_SIZE;230+ int64_t blockSize = compileInfo->blockSize;
233 auto dtypeSize = ge::GetSizeByDataType(tensorDesc->GetDataType());231 auto dtypeSize = ge::GetSizeByDataType(tensorDesc->GetDataType());
234 if (specifiedValue != 0) {232 if (specifiedValue != 0) {
235- return RoundUp(specifiedValue * dtypeSize, blockSize);233+ return RoundUp(specifiedValue * static_cast<int64_t>(dtypeSize), blockSize);
236 }234 }
237 235 
238 auto storageShape = context->GetInputShape(index);236 auto storageShape = context->GetInputShape(index);
239 OP_CHECK_NULL_WITH_CONTEXT(context, storageShape);237 OP_CHECK_NULL_WITH_CONTEXT(context, storageShape);
240 auto shape = storageShape->GetStorageShape();238 auto shape = storageShape->GetStorageShape();
241- uint64_t num = 1;239+ int64_t num = 1;
242- for (uint64_t i = 0; i < shape.GetDimNum(); i++) {240+ for (size_t i = 0; i < shape.GetDimNum(); i++) {
243 num = num * shape.GetDim(i);241 num = num * shape.GetDim(i);
244 }242 }
245- auto numUbSize = RoundUp(num * dtypeSize, blockSize);243+ auto numUbSize = RoundUp(num * static_cast<int64_t>(dtypeSize), blockSize);
246 return numUbSize;244 return numUbSize;
247}245}
248 246 
249-static void GetDichotomyAddParams(const gert::TilingContext *context, uint64_t r, uint64_t &power, uint64_t &dichotomyK,247+static void GetDichotomyAddParams(const gert::TilingContext *context, int64_t r, int64_t &power, int64_t &dichotomyK,
250- uint64_t &extraSize, uint64_t &lastNum)248+ int64_t &extraSize, int64_t &lastNum)
251{249{
252- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());250+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
253- uint32_t vl = VECTOR_LENGTH / FLOAT32_BYTES;251+ int64_t vl = compileInfo->vectorLength / FLOAT32_BYTES;
254- uint32_t blockSize = BLOCK_SIZE;252+ int64_t blockSize = compileInfo->blockSize;
255- uint64_t basePower = (1L << (ULONG_BIT_LEN - 1 - __builtin_clzl(r)));253+ int64_t basePower = (1L << (ULONG_BIT_LEN - 1 - __builtin_clzl(static_cast<uint64_t>(r))));
256 power = basePower == r ? basePower / DICHOTOMY_ADD_COEFF : basePower;254 power = basePower == r ? basePower / DICHOTOMY_ADD_COEFF : basePower;
257- uint64_t extraOriSize = power / vl;255+ int64_t extraOriSize = power / vl;
258 extraSize = RoundUp(extraOriSize * FLOAT32_BYTES, blockSize);256 extraSize = RoundUp(extraOriSize * FLOAT32_BYTES, blockSize);
259 dichotomyK = 0;257 dichotomyK = 0;
260 if (extraOriSize < vl) {258 if (extraOriSize < vl) {
261 lastNum = extraOriSize;259 lastNum = extraOriSize;
262 return;260 return;
263 }261 }
264- uint64_t totalNum = extraOriSize / vl;262+ int64_t totalNum = extraOriSize / vl;
265- uint64_t base = 1;263+ int64_t base = 1;
266 lastNum = vl;264 lastNum = vl;
267 while (base < totalNum) {265 while (base < totalNum) {
268 dichotomyK++;266 dichotomyK++;
@@ -286,9 +284,9 @@ static ge::graphStatus SetTilingParams(const gert::TilingContext *context, Group
286{284{
287 auto inputX = context->GetInputTensor(INPUT_IDX_X);285 auto inputX = context->GetInputTensor(INPUT_IDX_X);
288 auto xShape = inputX->GetStorageShape();286 auto xShape = inputX->GetStorageShape();
289- uint64_t hwNum = 1;287+ int64_t hwNum = 1;
290- uint64_t xDims = xShape.GetDimNum();288+ size_t xDims = xShape.GetDimNum();
291- for (uint64_t i = 2; i < xDims; i++) {289+ for (size_t i = 2; i < xDims; i++) {
292 hwNum = hwNum * xShape.GetDim(i);290 hwNum = hwNum * xShape.GetDim(i);
293 }291 }
294 tilingData.set_shapeC(xShape.GetDim(DIM_1));292 tilingData.set_shapeC(xShape.GetDim(DIM_1));
@@ -301,15 +299,15 @@ static ge::graphStatus SetTilingParams(const gert::TilingContext *context, Group
301 299 
302static ge::graphStatus SetBlockTiling(const gert::TilingContext *context, GroupNormV2TilingData &tilingData)300static ge::graphStatus SetBlockTiling(const gert::TilingContext *context, GroupNormV2TilingData &tilingData)
303{301{
304- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());302+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
305 auto inputX = context->GetInputTensor(INPUT_IDX_X);303 auto inputX = context->GetInputTensor(INPUT_IDX_X);
306 auto xShape = inputX->GetStorageShape();304 auto xShape = inputX->GetStorageShape();
307- uint64_t shapeN = xShape.GetDim(DIM_0);305+ int64_t shapeN = xShape.GetDim(DIM_0);
308 tilingData.set_numPerCore(CeilDiv(shapeN * tilingData.get_numGroups(), compileInfo->coreNum));306 tilingData.set_numPerCore(CeilDiv(shapeN * tilingData.get_numGroups(), compileInfo->coreNum));
309 tilingData.set_realCoreNum(CeilDiv(shapeN * tilingData.get_numGroups(), tilingData.get_numPerCore()));307 tilingData.set_realCoreNum(CeilDiv(shapeN * tilingData.get_numGroups(), tilingData.get_numPerCore()));
310 tilingData.set_numLastCore(shapeN * tilingData.get_numGroups() -308 tilingData.set_numLastCore(shapeN * tilingData.get_numGroups() -
311 tilingData.get_numPerCore() * (tilingData.get_realCoreNum() - 1));309 tilingData.get_numPerCore() * (tilingData.get_realCoreNum() - 1));
312- uint64_t xShapeSize = xShape.GetShapeSize();310+ int64_t xShapeSize = xShape.GetShapeSize();
313 if (xShapeSize == 0) {311 if (xShapeSize == 0) {
314 tilingData.set_realCoreNum(-1);312 tilingData.set_realCoreNum(-1);
315 }313 }
@@ -340,24 +338,24 @@ static void SetUbTiling(GroupNormV2TilingData &tilingData)
340 在非全载模板下,二分累加的UB额外空间不会影响normalize+swish阶段一次可载入的R轴大小338 在非全载模板下,二分累加的UB额外空间不会影响normalize+swish阶段一次可载入的R轴大小
341 当Gamma或者Beta非空,并且和输入数据类型不一致时,认为是mix type场景339 当Gamma或者Beta非空,并且和输入数据类型不一致时,认为是mix type场景
342*/340*/
343-static void SetTilingKey4Ascend950(const gert::TilingContext *context, uint64_t &maxReduceCount, uint64_t &ubRemain,341+static void SetTilingKey4Ascend950(const gert::TilingContext *context, int64_t &maxReduceCount, int64_t &ubRemain,
344 bool &isReduceFullLoad, GroupNormV2TilingData &tilingData)342 bool &isReduceFullLoad, GroupNormV2TilingData &tilingData)
345{343{
346- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());344+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
347- uint64_t ubSize = compileInfo->ubSize;345+ int64_t ubSize = compileInfo->ubSize;
348- uint32_t blockSize = BLOCK_SIZE;346+ int64_t blockSize = compileInfo->blockSize;
349- uint64_t reduceCount = tilingData.get_shapeD() * tilingData.get_hwNum();347+ int64_t reduceCount = tilingData.get_shapeD() * tilingData.get_hwNum();
350- uint64_t gammaUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA);348+ int64_t gammaUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA);
351- uint64_t betaUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA);349+ int64_t betaUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA);
352- uint64_t realNumPerCore = std::min(MAX_NUM_PER_CORE,350+ int64_t realNumPerCore = std::min(MAX_NUM_PER_CORE,
353- static_cast<uint64_t>(std::max(tilingData.get_numPerCore(), tilingData.get_numLastCore())));351+ static_cast<int64_t>(std::max(tilingData.get_numPerCore(), tilingData.get_numLastCore())));
354- uint64_t meanUbSize = RoundUp(realNumPerCore * FLOAT32_BYTES, blockSize);352+ int64_t meanUbSize = RoundUp(realNumPerCore * FLOAT32_BYTES, blockSize);
355- uint64_t rstdUbSize = RoundUp(realNumPerCore * FLOAT32_BYTES, blockSize);353+ int64_t rstdUbSize = RoundUp(realNumPerCore * FLOAT32_BYTES, blockSize);
356 354 
357- uint64_t otherUbSize = gammaUbSize + betaUbSize + meanUbSize + rstdUbSize;355+ int64_t otherUbSize = gammaUbSize + betaUbSize + meanUbSize + rstdUbSize;
358- uint64_t xDtypeSize = ge::GetSizeByDataType(context->GetInputDesc(INPUT_IDX_X)->GetDataType());356+ int64_t xDtypeSize = static_cast<int64_t>(ge::GetSizeByDataType(context->GetInputDesc(INPUT_IDX_X)->GetDataType()));
359- uint64_t meanUbExtraSize = 0;357+ int64_t meanUbExtraSize = 0;
360- uint64_t rstdUbExtraSize = 0;358+ int64_t rstdUbExtraSize = 0;
361 if (xDtypeSize != FLOAT32_BYTES) {359 if (xDtypeSize != FLOAT32_BYTES) {
362 meanUbExtraSize = RoundUp(realNumPerCore * xDtypeSize, blockSize);360 meanUbExtraSize = RoundUp(realNumPerCore * xDtypeSize, blockSize);
363 rstdUbExtraSize = RoundUp(realNumPerCore * xDtypeSize, blockSize);361 rstdUbExtraSize = RoundUp(realNumPerCore * xDtypeSize, blockSize);
@@ -365,17 +363,17 @@ static void SetTilingKey4Ascend950(const gert::TilingContext *context, uint64_t
365 }363 }
366 bool mixType = isMixType(context);364 bool mixType = isMixType(context);
367 365 
368- uint64_t dichotomyAddPower = 0;366+ int64_t dichotomyAddPower = 0;
369- uint64_t dichotomyAddK = 0;367+ int64_t dichotomyAddK = 0;
370- uint64_t dichotomyAddExtraSize = 0;368+ int64_t dichotomyAddExtraSize = 0;
371- uint64_t dichotomyAddLastNum = 0;369+ int64_t dichotomyAddLastNum = 0;
372 GetDichotomyAddParams(context, reduceCount, dichotomyAddPower, dichotomyAddK, dichotomyAddExtraSize,370 GetDichotomyAddParams(context, reduceCount, dichotomyAddPower, dichotomyAddK, dichotomyAddExtraSize,
373 dichotomyAddLastNum);371 dichotomyAddLastNum);
374 otherUbSize += dichotomyAddExtraSize;372 otherUbSize += dichotomyAddExtraSize;
375 373 
376 ubRemain = ubSize <= otherUbSize ? 0 : ubSize - otherUbSize;374 ubRemain = ubSize <= otherUbSize ? 0 : ubSize - otherUbSize;
377 OP_CHECK_IF((xDtypeSize == 0),375 OP_CHECK_IF((xDtypeSize == 0),
378- OP_LOGE(context->GetNodeName(), "XDtypeSize is zero."), return);376+ OP_LOGE(context->GetNodeName(), "xDtypeSize is zero."), return);
379 maxReduceCount = (ubRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;377 maxReduceCount = (ubRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;
380 378 
381 if (maxReduceCount > reduceCount) {379 if (maxReduceCount > reduceCount) {
@@ -385,16 +383,16 @@ static void SetTilingKey4Ascend950(const gert::TilingContext *context, uint64_t
385 tilingData.set_tilingKey(tilingKey);383 tilingData.set_tilingKey(tilingKey);
386 return;384 return;
387 }385 }
388- bool isLargeChannel = static_cast<uint64_t>(tilingData.get_shapeC()) > MAX_CHANNEL_SIZE;386+ bool isLargeChannel = static_cast<int64_t>(tilingData.get_shapeC()) > MAX_CHANNEL_SIZE;
389- uint64_t newUbRemain = ubRemain;387+ int64_t newUbRemain = ubRemain;
390 // 对于Channel轴过大的场景,尝试将gamma大小限制为ShapeD,beta大小限制为shapeD,重新计算最大可全载的R轴388 // 对于Channel轴过大的场景,尝试将gamma大小限制为ShapeD,beta大小限制为shapeD,重新计算最大可全载的R轴
391 // 如果此时仍然无法全载,则走channel过大的非全载模板389 // 如果此时仍然无法全载,则走channel过大的非全载模板
392 if (isLargeChannel) {390 if (isLargeChannel) {
393- uint64_t gammaSplitUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA, tilingData.get_shapeD());391+ int64_t gammaSplitUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA, tilingData.get_shapeD());
394- uint64_t betaSplitUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA, tilingData.get_shapeD());392+ int64_t betaSplitUbSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA, tilingData.get_shapeD());
395 otherUbSize = otherUbSize - gammaUbSize - betaUbSize + gammaSplitUbSize + betaSplitUbSize;393 otherUbSize = otherUbSize - gammaUbSize - betaUbSize + gammaSplitUbSize + betaSplitUbSize;
396 newUbRemain = ubSize <= otherUbSize ? 0 : ubSize - otherUbSize;394 newUbRemain = ubSize <= otherUbSize ? 0 : ubSize - otherUbSize;
397- uint64_t newMaxReduceCount = (newUbRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;395+ int64_t newMaxReduceCount = (newUbRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;
398 if (newMaxReduceCount > reduceCount) {396 if (newMaxReduceCount > reduceCount) {
399 isReduceFullLoad = true;397 isReduceFullLoad = true;
400 maxReduceCount = newMaxReduceCount;398 maxReduceCount = newMaxReduceCount;
@@ -426,11 +424,11 @@ static void SetTilingKey4Ascend950(const gert::TilingContext *context, uint64_t
426 424 
427static void SetDichotomyAddParams(const gert::TilingContext *context, GroupNormV2TilingData &tilingData)425static void SetDichotomyAddParams(const gert::TilingContext *context, GroupNormV2TilingData &tilingData)
428{426{
429- uint64_t reduceCount = tilingData.get_shapeD() * tilingData.get_hwNum();427+ int64_t reduceCount = tilingData.get_shapeD() * tilingData.get_hwNum();
430- uint64_t dichotomyAddPower = 0;428+ int64_t dichotomyAddPower = 0;
431- uint64_t dichotomyAddK = 0;429+ int64_t dichotomyAddK = 0;
432- uint64_t dichotomyAddExtraSize = 0;430+ int64_t dichotomyAddExtraSize = 0;
433- uint64_t dichotomyAddLastNum = 0;431+ int64_t dichotomyAddLastNum = 0;
434 GetDichotomyAddParams(context, reduceCount, dichotomyAddPower, dichotomyAddK, dichotomyAddExtraSize,432 GetDichotomyAddParams(context, reduceCount, dichotomyAddPower, dichotomyAddK, dichotomyAddExtraSize,
435 dichotomyAddLastNum);433 dichotomyAddLastNum);
436 tilingData.set_dichotomyAddPower(dichotomyAddPower);434 tilingData.set_dichotomyAddPower(dichotomyAddPower);
@@ -438,27 +436,27 @@ static void SetDichotomyAddParams(const gert::TilingContext *context, GroupNormV
438 tilingData.set_dichotomyAddLastNum(dichotomyAddLastNum);436 tilingData.set_dichotomyAddLastNum(dichotomyAddLastNum);
439}437}
440 438 
441-static void SetWelfordParallelN(const gert::TilingContext *context, uint64_t xDtypeSize, uint64_t ubRemain,439+static void SetWelfordParallelN(const gert::TilingContext *context, int64_t xDtypeSize, int64_t ubRemain,
442 GroupNormV2TilingData &tilingData)440 GroupNormV2TilingData &tilingData)
443{441{
444- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());442+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
445- uint32_t blockSize = BLOCK_SIZE;443+ int64_t blockSize = compileInfo->blockSize;
446 OP_CHECK_IF((xDtypeSize == 0),444 OP_CHECK_IF((xDtypeSize == 0),
447- OP_LOGE(context->GetNodeName(), "XDtypeSize is zero."), return);445+ OP_LOGE(context->GetNodeName(), "xDtypeSize is zero."), return);
448- uint32_t coeff = FLOAT32_BYTES / xDtypeSize;446+ int64_t coeff = FLOAT32_BYTES / xDtypeSize;
449- uint32_t totalNum = BUFFER_NUM * (coeff + 1);447+ int64_t totalNum = BUFFER_NUM * (coeff + 1);
450- uint32_t welfordBase = blockSize / xDtypeSize;448+ int64_t welfordBase = blockSize / xDtypeSize;
451 OP_CHECK_IF((totalNum == 0),449 OP_CHECK_IF((totalNum == 0),
452 OP_LOGE(context->GetNodeName(), "TotalNum is zero."), return);450 OP_LOGE(context->GetNodeName(), "TotalNum is zero."), return);
453- uint32_t maxParallelN = DownAlign((ubRemain / xDtypeSize) / totalNum, welfordBase);451+ int64_t maxParallelN = DownAlign((ubRemain / xDtypeSize) / totalNum, welfordBase);
454 452 
455- uint64_t dichotomyAddPower = 0;453+ int64_t dichotomyAddPower = 0;
456- uint64_t dichotomyAddK = 0;454+ int64_t dichotomyAddK = 0;
457- uint64_t dichotomyAddExtraSize = 0;455+ int64_t dichotomyAddExtraSize = 0;
458- uint64_t dichotomyAddLastNum = 0;456+ int64_t dichotomyAddLastNum = 0;
459 GetDichotomyAddParams(context, maxParallelN, dichotomyAddPower, dichotomyAddK, dichotomyAddExtraSize,457 GetDichotomyAddParams(context, maxParallelN, dichotomyAddPower, dichotomyAddK, dichotomyAddExtraSize,
460 dichotomyAddLastNum);458 dichotomyAddLastNum);
461- uint32_t ubCurUse =459+ int64_t ubCurUse =
462 maxParallelN * BUFFER_NUM * xDtypeSize + dichotomyAddExtraSize + maxParallelN * BUFFER_NUM * FLOAT32_BYTES;460 maxParallelN * BUFFER_NUM * xDtypeSize + dichotomyAddExtraSize + maxParallelN * BUFFER_NUM * FLOAT32_BYTES;
463 while (ubCurUse > ubRemain) {461 while (ubCurUse > ubRemain) {
464 maxParallelN -= welfordBase;462 maxParallelN -= welfordBase;
@@ -480,28 +478,28 @@ static void SetWelfordParallelN(const gert::TilingContext *context, uint64_t xDt
480}478}
481 479 
482static void SetUbTiling4TwoPass(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,480static void SetUbTiling4TwoPass(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,
483- uint64_t maxReduceCount, uint32_t xDtypeSize)481+ int64_t maxReduceCount, int64_t xDtypeSize)
484{482{
485- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());483+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
486- uint32_t blockSize = BLOCK_SIZE;484+ int64_t blockSize = compileInfo->blockSize;
487- uint64_t elemNum = tilingData.get_elemNum();485+ int64_t elemNum = tilingData.get_elemNum();
488 OP_CHECK_IF((xDtypeSize == 0),486 OP_CHECK_IF((xDtypeSize == 0),
489- OP_LOGE(context->GetNodeName(), "XDtypeSize is zero."), return);487+ OP_LOGE(context->GetNodeName(), "xDtypeSize is zero."), return);
490- uint64_t elemNumAlign = RoundUp(elemNum, blockSize / xDtypeSize);488+ int64_t elemNumAlign = RoundUp(elemNum, blockSize / xDtypeSize);
491 SetDichotomyAddParams(context, tilingData);489 SetDichotomyAddParams(context, tilingData);
492 OP_CHECK_IF((elemNumAlign == 0),490 OP_CHECK_IF((elemNumAlign == 0),
493 OP_LOGE(context->GetNodeName(), "ElemNumAlign is zero."), return);491 OP_LOGE(context->GetNodeName(), "ElemNumAlign is zero."), return);
494- uint64_t count = maxReduceCount / elemNumAlign;492+ int64_t count = maxReduceCount / elemNumAlign;
495- uint64_t processSize = count * elemNumAlign;493+ int64_t processSize = count * elemNumAlign;
496 tilingData.set_processSize(processSize);494 tilingData.set_processSize(processSize);
497}495}
498 496 
499static WelfordTilingInitResult InitWelfordTilingCommon(const gert::TilingContext *context, 497static WelfordTilingInitResult InitWelfordTilingCommon(const gert::TilingContext *context,
500- GroupNormV2TilingData &tilingData, uint32_t blockSize, uint32_t xDtypeSize) {498+ GroupNormV2TilingData &tilingData, int64_t blockSize, int64_t xDtypeSize) {
501 WelfordTilingInitResult result{};499 WelfordTilingInitResult result{};
502 result.hwNum = tilingData.get_hwNum();500 result.hwNum = tilingData.get_hwNum();
503 OP_CHECK_IF((xDtypeSize == 0),501 OP_CHECK_IF((xDtypeSize == 0),
504- OP_LOGE(context->GetNodeName(), "XDtypeSize is zero."), return result);502+ OP_LOGE(context->GetNodeName(), "xDtypeSize is zero."), return result);
505 result.hwNumAlign = RoundUp(result.hwNum, blockSize / xDtypeSize);503 result.hwNumAlign = RoundUp(result.hwNum, blockSize / xDtypeSize);
506 OP_CHECK_IF((result.hwNumAlign == 0),504 OP_CHECK_IF((result.hwNumAlign == 0),
507 OP_LOGE(context->GetNodeName(), "HwNumAlign is zero."), return result);505 OP_LOGE(context->GetNodeName(), "HwNumAlign is zero."), return result);
@@ -510,15 +508,15 @@ static WelfordTilingInitResult InitWelfordTilingCommon(const gert::TilingContext
510}508}
511 509 
512static void SetUbTiling4WelfordPerf(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,510static void SetUbTiling4WelfordPerf(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,
513- uint64_t maxReduceCount, uint32_t ubRemain, uint32_t xDtypeSize)511+ int64_t maxReduceCount, int64_t ubRemain, int64_t xDtypeSize)
514{512{
515- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());513+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
516- uint32_t blockSize = BLOCK_SIZE;514+ int64_t blockSize = compileInfo->blockSize;
517 SetWelfordParallelN(context, xDtypeSize, ubRemain, tilingData);515 SetWelfordParallelN(context, xDtypeSize, ubRemain, tilingData);
518 WelfordTilingInitResult result = InitWelfordTilingCommon(context, tilingData, blockSize, xDtypeSize);516 WelfordTilingInitResult result = InitWelfordTilingCommon(context, tilingData, blockSize, xDtypeSize);
519 OP_CHECK_IF((result.checkResult == false),517 OP_CHECK_IF((result.checkResult == false),
520 OP_LOGE(context->GetNodeName(), "InitWelfordTilingCommon Failed."), return);518 OP_LOGE(context->GetNodeName(), "InitWelfordTilingCommon Failed."), return);
521- uint64_t count = maxReduceCount / result.hwNumAlign;519+ int64_t count = maxReduceCount / result.hwNumAlign;
522 if (count >= 1) {520 if (count >= 1) {
523 result.loopNum = CeilDiv(tilingData.get_shapeD(), count);521 result.loopNum = CeilDiv(tilingData.get_shapeD(), count);
524 result.loopTail = (tilingData.get_shapeD() - (result.loopNum - 1) * count) * result.hwNumAlign;522 result.loopTail = (tilingData.get_shapeD() - (result.loopNum - 1) * count) * result.hwNumAlign;
@@ -539,22 +537,22 @@ static void SetUbTiling4WelfordPerf(const gert::TilingContext *context, GroupNor
539 tilingData.set_innerLoopTail(result.innerLoopTail);537 tilingData.set_innerLoopTail(result.innerLoopTail);
540}538}
541static void SetUbTiling4WelfordGeneralized(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,539static void SetUbTiling4WelfordGeneralized(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,
542- uint32_t ubRemain, uint32_t xDtypeSize)540+ int64_t ubRemain, int64_t xDtypeSize)
543{541{
544- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());542+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
545- uint32_t blockSize = BLOCK_SIZE;543+ int64_t blockSize = compileInfo->blockSize;
546 WelfordTilingInitResult result = InitWelfordTilingCommon(context, tilingData, blockSize, xDtypeSize);544 WelfordTilingInitResult result = InitWelfordTilingCommon(context, tilingData, blockSize, xDtypeSize);
547 OP_CHECK_IF((result.checkResult == false),545 OP_CHECK_IF((result.checkResult == false),
548 OP_LOGE(context->GetNodeName(), "InitWelfordTilingCommon Failed."), return);546 OP_LOGE(context->GetNodeName(), "InitWelfordTilingCommon Failed."), return);
549- uint64_t maxReduceCount = (ubRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;547+ int64_t maxReduceCount = (ubRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;
550- uint64_t count = maxReduceCount / result.hwNumAlign;548+ int64_t count = maxReduceCount / result.hwNumAlign;
551- uint64_t gammaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA, count);549+ int64_t gammaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA, count);
552- uint64_t betaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA, count);550+ int64_t betaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA, count);
553- uint64_t curUbSize = gammaRealSize + betaRealSize + count * result.hwNumAlign * xDtypeSize * BUFFER_NUM * DOUBLE_BUFFER;551+ int64_t curUbSize = gammaRealSize + betaRealSize + count * result.hwNumAlign * xDtypeSize * BUFFER_NUM * DOUBLE_BUFFER;
554 while (curUbSize > ubRemain && count >= 1) {552 while (curUbSize > ubRemain && count >= 1) {
555 count--;553 count--;
556- uint64_t gammaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA, count);554+ int64_t gammaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_GAMMA, count);
557- uint64_t betaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA, count);555+ int64_t betaRealSize = GetOptionalInputTensorSize(context, INPUT_IDX_BETA, count);
558 curUbSize = gammaRealSize + betaRealSize + count * result.hwNumAlign * xDtypeSize * BUFFER_NUM * DOUBLE_BUFFER;556 curUbSize = gammaRealSize + betaRealSize + count * result.hwNumAlign * xDtypeSize * BUFFER_NUM * DOUBLE_BUFFER;
559 }557 }
560 if (count >= 1) {558 if (count >= 1) {
@@ -568,7 +566,7 @@ static void SetUbTiling4WelfordGeneralized(const gert::TilingContext *context, G
568 betaRealSize = blockSize;566 betaRealSize = blockSize;
569 ubRemain = ubRemain - gammaRealSize - betaRealSize;567 ubRemain = ubRemain - gammaRealSize - betaRealSize;
570 maxReduceCount = (ubRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;568 maxReduceCount = (ubRemain / (DOUBLE_BUFFER * BUFFER_NUM)) / xDtypeSize;
571- uint64_t maxReduceCountDownAlign = DownAlign(maxReduceCount, blockSize / xDtypeSize);569+ int64_t maxReduceCountDownAlign = DownAlign(maxReduceCount, blockSize / xDtypeSize);
572 result.innerLoopNum = CeilDiv(result.hwNum, maxReduceCountDownAlign);570 result.innerLoopNum = CeilDiv(result.hwNum, maxReduceCountDownAlign);
573 result.innerLoopTail = result.hwNum - maxReduceCountDownAlign * (result.innerLoopNum - 1);571 result.innerLoopTail = result.hwNum - maxReduceCountDownAlign * (result.innerLoopNum - 1);
574 result.processSize = maxReduceCountDownAlign;572 result.processSize = maxReduceCountDownAlign;
@@ -584,7 +582,7 @@ static void SetUbTiling4WelfordGeneralized(const gert::TilingContext *context, G
584}582}
585 583 
586static void SetUbTiling4Welford(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,584static void SetUbTiling4Welford(const gert::TilingContext *context, GroupNormV2TilingData &tilingData,
587- uint64_t maxReduceCount, uint64_t ubRemain, uint32_t xDtypeSize)585+ int64_t maxReduceCount, int64_t ubRemain, int64_t xDtypeSize)
588{586{
589 if (tilingData.get_tilingKey() == static_cast<int64_t>(GroupNormV2TilingKey::TILINGKEY_WELFORD_PERF) ||587 if (tilingData.get_tilingKey() == static_cast<int64_t>(GroupNormV2TilingKey::TILINGKEY_WELFORD_PERF) ||
590 tilingData.get_tilingKey() == static_cast<int64_t>(GroupNormV2TilingKey::TILINGKEY_WELFORD_PERF_MIX_TYPE)) {588 tilingData.get_tilingKey() == static_cast<int64_t>(GroupNormV2TilingKey::TILINGKEY_WELFORD_PERF_MIX_TYPE)) {
@@ -593,12 +591,12 @@ static void SetUbTiling4Welford(const gert::TilingContext *context, GroupNormV2T
593 return SetUbTiling4WelfordGeneralized(context, tilingData, ubRemain, xDtypeSize);591 return SetUbTiling4WelfordGeneralized(context, tilingData, ubRemain, xDtypeSize);
594}592}
595 593 
596-static void SetUbTiling4Ascend950(const gert::TilingContext *context, uint64_t maxReduceCount, uint64_t ubRemain,594+static void SetUbTiling4Ascend950(const gert::TilingContext *context, int64_t maxReduceCount, int64_t ubRemain,
597 bool isReduceFullLoad, GroupNormV2TilingData &tilingData)595 bool isReduceFullLoad, GroupNormV2TilingData &tilingData)
598{596{
599- auto compileInfo = reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());597+ auto compileInfo = context->GetCompileInfo<GroupNormV2CompileInfo>();
600- int32_t ubSize = compileInfo->ubSize;598+ int64_t ubSize = compileInfo->ubSize;
601- uint64_t xDtypeSize = ge::GetSizeByDataType(context->GetInputDesc(INPUT_IDX_X)->GetDataType());599+ int64_t xDtypeSize = static_cast<int64_t>(ge::GetSizeByDataType(context->GetInputDesc(INPUT_IDX_X)->GetDataType()));
602 tilingData.set_ubSize(ubSize);600 tilingData.set_ubSize(ubSize);
603 if (!isReduceFullLoad) {601 if (!isReduceFullLoad) {
604 SetUbTiling4Welford(context, tilingData, maxReduceCount, ubRemain, xDtypeSize);602 SetUbTiling4Welford(context, tilingData, maxReduceCount, ubRemain, xDtypeSize);
@@ -609,8 +607,8 @@ static void SetUbTiling4Ascend950(const gert::TilingContext *context, uint64_t m
609 607 
610static void SetTilingForAscend950(const gert::TilingContext *context, GroupNormV2TilingData &tilingData)608static void SetTilingForAscend950(const gert::TilingContext *context, GroupNormV2TilingData &tilingData)
611{609{
612- uint64_t maxReduceCount = 0;610+ int64_t maxReduceCount = 0;
613- uint64_t ubRemain = 0;611+ int64_t ubRemain = 0;
614 bool reduceFullLoad = false;612 bool reduceFullLoad = false;
615 SetTilingKey4Ascend950(context, maxReduceCount, ubRemain, reduceFullLoad, tilingData);613 SetTilingKey4Ascend950(context, maxReduceCount, ubRemain, reduceFullLoad, tilingData);
616 SetUbTiling4Ascend950(context, maxReduceCount, ubRemain, reduceFullLoad, tilingData);614 SetUbTiling4Ascend950(context, maxReduceCount, ubRemain, reduceFullLoad, tilingData);
@@ -645,8 +643,7 @@ ge::graphStatus SetTilingData(gert::TilingContext *context)
645 643 
646static ge::graphStatus Tiling4GroupNormV2(gert::TilingContext *context)644static ge::graphStatus Tiling4GroupNormV2(gert::TilingContext *context)
647{645{
648- const GroupNormV2CompileInfo *compile_info =646+ auto compile_info = context->GetCompileInfo<GroupNormV2CompileInfo>();
649- reinterpret_cast<const GroupNormV2CompileInfo *>(context->GetCompileInfo());
650 OP_CHECK_NULL_WITH_CONTEXT(context, compile_info);647 OP_CHECK_NULL_WITH_CONTEXT(context, compile_info);
651 648 
652 // get input shape info649 // get input shape info
@@ -682,9 +679,23 @@ static ge::graphStatus TilingPrepare4GroupNormV2(gert::TilingParseContext *conte
682 OP_CHECK_NULL_WITH_CONTEXT(context, platform_info);679 OP_CHECK_NULL_WITH_CONTEXT(context, platform_info);
683 auto ascendc_platform = platform_ascendc::PlatformAscendC(platform_info);680 auto ascendc_platform = platform_ascendc::PlatformAscendC(platform_info);
684 compile_info->coreNum = ascendc_platform.GetCoreNumAiv();681 compile_info->coreNum = ascendc_platform.GetCoreNumAiv();
682+ OP_CHECK_IF((compile_info->coreNum <= 0),
683+ OP_LOGE(context->GetNodeName(), "Get coreNum failed, coreNum: %d", compile_info->coreNum),
684+ return ge::GRAPH_FAILED);
685 uint64_t ubSize;685 uint64_t ubSize;
686 ascendc_platform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);686 ascendc_platform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
687 compile_info->ubSize = static_cast<int64_t>(ubSize);687 compile_info->ubSize = static_cast<int64_t>(ubSize);
688+ OP_CHECK_IF((compile_info->ubSize <= 0),
689+ OP_LOGE(context->GetNodeName(), "Get ubSize failed, ubSize: %ld", compile_info->ubSize),
690+ return ge::GRAPH_FAILED);
691+ compile_info->blockSize = Ops::Base::GetUbBlockSize(context);
692+ OP_CHECK_IF((compile_info->blockSize <= 0),
693+ OP_LOGE(context->GetNodeName(), "Get blockSize failed, blockSize: %ld", compile_info->blockSize),
694+ return ge::GRAPH_FAILED);
695+ compile_info->vectorLength = Ops::Base::GetVRegSize(context);
696+ OP_CHECK_IF((compile_info->vectorLength <= 0),
697+ OP_LOGE(context->GetNodeName(), "Get vectorLength failed, vectorLength: %ld", compile_info->vectorLength),
698+ return ge::GRAPH_FAILED);
688 return ge::GRAPH_SUCCESS;699 return ge::GRAPH_SUCCESS;
689 }700 }
690 return ge::GRAPH_FAILED;701 return ge::GRAPH_FAILED;
@@ -21,11 +21,14 @@
21#include "register/tilingdata_base.h"21#include "register/tilingdata_base.h"
22#include "op_host/tiling_base.h"22#include "op_host/tiling_base.h"
23#include "op_api/runtime2_util.h"23#include "op_api/runtime2_util.h"
24+#include "op_common/op_host/util/platform_util.h"
24 25 
25namespace optiling {26namespace optiling {
26struct GroupNormV2CompileInfo {27struct GroupNormV2CompileInfo {
27- int32_t coreNum;28+ int32_t coreNum = 0;
28- uint64_t ubSize;29+ int64_t ubSize = 0;
30+ int64_t blockSize = 0;
31+ int64_t vectorLength = 0;
29};32};
30 33 
31BEGIN_TILING_DATA_DEF(GroupNormV2TilingData)34BEGIN_TILING_DATA_DEF(GroupNormV2TilingData)
@@ -88,6 +88,8 @@ TEST_F(GroupNormV2Tiling, GroupNormV2_tiling_0)
88 optiling::GroupNormV2CompileInfo compile_info;88 optiling::GroupNormV2CompileInfo compile_info;
89 compile_info.coreNum = 64;89 compile_info.coreNum = 64;
90 compile_info.ubSize = 245760;90 compile_info.ubSize = 245760;
91+ compile_info.blockSize = 32;
92+ compile_info.vectorLength = 256;
91 93 
92 // tilingFunc simulate94 // tilingFunc simulate
93 auto param = gert::TilingData::CreateCap(4096);95 auto param = gert::TilingData::CreateCap(4096);
@@ -156,6 +158,8 @@ TEST_F(GroupNormV2Tiling, GroupNormV2_tiling_1)
156 optiling::GroupNormV2CompileInfo compile_info;158 optiling::GroupNormV2CompileInfo compile_info;
157 compile_info.coreNum = 64;159 compile_info.coreNum = 64;
158 compile_info.ubSize = 262114;160 compile_info.ubSize = 262114;
161+ compile_info.blockSize = 32;
162+ compile_info.vectorLength = 256;
159 163 
160 // tilingFunc simulate164 // tilingFunc simulate
161 auto param = gert::TilingData::CreateCap(4096);165 auto param = gert::TilingData::CreateCap(4096);
@@ -220,6 +224,8 @@ TEST_F(GroupNormV2Tiling, GroupNormV2_tiling_2)
220 optiling::GroupNormV2CompileInfo compile_info;224 optiling::GroupNormV2CompileInfo compile_info;
221 compile_info.coreNum = 64;225 compile_info.coreNum = 64;
222 compile_info.ubSize = 262114;226 compile_info.ubSize = 262114;
227+ compile_info.blockSize = 32;
228+ compile_info.vectorLength = 256;
223 229 
224 // tilingFunc simulate230 // tilingFunc simulate
225 auto param = gert::TilingData::CreateCap(4096);231 auto param = gert::TilingData::CreateCap(4096);
@@ -284,6 +290,8 @@ TEST_F(GroupNormV2Tiling, GroupNormV2_tiling_3)
284 optiling::GroupNormV2CompileInfo compile_info;290 optiling::GroupNormV2CompileInfo compile_info;
285 compile_info.coreNum = 64;291 compile_info.coreNum = 64;
286 compile_info.ubSize = 245760;292 compile_info.ubSize = 245760;
293+ compile_info.blockSize = 32;
294+ compile_info.vectorLength = 256;
287 295 
288 // tilingFunc simulate296 // tilingFunc simulate
289 auto param = gert::TilingData::CreateCap(4096);297 auto param = gert::TilingData::CreateCap(4096);
@@ -24,8 +24,6 @@
24#include "norm/lp_norm_v2/op_kernel/arch35/lp_norm_v2_dag.h"24#include "norm/lp_norm_v2/op_kernel/arch35/lp_norm_v2_dag.h"
25#include "norm/lp_norm_v2/op_kernel/arch35/lp_norm_v2_tiling_key.h"25#include "norm/lp_norm_v2/op_kernel/arch35/lp_norm_v2_tiling_key.h"
26 26 
27-using namespace Ops::Base;
28- 
29namespace optiling27namespace optiling
30{28{
31using namespace LpNormV2;29using namespace LpNormV2;
@@ -423,7 +421,7 @@ ge::graphStatus Tiling4LpNormV2Func(gert::TilingContext* context)
423{421{
424 OP_LOGD(context->GetNodeName(), "Tiling4LpNormV2Func running begin");422 OP_LOGD(context->GetNodeName(), "Tiling4LpNormV2Func running begin");
425 423 
426- auto compileInfo = reinterpret_cast<const ReduceOpCompileInfo*>(context->GetCompileInfo());424+ auto compileInfo = context->GetCompileInfo<ReduceOpCompileInfo>();
427 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);425 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
428 426 
429 OP_LOGD(context->GetNodeName(), "Tiling4LpNormV2Func enter LpNormV2 tiling.");427 OP_LOGD(context->GetNodeName(), "Tiling4LpNormV2Func enter LpNormV2 tiling.");
@@ -27,7 +27,7 @@ namespace optiling
27{27{
28struct LpNormV2TilingKey {28struct LpNormV2TilingKey {
29 ReduceTilingKey reduceTiling;29 ReduceTilingKey reduceTiling;
30- uint32_t templateNum;30+ uint32_t templateNum = 0;
31};31};
32 32 
33class LpNormV2Tiling33class LpNormV2Tiling
@@ -57,10 +57,10 @@ private:
57 bool ChechReduceAxisIsOne();57 bool ChechReduceAxisIsOne();
58 58 
59private:59private:
60- ge::DataType xDtype_;60+ ge::DataType xDtype_ = ge::DT_FLOAT;
61- gert::TilingContext* tilingContext_;61+ gert::TilingContext* tilingContext_ = nullptr;
62 LpNormV2TilingKey key_;62 LpNormV2TilingKey key_;
63- LpNormV2TilingData* tilingData_;63+ LpNormV2TilingData* tilingData_ = nullptr;
64 float p_ = 0.0f;64 float p_ = 0.0f;
65 float recp_ = 0.0f;65 float recp_ = 0.0f;
66 float epsilon_ = 0.0f;66 float epsilon_ = 0.0f;
@@ -17,15 +17,13 @@
17 17 
18#include "atvoss/reduce/reduce_tiling_data.h"18#include "atvoss/reduce/reduce_tiling_data.h"
19 19 
20-using namespace Ops::Base;
21- 
22namespace optiling20namespace optiling
23{21{
24struct LpNormV2TilingData {22struct LpNormV2TilingData {
25- ReduceOpTilingData reduceTiling;23+ Ops::Base::ReduceOpTilingData reduceTiling;
26- float epsilon;24+ float epsilon = 0.0f;
27- float p;25+ float p = 0.0f;
28- float recp;26+ float recp = 0.0f;
29};27};
30} // namespace optiling28} // namespace optiling
31 29