已合并
fix: 规范 random 目录算子日志输出(级别错配/拼写/语法/格式符/中文残留) #5347
fix: 规范 random 目录算子日志输出(级别错配/拼写/语法/格式符/中文残留) #5347
已合并
liangtongxue创建于 9月3日
共 63 个文件变更+1607-1529
@@ -87,11 +87,11 @@ static inline bool CheckProbability(double scale)
87{87{
88 double p = ComputeProb(scale);88 double p = ComputeProb(scale);
89 if (p > 1 || p < 0) {89 if (p > 1 || p < 0) {
90- OP_LOGE(90+ OP_LOGE(ACLNN_ERR_PARAM_INVALID,
91- ACLNN_ERR_PARAM_INVALID,91+ "The value of scale is invalid, p = (scale == 0.0) ? 1 : (1 - 1 / scale) has to be between 0 and 1, "
92- "The value of scale is error, p = (scale == 0.0) ? 1 : (1 - 1 / scale) has to be between 0 and 1, but got "92+ "but got "
93- "scale %f.",93+ "scale %f.",
94- scale);94+ scale);
95 return false;95 return false;
96 }96 }
97 return true;97 return true;
@@ -27,7 +27,7 @@ const int64_t TWENTY_FIVE_NUM = 25;
27 27 
28ge::graphStatus DropOutDoMaskTilingFunc(gert::TilingContext* context)28ge::graphStatus DropOutDoMaskTilingFunc(gert::TilingContext* context)
29{29{
30- OP_LOGD(context->GetNodeName(), "DropOutDoMaskTiling running begin");30+ OP_LOGD(context->GetNodeName(), "DropOutDoMaskTiling started");
31 auto compileInfo = context->GetCompileInfo<DropOutDoMaskCompileInfo>();31 auto compileInfo = context->GetCompileInfo<DropOutDoMaskCompileInfo>();
32 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);32 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
33 return DropOutDoMaskTilingForAscendC(context);33 return DropOutDoMaskTilingForAscendC(context);
@@ -42,20 +42,19 @@ ge::graphStatus TilingPrepareDropOutDoMaskForAscendC(gert::TilingParseContext* c
42 OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo);42 OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
43 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);43 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
44 compileInfo->coreNum = ascendcPlatform.GetCoreNumAiv();44 compileInfo->coreNum = ascendcPlatform.GetCoreNumAiv();
45- OP_CHECK_IF(45+ OP_CHECK_IF((compileInfo->coreNum <= 0), OP_LOGE(context->GetNodeName(), "Failed to get core num."),
46- (compileInfo->coreNum <= 0), OP_LOGE(context->GetNodeName(), "Failed to get core num."),46+ return ge::GRAPH_FAILED);
47- return ge::GRAPH_FAILED);
48 uint64_t ubSize;47 uint64_t ubSize;
49 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);48 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
50 compileInfo->ubSize = static_cast<int64_t>(ubSize);49 compileInfo->ubSize = static_cast<int64_t>(ubSize);
51- OP_CHECK_IF(50+ OP_CHECK_IF((compileInfo->ubSize <= 0), OP_LOGE(context->GetNodeName(), "Failed to get ub size."),
52- (compileInfo->ubSize <= 0), OP_LOGE(context->GetNodeName(), "Failed to get ub size."), return ge::GRAPH_FAILED);51+ return ge::GRAPH_FAILED);
53 return ge::GRAPH_SUCCESS;52 return ge::GRAPH_SUCCESS;
54}53}
55 54 
56ge::graphStatus TilingPrepareForDropOutDoMask(gert::TilingParseContext* context)55ge::graphStatus TilingPrepareForDropOutDoMask(gert::TilingParseContext* context)
57{56{
58- OP_LOGD(context->GetNodeName(), "TilingPrepareForDropOutDoMask running begin");57+ OP_LOGD(context->GetNodeName(), "TilingPrepareForDropOutDoMask started");
59 auto compileInfo = context->GetCompiledInfo<DropOutDoMaskCompileInfo>();58 auto compileInfo = context->GetCompiledInfo<DropOutDoMaskCompileInfo>();
60 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);59 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
61 return TilingPrepareDropOutDoMaskForAscendC(context);60 return TilingPrepareDropOutDoMaskForAscendC(context);
@@ -41,10 +41,7 @@ constexpr int64_t UB_MIN_FACTOR = 2048;
41 41 
42static const std::set<ge::DataType> DROP_SUPPORTED_DTYPE = {ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16};42static const std::set<ge::DataType> DROP_SUPPORTED_DTYPE = {ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16};
43 43 
44-bool DropOutDoMaskTiling::IsCapable()44+bool DropOutDoMaskTiling::IsCapable() { return true; }
45-{
46- return true;
47-}
48 45 
49ge::graphStatus DropOutDoMaskTiling::GetPlatformInfo()46ge::graphStatus DropOutDoMaskTiling::GetPlatformInfo()
50{47{
@@ -69,7 +66,8 @@ ge::graphStatus DropOutDoMaskTiling::CheckInputShape()
69 if (keepProbAxis != 1) {66 if (keepProbAxis != 1) {
70 std::string valueStr = std::to_string(keepProbAxis);67 std::string valueStr = std::to_string(keepProbAxis);
71 std::string reasonMsg = "size of keep_prob has to be 1";68 std::string reasonMsg = "size of keep_prob has to be 1";
72- OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(context_->GetNodeName(), "input keep_prob", valueStr.c_str(), reasonMsg.c_str());69+ OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(context_->GetNodeName(), "input keep_prob", valueStr.c_str(),
70+ reasonMsg.c_str());
73 return ge::GRAPH_FAILED;71 return ge::GRAPH_FAILED;
74 }72 }
75 return ge::GRAPH_SUCCESS;73 return ge::GRAPH_SUCCESS;
@@ -82,9 +80,9 @@ ge::graphStatus DropOutDoMaskTiling::GetShapeAttrsInfo()
82 dType_ = xPtr->GetDataType();80 dType_ = xPtr->GetDataType();
83 if (DROP_SUPPORTED_DTYPE.find(dType_) == DROP_SUPPORTED_DTYPE.end()) {81 if (DROP_SUPPORTED_DTYPE.find(dType_) == DROP_SUPPORTED_DTYPE.end()) {
84 std::string valueStr = ToString(dType_);82 std::string valueStr = ToString(dType_);
85- std::string reasonMsg = "x dtype only support float32, float16, bfloat16";83+ std::string reasonMsg = "x dtype only supports float32, float16, bfloat16";
86- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(84+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor x", valueStr.c_str(),
87- context_->GetNodeName(), "input tensor x", valueStr.c_str(), reasonMsg.c_str());85+ reasonMsg.c_str());
88 return ge::GRAPH_FAILED;86 return ge::GRAPH_FAILED;
89 }87 }
90 88 
@@ -94,9 +92,9 @@ ge::graphStatus DropOutDoMaskTiling::GetShapeAttrsInfo()
94 bool dtypeInValid = (maskDtype != ge::DT_UINT8 && maskDtype != ge::DT_UINT1);92 bool dtypeInValid = (maskDtype != ge::DT_UINT8 && maskDtype != ge::DT_UINT1);
95 if (dtypeInValid) {93 if (dtypeInValid) {
96 std::string valueStr = ToString(maskDtype);94 std::string valueStr = ToString(maskDtype);
97- std::string reasonMsg = "mask dtype only support uint8, uint1 currently";95+ std::string reasonMsg = "mask dtype only supports uint8, uint1 currently";
98- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(96+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor mask", valueStr.c_str(),
99- context_->GetNodeName(), "input tensor mask", valueStr.c_str(), reasonMsg.c_str());97+ reasonMsg.c_str());
100 return ge::GRAPH_FAILED;98 return ge::GRAPH_FAILED;
101 }99 }
102 100 
@@ -106,21 +104,21 @@ ge::graphStatus DropOutDoMaskTiling::GetShapeAttrsInfo()
106 if (probPtrDtype != dType_) {104 if (probPtrDtype != dType_) {
107 std::string valueStr = ToString(dType_) + " and " + ToString(probPtrDtype);105 std::string valueStr = ToString(dType_) + " and " + ToString(probPtrDtype);
108 std::string reasonMsg = "keep_prob dtype must be equal to x dtype";106 std::string reasonMsg = "keep_prob dtype must be equal to x dtype";
109- OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(107+ OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(context_->GetNodeName(), "input keep_prob and input tensor x",
110- context_->GetNodeName(), "input keep_prob and input tensor x", valueStr.c_str(), reasonMsg.c_str());108+ valueStr.c_str(), reasonMsg.c_str());
111 return ge::GRAPH_FAILED;109 return ge::GRAPH_FAILED;
112 }110 }
113 111 
114- OP_CHECK_IF(112+ OP_CHECK_IF(CheckInputShape() != ge::GRAPH_SUCCESS, OP_LOGE(context_->GetNodeName(), "input shape check failed."),
115- CheckInputShape() != ge::GRAPH_SUCCESS, OP_LOGE(context_->GetNodeName(), "input shape check failed."),113+ return ge::GRAPH_FAILED);
116- return ge::GRAPH_FAILED);
117 return ge::GRAPH_SUCCESS;114 return ge::GRAPH_SUCCESS;
118}115}
119 116 
120ge::graphStatus DropOutDoMaskTiling::DoOpTiling()117ge::graphStatus DropOutDoMaskTiling::DoOpTiling()
121{118{
122 typeSize_ = ge::GetSizeByDataType(dType_);119 typeSize_ = ge::GetSizeByDataType(dType_);
123- OP_CHECK_IF(typeSize_ <= 0, OP_LOGE(context_->GetNodeName(), "get dataType size fail."), return ge::GRAPH_FAILED);120+ OP_CHECK_IF(typeSize_ <= 0, OP_LOGE(context_->GetNodeName(), "Failed to get dataType size."),
121+ return ge::GRAPH_FAILED);
124 // total: ub/db122 // total: ub/db
125 // used: x*typesize (x), x/8 (mask), x*typesize(out)123 // used: x*typesize (x), x/8 (mask), x*typesize(out)
126 int64_t ubBlock = GetUbBlockSize(context_);124 int64_t ubBlock = GetUbBlockSize(context_);
@@ -149,10 +147,7 @@ ge::graphStatus DropOutDoMaskTiling::DoOpTiling()
149 return ge::GRAPH_SUCCESS;147 return ge::GRAPH_SUCCESS;
150}148}
151 149 
152-ge::graphStatus DropOutDoMaskTiling::DoLibApiTiling()150+ge::graphStatus DropOutDoMaskTiling::DoLibApiTiling() { return ge::GRAPH_SUCCESS; }
153-{
154- return ge::GRAPH_SUCCESS;
155-}
156 151 
157uint64_t DropOutDoMaskTiling::GetTilingKey() const152uint64_t DropOutDoMaskTiling::GetTilingKey() const
158{153{
@@ -160,10 +155,7 @@ uint64_t DropOutDoMaskTiling::GetTilingKey() const
160 return tilingKey;155 return tilingKey;
161}156}
162 157 
163-ge::graphStatus DropOutDoMaskTiling::GetWorkspaceSize()158+ge::graphStatus DropOutDoMaskTiling::GetWorkspaceSize() { return ge::GRAPH_SUCCESS; }
164-{
165- return ge::GRAPH_SUCCESS;
166-}
167 159 
168ge::graphStatus DropOutDoMaskTiling::PostTiling()160ge::graphStatus DropOutDoMaskTiling::PostTiling()
169{161{
@@ -42,27 +42,27 @@ using std::map;
42using std::string;42using std::string;
43using std::vector;43using std::vector;
44 44 
45-#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape) \45+#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape) \
46- do { \46+ do { \
47- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \47+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
48- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \48+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
49- TensorDesc placeholder##intputIndex##_desc = \49+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), \
50- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \50+ FORMAT_ND, intputDtype); \
51- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \51+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
52- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \52+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
53- Tensor tensor_placeholder##intputIndex; \53+ Tensor tensor_placeholder##intputIndex; \
54- ret = GenOnesDataFloat32( \54+ ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
55- placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, placeholder##intputIndex##_desc, 2); \55+ placeholder##intputIndex##_desc, 2); \
56- if (ret != SUCCESS) { \56+ if (ret != SUCCESS) { \
57- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \57+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
58- return FAILED; \58+ return FAILED; \
59- } \59+ } \
60- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \60+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
61- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \61+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
62- input.push_back(tensor_placeholder##intputIndex); \62+ input.push_back(tensor_placeholder##intputIndex); \
63- graph.AddOp(placeholder##intputIndex); \63+ graph.AddOp(placeholder##intputIndex); \
64- dropoutdomaskv3.set_input_##intputName(placeholder##intputIndex); \64+ dropoutdomaskv3.set_input_##intputName(placeholder##intputIndex); \
65- inputs.push_back(placeholder##intputIndex); \65+ inputs.push_back(placeholder##intputIndex); \
66 } while (0)66 } while (0)
67 67 
68#define LOG_PRINT(message, ...) \68#define LOG_PRINT(message, ...) \
@@ -130,8 +130,8 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorD
130 return SUCCESS;130 return SUCCESS;
131}131}
132 132 
133-int32_t GenOnesData(133+int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type,
134- vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type, int value)134+ int value)
135{135{
136 input_tensor_desc.SetRealDimCnt(shapes.size());136 input_tensor_desc.SetRealDimCnt(shapes.size());
137 size_t size = 1;137 size_t size = 1;
@@ -139,11 +139,11 @@ int32_t GenOnesData(
139 size *= shapes[i];139 size *= shapes[i];
140 }140 }
141 uint32_t data_len = size * GetDataTypeSize(data_type);141 uint32_t data_len = size * GetDataTypeSize(data_type);
142- int64_t *pData = new (std::nothrow) int64_t[size];142+ int64_t* pData = new (std::nothrow) int64_t[size];
143 for (uint32_t i = 0; i < size; ++i) {143 for (uint32_t i = 0; i < size; ++i) {
144 pData[i] = static_cast<int64_t>(value);144 pData[i] = static_cast<int64_t>(value);
145 }145 }
146- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);146+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
147 delete[] pData;147 delete[] pData;
148 return SUCCESS;148 return SUCCESS;
149}149}
@@ -156,9 +156,8 @@ int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
156 return SUCCESS;156 return SUCCESS;
157}157}
158 158 
159-int CreateOppInGraph(159+int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs,
160- DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs, std::vector<Operator>& outputs,160+ std::vector<Operator>& outputs, Graph& graph)
161- Graph& graph)
162{161{
163 Status ret = SUCCESS;162 Status ret = SUCCESS;
164 // 自定义代码:添加单算子定义到图中163 // 自定义代码:添加单算子定义到图中
@@ -172,7 +171,7 @@ int CreateOppInGraph(
172 TensorDesc desc2(ge::Shape(maskShape), FORMAT_ND, ge::DT_UINT8);171 TensorDesc desc2(ge::Shape(maskShape), FORMAT_ND, ge::DT_UINT8);
173 desc2.SetPlacement(ge::kPlacementHost);172 desc2.SetPlacement(ge::kPlacementHost);
174 desc2.SetFormat(FORMAT_ND);173 desc2.SetFormat(FORMAT_ND);
175- uint8_t *mask_data = new (std::nothrow) uint8_t[128];174+ uint8_t* mask_data = new (std::nothrow) uint8_t[128];
176 memset(mask_data, 1, 128);175 memset(mask_data, 1, 128);
177 Tensor tensor2(desc2, mask_data, 128);176 Tensor tensor2(desc2, mask_data, 128);
178 delete[] mask_data;177 delete[] mask_data;
@@ -199,7 +198,7 @@ void SaveInputOutput(std::vector<ge::Tensor>& input, std::vector<ge::Tensor>& ou
199 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";198 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
200 uint8_t* input_data_i = input[i].GetData();199 uint8_t* input_data_i = input[i].GetData();
201 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();200 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
202- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;201+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
203 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());202 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
204 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);203 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
205 }204 }
@@ -210,7 +209,7 @@ void SaveInputOutput(std::vector<ge::Tensor>& input, std::vector<ge::Tensor>& ou
210 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";209 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
211 uint8_t* output_data_i = output[i].GetData();210 uint8_t* output_data_i = output[i].GetData();
212 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();211 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
213- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;212+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
214 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());213 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
215 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);214 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
216 float* resultData = (float*)output_data_i;215 float* resultData = (float*)output_data_i;
@@ -230,7 +229,7 @@ int main(int argc, char* argv[])
230 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};229 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
231 Status ret = ge::GEInitialize(global_options);230 Status ret = ge::GEInitialize(global_options);
232 if (ret != SUCCESS) {231 if (ret != SUCCESS) {
233- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());232+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
234 return FAILED;233 return FAILED;
235 }234 }
236 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());235 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -281,7 +280,7 @@ int main(int argc, char* argv[])
281 std::vector<ge::Tensor> output;280 std::vector<ge::Tensor> output;
282 ret = session->RunGraph(graph_id, input, output);281 ret = session->RunGraph(graph_id, input, output);
283 if (ret != SUCCESS) {282 if (ret != SUCCESS) {
284- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());283+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
285 delete session;284 delete session;
286 GEFinalize();285 GEFinalize();
287 return FAILED;286 return FAILED;
@@ -299,7 +298,7 @@ int main(int argc, char* argv[])
299 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());298 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
300 ret = ge::GEFinalize();299 ret = ge::GEFinalize();
301 if (ret != SUCCESS) {300 if (ret != SUCCESS) {
302- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());301+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
303 return FAILED;302 return FAILED;
304 }303 }
305 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());304 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -44,51 +44,51 @@ using std::string;
44using std::vector;44using std::vector;
45 45 
46#undef ADD_INPUT46#undef ADD_INPUT
47-#define ADD_INPUT(intputIndex, opInstance, intputName, intputDtype, inputShape) \47+#define ADD_INPUT(intputIndex, opInstance, intputName, intputDtype, inputShape) \
48- do { \
49- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
50- auto placeholder##intputIndex = op::Data(("placeholder" + std::to_string(intputIndex)).c_str()).set_attr_index(intputIndex-1); \
51- TensorDesc placeholder##intputIndex##_desc = \
52- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \
53- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
54- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
55- Tensor tensor_placeholder##intputIndex; \
56- ret = GenOnesDataFloat32( \
57- placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, placeholder##intputIndex##_desc, 2); \
58- if (ret != SUCCESS) { \
59- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
60- return FAILED; \
61- } \
62- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
63- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
64- input.push_back(tensor_placeholder##intputIndex); \
65- graph.AddOp(placeholder##intputIndex); \
66- opInstance.set_input_##intputName(placeholder##intputIndex); \
67- inputs.push_back(placeholder##intputIndex); \
68- } while (0)
69- 
70-#define ADD_CONST_INPUT(inputIndex, inputName, inputDtype, inputShape, value) \
71 do { \48 do { \
72- vector<int64_t> placeholder##inputIndex##_shape = inputShape; \49+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
73- auto placeholder##inputIndex = op::Const("placeholder" + inputIndex); \50+ auto placeholder##intputIndex = op::Data(("placeholder" + std::to_string(intputIndex)).c_str()) \
74- TensorDesc placeholder##inputIndex##_desc = \51+ .set_attr_index(intputIndex - 1); \
75- TensorDesc(ge::Shape(placeholder##inputIndex##_shape), FORMAT_ND, inputDtype); \52+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), \
76- placeholder##inputIndex##_desc.SetPlacement(ge::kPlacementHost); \53+ FORMAT_ND, intputDtype); \
77- placeholder##inputIndex##_desc.SetFormat(FORMAT_ND); \54+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
78- Tensor tensor_placeholder##inputIndex; \55+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
79- ret = GenOnesData( \56+ Tensor tensor_placeholder##intputIndex; \
80- placeholder##inputIndex##_shape, tensor_placeholder##inputIndex, placeholder##inputIndex##_desc, \57+ ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
81- inputDtype, value); \58+ placeholder##intputIndex##_desc, 2); \
82 if (ret != SUCCESS) { \59 if (ret != SUCCESS) { \
83- printf("%s - ERROR - [XIR]: Generate const input data failed\n", GetTime().c_str()); \60+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
84 return FAILED; \61 return FAILED; \
85 } \62 } \
86- placeholder##inputIndex.SetAttr("value", tensor_placeholder##inputIndex); \63+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
87- placeholder##inputIndex.update_output_desc_y(placeholder##inputIndex##_desc); \64+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
88- graph.AddOp(placeholder##inputIndex); \65+ input.push_back(tensor_placeholder##intputIndex); \
89- add1.set_input_##inputName(placeholder##inputIndex); \66+ graph.AddOp(placeholder##intputIndex); \
90- add1.update_input_desc_##inputName(placeholder##inputIndex##_desc); \67+ opInstance.set_input_##intputName(placeholder##intputIndex); \
91- inputs.push_back(placeholder##inputIndex); \68+ inputs.push_back(placeholder##intputIndex); \
69+ } while (0)
70+ 
71+#define ADD_CONST_INPUT(inputIndex, inputName, inputDtype, inputShape, value) \
72+ do { \
73+ vector<int64_t> placeholder##inputIndex##_shape = inputShape; \
74+ auto placeholder##inputIndex = op::Const("placeholder" + inputIndex); \
75+ TensorDesc placeholder##inputIndex##_desc = TensorDesc(ge::Shape(placeholder##inputIndex##_shape), FORMAT_ND, \
76+ inputDtype); \
77+ placeholder##inputIndex##_desc.SetPlacement(ge::kPlacementHost); \
78+ placeholder##inputIndex##_desc.SetFormat(FORMAT_ND); \
79+ Tensor tensor_placeholder##inputIndex; \
80+ ret = GenOnesData(placeholder##inputIndex##_shape, tensor_placeholder##inputIndex, \
81+ placeholder##inputIndex##_desc, inputDtype, value); \
82+ if (ret != SUCCESS) { \
83+ printf("%s - ERROR - [XIR]: Generate const input data failed\n", GetTime().c_str()); \
84+ return FAILED; \
85+ } \
86+ placeholder##inputIndex.SetAttr("value", tensor_placeholder##inputIndex); \
87+ placeholder##inputIndex.update_output_desc_y(placeholder##inputIndex##_desc); \
88+ graph.AddOp(placeholder##inputIndex); \
89+ add1.set_input_##inputName(placeholder##inputIndex); \
90+ add1.update_input_desc_##inputName(placeholder##inputIndex##_desc); \
91+ inputs.push_back(placeholder##inputIndex); \
92 } while (0)92 } while (0)
93 93 
94#define LOG_PRINT(message, ...) \94#define LOG_PRINT(message, ...) \
@@ -155,8 +155,8 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorD
155 return SUCCESS;155 return SUCCESS;
156}156}
157 157 
158-int32_t GenOnesData(158+int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type,
159- vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type, int value)159+ int value)
160{160{
161 input_tensor_desc.SetRealDimCnt(shapes.size());161 input_tensor_desc.SetRealDimCnt(shapes.size());
162 size_t size = 1;162 size_t size = 1;
@@ -184,19 +184,18 @@ int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
184#define ADD_INPUT_ATTR(opInstance, attrName, attrValue) opInstance.set_attr_##attrName(attrValue)184#define ADD_INPUT_ATTR(opInstance, attrName, attrValue) opInstance.set_attr_##attrName(attrValue)
185 185 
186// --- 2. 修改业务函数 ---186// --- 2. 修改业务函数 ---
187-int CreateOppInGraph(187+int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs,
188- DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs, std::vector<Operator>& outputs,188+ std::vector<Operator>& outputs, Graph& graph)
189- Graph& graph)
190{189{
191 Status ret = SUCCESS;190 Status ret = SUCCESS;
192- 191+ 
193 // 定义算子实例192 // 定义算子实例
194 auto dropoutdomaskv3d = op::DropOutDoMaskV3D("dropoutdomaskv3d");193 auto dropoutdomaskv3d = op::DropOutDoMaskV3D("dropoutdomaskv3d");
195- 194+ 
196 std::vector<int64_t> xShape = {32};195 std::vector<int64_t> xShape = {32};
197 std::vector<int64_t> maskShape = {128};196 std::vector<int64_t> maskShape = {128};
198 197 
199- ADD_INPUT(1, dropoutdomaskv3d, x, inDtype, xShape); 198+ ADD_INPUT(1, dropoutdomaskv3d, x, inDtype, xShape);
200 ADD_INPUT(2, dropoutdomaskv3d, mask, ge::DT_UINT8, maskShape);199 ADD_INPUT(2, dropoutdomaskv3d, mask, ge::DT_UINT8, maskShape);
201 200 
202 ADD_INPUT_ATTR(dropoutdomaskv3d, keep_prob, 0.56);201 ADD_INPUT_ATTR(dropoutdomaskv3d, keep_prob, 0.56);
@@ -212,7 +211,7 @@ void SaveInputOutput(std::vector<ge::Tensor>& input, std::vector<ge::Tensor>& ou
212 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";211 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
213 uint8_t* input_data_i = input[i].GetData();212 uint8_t* input_data_i = input[i].GetData();
214 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();213 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
215- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;214+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
216 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());215 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
217 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);216 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
218 }217 }
@@ -223,7 +222,7 @@ void SaveInputOutput(std::vector<ge::Tensor>& input, std::vector<ge::Tensor>& ou
223 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";222 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
224 uint8_t* output_data_i = output[i].GetData();223 uint8_t* output_data_i = output[i].GetData();
225 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();224 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
226- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;225+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
227 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());226 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
228 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);227 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
229 float* resultData = (float*)output_data_i;228 float* resultData = (float*)output_data_i;
@@ -243,7 +242,7 @@ int main(int argc, char* argv[])
243 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};242 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
244 Status ret = ge::GEInitialize(global_options);243 Status ret = ge::GEInitialize(global_options);
245 if (ret != SUCCESS) {244 if (ret != SUCCESS) {
246- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());245+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
247 return FAILED;246 return FAILED;
248 }247 }
249 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());248 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -294,7 +293,7 @@ int main(int argc, char* argv[])
294 std::vector<ge::Tensor> output;293 std::vector<ge::Tensor> output;
295 ret = session->RunGraph(graph_id, input, output);294 ret = session->RunGraph(graph_id, input, output);
296 if (ret != SUCCESS) {295 if (ret != SUCCESS) {
297- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());296+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
298 delete session;297 delete session;
299 GEFinalize();298 GEFinalize();
300 return FAILED;299 return FAILED;
@@ -312,9 +311,9 @@ int main(int argc, char* argv[])
312 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());311 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
313 ret = ge::GEFinalize();312 ret = ge::GEFinalize();
314 if (ret != SUCCESS) {313 if (ret != SUCCESS) {
315- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());314+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
316 return FAILED;315 return FAILED;
317 }316 }
318 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());317 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
319 return SUCCESS;318 return SUCCESS;
320-}319+}
@@ -43,27 +43,27 @@ using std::map;
43using std::string;43using std::string;
44using std::vector;44using std::vector;
45 45 
46-#define ADD_INPUT(inputIndex, inputName, inputDtype, inputShape, value) \46+#define ADD_INPUT(inputIndex, inputName, inputDtype, inputShape, value) \
47- do { \47+ do { \
48- vector<int64_t> placeholder##inputIndex##_shape = inputShape; \48+ vector<int64_t> placeholder##inputIndex##_shape = inputShape; \
49- auto placeholder##inputIndex = op::Data("placeholder" + inputIndex).set_attr_index(0); \49+ auto placeholder##inputIndex = op::Data("placeholder" + inputIndex).set_attr_index(0); \
50- TensorDesc placeholder##inputIndex##_desc = \50+ TensorDesc placeholder##inputIndex##_desc = TensorDesc(ge::Shape(placeholder##inputIndex##_shape), FORMAT_ND, \
51- TensorDesc(ge::Shape(placeholder##inputIndex##_shape), FORMAT_ND, inputDtype); \51+ inputDtype); \
52- placeholder##inputIndex##_desc.SetPlacement(ge::kPlacementHost); \52+ placeholder##inputIndex##_desc.SetPlacement(ge::kPlacementHost); \
53- placeholder##inputIndex##_desc.SetFormat(FORMAT_ND); \53+ placeholder##inputIndex##_desc.SetFormat(FORMAT_ND); \
54- Tensor tensor_placeholder##inputIndex; \54+ Tensor tensor_placeholder##inputIndex; \
55- ret = GenOnesData<decltype(value)>( \55+ ret = GenOnesData<decltype(value)>(placeholder##inputIndex##_shape, tensor_placeholder##inputIndex, \
56- placeholder##inputIndex##_shape, tensor_placeholder##inputIndex, placeholder##inputIndex##_desc, value); \56+ placeholder##inputIndex##_desc, value); \
57- if (ret != SUCCESS) { \57+ if (ret != SUCCESS) { \
58- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \58+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
59- return FAILED; \59+ return FAILED; \
60- } \60+ } \
61- placeholder##inputIndex.update_input_desc_x(placeholder##inputIndex##_desc); \61+ placeholder##inputIndex.update_input_desc_x(placeholder##inputIndex##_desc); \
62- placeholder##inputIndex.update_output_desc_y(placeholder##inputIndex##_desc); \62+ placeholder##inputIndex.update_output_desc_y(placeholder##inputIndex##_desc); \
63- input.push_back(tensor_placeholder##inputIndex); \63+ input.push_back(tensor_placeholder##inputIndex); \
64- graph.AddOp(placeholder##inputIndex); \64+ graph.AddOp(placeholder##inputIndex); \
65- dropoutv3.set_input_##inputName(placeholder##inputIndex); \65+ dropoutv3.set_input_##inputName(placeholder##inputIndex); \
66- inputs.push_back(placeholder##inputIndex); \66+ inputs.push_back(placeholder##inputIndex); \
67 } while (0)67 } while (0)
68 68 
69#define LOG_PRINT(message, ...) \69#define LOG_PRINT(message, ...) \
@@ -112,7 +112,7 @@ uint32_t GetDataTypeSize(DataType dt)
112 return dilation;112 return dilation;
113}113}
114 114 
115-template<typename T>115+template <typename T>
116int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, T value)116int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, T value)
117{117{
118 input_tensor_desc.SetRealDimCnt(shapes.size());118 input_tensor_desc.SetRealDimCnt(shapes.size());
@@ -138,8 +138,8 @@ int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
138 return SUCCESS;138 return SUCCESS;
139}139}
140 140 
141-int CreateOppInGraph(141+int CreateOppInGraph(std::vector<ge::Tensor>& input, std::vector<Operator>& inputs, std::vector<Operator>& outputs,
142- std::vector<ge::Tensor>& input, std::vector<Operator>& inputs, std::vector<Operator>& outputs, Graph& graph)142+ Graph& graph)
143{143{
144 Status ret = SUCCESS;144 Status ret = SUCCESS;
145 auto dropoutv3 = op::DropOutV3("dropoutv3");145 auto dropoutv3 = op::DropOutV3("dropoutv3");
@@ -163,7 +163,7 @@ void SaveInputOutput(std::vector<ge::Tensor>& input, std::vector<ge::Tensor>& ou
163 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";163 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
164 uint8_t* input_data_i = input[i].GetData();164 uint8_t* input_data_i = input[i].GetData();
165 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();165 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
166- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;166+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
167 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());167 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
168 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);168 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
169 }169 }
@@ -174,7 +174,7 @@ void SaveInputOutput(std::vector<ge::Tensor>& input, std::vector<ge::Tensor>& ou
174 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";174 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
175 uint8_t* output_data_i = output[i].GetData();175 uint8_t* output_data_i = output[i].GetData();
176 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();176 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
177- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;177+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
178 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());178 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
179 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);179 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
180 float* resultData = (float*)output_data_i;180 float* resultData = (float*)output_data_i;
@@ -194,7 +194,7 @@ int main(int argc, char* argv[])
194 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};194 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
195 Status ret = ge::GEInitialize(global_options);195 Status ret = ge::GEInitialize(global_options);
196 if (ret != SUCCESS) {196 if (ret != SUCCESS) {
197- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());197+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
198 return FAILED;198 return FAILED;
199 }199 }
200 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());200 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -241,7 +241,7 @@ int main(int argc, char* argv[])
241 std::vector<ge::Tensor> output;241 std::vector<ge::Tensor> output;
242 ret = session->RunGraph(graph_id, input, output);242 ret = session->RunGraph(graph_id, input, output);
243 if (ret != SUCCESS) {243 if (ret != SUCCESS) {
244- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());244+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
245 delete session;245 delete session;
246 GEFinalize();246 GEFinalize();
247 return FAILED;247 return FAILED;
@@ -259,7 +259,7 @@ int main(int argc, char* argv[])
259 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());259 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
260 ret = ge::GEFinalize();260 ret = ge::GEFinalize();
261 if (ret != SUCCESS) {261 if (ret != SUCCESS) {
262- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());262+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
263 return FAILED;263 return FAILED;
264 }264 }
265 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());265 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -61,7 +61,8 @@ static inline bool CheckNotNull(const aclTensor* input, const aclTensor* out, co
61static inline bool CheckIsNullptr(const aclTensor* optionalNoiseShape)61static inline bool CheckIsNullptr(const aclTensor* optionalNoiseShape)
62{62{
63 if (optionalNoiseShape != nullptr) {63 if (optionalNoiseShape != nullptr) {
64- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "currently, the input of noise_shape must be nullptr, please check.");64+ OP_LOGE(ACLNN_ERR_PARAM_INVALID,
65+ "currently, the input of noise_shape must be nullptr, but got a non-null tensor, please check.");
65 return false;66 return false;
66 }67 }
67 return true;68 return true;
@@ -243,7 +243,7 @@ bool DropOutV3FusionPass::MeetRequirements(const std::unique_ptr<MatchResult>& m
243 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);243 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);
244 }244 }
245 if (version < GE_COMPILER_VERSION_900) {245 if (version < GE_COMPILER_VERSION_900) {
246- OP_LOGD(kPassName.c_str(), "GE runtime version %d < 90000000, skip pass.", version);246+ OP_LOGD(kPassName.c_str(), "GE runtime version %d < 9.0.0, skip pass.", version);
247 return false;247 return false;
248 }248 }
249 249 
@@ -275,35 +275,35 @@ bool DropOutV3FusionPass::CheckGenMaskNode(const std::unique_ptr<MatchResult>& m
275 }275 }
276 276 
277 if (genMaskIo.node.GetInputsSize() != kGenMaskInputCount) {277 if (genMaskIo.node.GetInputsSize() != kGenMaskInputCount) {
278- OP_LOGE(kPassName.c_str(), "GenMask input size != 5");278+ OP_LOGE(kPassName.c_str(), "GenMask input size is %zu, expected 5", genMaskIo.node.GetInputsSize());
279 return false;279 return false;
280 }280 }
281 281 
282 TensorDesc probDesc;282 TensorDesc probDesc;
283 genMaskIo.node.GetInputDesc(kGenMaskIdxProb, probDesc);283 genMaskIo.node.GetInputDesc(kGenMaskIdxProb, probDesc);
284 if (!CheckDtype(probDesc.GetDataType(), {DT_FLOAT, DT_FLOAT16, DT_BF16})) {284 if (!CheckDtype(probDesc.GetDataType(), {DT_FLOAT, DT_FLOAT16, DT_BF16})) {
285- OP_LOGE(kPassName.c_str(), "GenMask prob dtype not supported");285+ OP_LOGE(kPassName.c_str(), "GenMask prob dtype %d not supported", static_cast<int>(probDesc.GetDataType()));
286 return false;286 return false;
287 }287 }
288 288 
289 TensorDesc seedDesc;289 TensorDesc seedDesc;
290 genMaskIo.node.GetInputDesc(kGenMaskIdxSeed, seedDesc);290 genMaskIo.node.GetInputDesc(kGenMaskIdxSeed, seedDesc);
291 if (!CheckDtype(seedDesc.GetDataType(), {DT_INT32, DT_INT64})) {291 if (!CheckDtype(seedDesc.GetDataType(), {DT_INT32, DT_INT64})) {
292- OP_LOGE(kPassName.c_str(), "GenMask seed dtype not supported");292+ OP_LOGE(kPassName.c_str(), "GenMask seed dtype %d not supported", static_cast<int>(seedDesc.GetDataType()));
293 return false;293 return false;
294 }294 }
295 295 
296 TensorDesc offsetDesc;296 TensorDesc offsetDesc;
297 genMaskIo.node.GetInputDesc(kGenMaskIdxOffset, offsetDesc);297 genMaskIo.node.GetInputDesc(kGenMaskIdxOffset, offsetDesc);
298 if (offsetDesc.GetDataType() != DT_INT64) {298 if (offsetDesc.GetDataType() != DT_INT64) {
299- OP_LOGE(kPassName.c_str(), "GenMask offset dtype != DT_INT64");299+ OP_LOGE(kPassName.c_str(), "GenMask offset dtype %d != DT_INT64", static_cast<int>(offsetDesc.GetDataType()));
300 return false;300 return false;
301 }301 }
302 302 
303 TensorDesc outputDesc;303 TensorDesc outputDesc;
304 genMaskIo.node.GetOutputDesc(0, outputDesc);304 genMaskIo.node.GetOutputDesc(0, outputDesc);
305 if (outputDesc.GetDataType() != DT_UINT8) {305 if (outputDesc.GetDataType() != DT_UINT8) {
306- OP_LOGE(kPassName.c_str(), "GenMask output dtype != DT_UINT8");306+ OP_LOGE(kPassName.c_str(), "GenMask output dtype %d != DT_UINT8", static_cast<int>(outputDesc.GetDataType()));
307 return false;307 return false;
308 }308 }
309 return true;309 return true;
@@ -325,21 +325,21 @@ bool DropOutV3FusionPass::CheckDoMaskNode(const std::unique_ptr<MatchResult>& ma
325 }325 }
326 326 
327 if (doMaskIo.node.GetInputsSize() != kDoMaskInputCount) {327 if (doMaskIo.node.GetInputsSize() != kDoMaskInputCount) {
328- OP_LOGE(kPassName.c_str(), "DoMask input size != 3");328+ OP_LOGE(kPassName.c_str(), "DoMask input size is %zu, expected 3", doMaskIo.node.GetInputsSize());
329 return false;329 return false;
330 }330 }
331 331 
332 TensorDesc inputDesc;332 TensorDesc inputDesc;
333 doMaskIo.node.GetInputDesc(0, inputDesc);333 doMaskIo.node.GetInputDesc(0, inputDesc);
334 if (!CheckDtype(inputDesc.GetDataType(), {DT_FLOAT, DT_FLOAT16, DT_BF16})) {334 if (!CheckDtype(inputDesc.GetDataType(), {DT_FLOAT, DT_FLOAT16, DT_BF16})) {
335- OP_LOGE(kPassName.c_str(), "DoMask x dtype not supported");335+ OP_LOGE(kPassName.c_str(), "DoMask x dtype %d not supported", static_cast<int>(inputDesc.GetDataType()));
336 return false;336 return false;
337 }337 }
338 338 
339 TensorDesc outputDesc;339 TensorDesc outputDesc;
340 doMaskIo.node.GetOutputDesc(0, outputDesc);340 doMaskIo.node.GetOutputDesc(0, outputDesc);
341 if (!CheckDtype(outputDesc.GetDataType(), {DT_FLOAT, DT_FLOAT16, DT_BF16})) {341 if (!CheckDtype(outputDesc.GetDataType(), {DT_FLOAT, DT_FLOAT16, DT_BF16})) {
342- OP_LOGE(kPassName.c_str(), "DoMask y dtype not supported");342+ OP_LOGE(kPassName.c_str(), "DoMask y dtype %d not supported", static_cast<int>(outputDesc.GetDataType()));
343 return false;343 return false;
344 }344 }
345 return true;345 return true;
@@ -363,7 +363,7 @@ GraphUniqPtr DropOutV3FusionPass::Replacement(const std::unique_ptr<MatchResult>
363 363 
364 GraphUniqPtr replaceGraph = builder.BuildAndReset({output.y});364 GraphUniqPtr replaceGraph = builder.BuildAndReset({output.y});
365 if (InferShape(replaceGraph, subgraphInputs) != SUCCESS) {365 if (InferShape(replaceGraph, subgraphInputs) != SUCCESS) {
366- OP_LOGE(kPassName.c_str(), "Infershape failed.");366+ OP_LOGE(kPassName.c_str(), "InferShape failed.");
367 return nullptr;367 return nullptr;
368 }368 }
369 return replaceGraph;369 return replaceGraph;
@@ -112,7 +112,7 @@ bool DropOutV3SplitFusionPass::CheckDtypes(const GNode& node) const
112 node.GetInputDesc(kIdxX, xDesc);112 node.GetInputDesc(kIdxX, xDesc);
113 DataType xDtype = xDesc.GetDataType();113 DataType xDtype = xDesc.GetDataType();
114 if (!CheckDtype(xDtype, {DT_FLOAT, DT_FLOAT16, DT_BF16})) {114 if (!CheckDtype(xDtype, {DT_FLOAT, DT_FLOAT16, DT_BF16})) {
115- OP_LOGE(kPassName.c_str(), "x dtype only support float/float16/bf16, actual: %d", static_cast<int>(xDtype));115+ OP_LOGE(kPassName.c_str(), "x dtype only supports float/float16/bf16, actual: %d", static_cast<int>(xDtype));
116 return false;116 return false;
117 }117 }
118 118 
@@ -120,7 +120,7 @@ bool DropOutV3SplitFusionPass::CheckDtypes(const GNode& node) const
120 node.GetInputDesc(kIdxP, pDesc);120 node.GetInputDesc(kIdxP, pDesc);
121 DataType pDtype = pDesc.GetDataType();121 DataType pDtype = pDesc.GetDataType();
122 if (!CheckDtype(pDtype, {DT_FLOAT, DT_FLOAT16, DT_BF16})) {122 if (!CheckDtype(pDtype, {DT_FLOAT, DT_FLOAT16, DT_BF16})) {
123- OP_LOGE(kPassName.c_str(), "p dtype only support float/float16/bf16, actual: %d", static_cast<int>(pDtype));123+ OP_LOGE(kPassName.c_str(), "p dtype only supports float/float16/bf16, actual: %d", static_cast<int>(pDtype));
124 return false;124 return false;
125 }125 }
126 126 
@@ -128,7 +128,7 @@ bool DropOutV3SplitFusionPass::CheckDtypes(const GNode& node) const
128 node.GetInputDesc(kIdxSeed, seedDesc);128 node.GetInputDesc(kIdxSeed, seedDesc);
129 DataType seedDtype = seedDesc.GetDataType();129 DataType seedDtype = seedDesc.GetDataType();
130 if (!CheckDtype(seedDtype, {DT_INT32, DT_INT64})) {130 if (!CheckDtype(seedDtype, {DT_INT32, DT_INT64})) {
131- OP_LOGE(kPassName.c_str(), "seed dtype only support int32/int64, actual: %d", static_cast<int>(seedDtype));131+ OP_LOGE(kPassName.c_str(), "seed dtype only supports int32/int64, actual: %d", static_cast<int>(seedDtype));
132 return false;132 return false;
133 }133 }
134 134 
@@ -136,12 +136,12 @@ bool DropOutV3SplitFusionPass::CheckDtypes(const GNode& node) const
136 node.GetOutputDesc(0, yDesc);136 node.GetOutputDesc(0, yDesc);
137 DataType yDtype = yDesc.GetDataType();137 DataType yDtype = yDesc.GetDataType();
138 if (!CheckDtype(yDtype, {DT_FLOAT, DT_FLOAT16, DT_BF16})) {138 if (!CheckDtype(yDtype, {DT_FLOAT, DT_FLOAT16, DT_BF16})) {
139- OP_LOGE(kPassName.c_str(), "y dtype only support float/float16/bf16, actual: %d", static_cast<int>(yDtype));139+ OP_LOGE(kPassName.c_str(), "y dtype only supports float/float16/bf16, actual: %d", static_cast<int>(yDtype));
140 return false;140 return false;
141 }141 }
142 142 
143 if (xDtype != yDtype) {143 if (xDtype != yDtype) {
144- OP_LOGE(kPassName.c_str(), "x dtype should same with y dtype, x: %d, y: %d", static_cast<int>(xDtype),144+ OP_LOGE(kPassName.c_str(), "x dtype should be the same as y dtype, x: %d, y: %d", static_cast<int>(xDtype),
145 static_cast<int>(yDtype));145 static_cast<int>(yDtype));
146 return false;146 return false;
147 }147 }
@@ -316,7 +316,7 @@ Status DropOutV3SplitFusionPass::Run(GraphPtr& graph, [[maybe_unused]] CustomPas
316 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);316 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);
317 }317 }
318 if (version < GE_COMPILER_VERSION_900) {318 if (version < GE_COMPILER_VERSION_900) {
319- OP_LOGD(kPassName.c_str(), "GE runtime version %d < 90000000, skip pass.", version);319+ OP_LOGD(kPassName.c_str(), "GE runtime version %d < 9.0.0, skip pass.", version);
320 return GRAPH_NOT_CHANGED;320 return GRAPH_NOT_CHANGED;
321 }321 }
322 322 
@@ -78,12 +78,14 @@ OpTilingConfig DropOutV3Tiling::BuildOpConfig(gert::TilingContext* context)
78 config.getSeedAndOffset = [](gert::TilingContext* ctx, int64_t& seed, int64_t& offset) {78 config.getSeedAndOffset = [](gert::TilingContext* ctx, int64_t& seed, int64_t& offset) {
79 gert::Shape seedShape;79 gert::Shape seedShape;
80 auto ret = ExtractTensorValue(ctx, INPUT_IDX_SEED, seedShape);80 auto ret = ExtractTensorValue(ctx, INPUT_IDX_SEED, seedShape);
81- OP_CHECK_IF(ret != ge::GRAPH_SUCCESS, OP_LOGE(ctx->GetNodeName(), "get seed value failed"),81+ OP_CHECK_IF(ret != ge::GRAPH_SUCCESS,
82+ OP_LOGE(ctx->GetNodeName(), "get seed value failed, ret = %d", static_cast<int32_t>(ret)),
82 return ge::GRAPH_FAILED);83 return ge::GRAPH_FAILED);
83 seed = static_cast<int64_t>(seedShape.GetDim(0));84 seed = static_cast<int64_t>(seedShape.GetDim(0));
84 gert::Shape offsetShape;85 gert::Shape offsetShape;
85 ret = ExtractTensorValue(ctx, INPUT_IDX_OFFSET, offsetShape);86 ret = ExtractTensorValue(ctx, INPUT_IDX_OFFSET, offsetShape);
86- OP_CHECK_IF(ret != ge::GRAPH_SUCCESS, OP_LOGE(ctx->GetNodeName(), "get offset value failed"),87+ OP_CHECK_IF(ret != ge::GRAPH_SUCCESS,
88+ OP_LOGE(ctx->GetNodeName(), "get offset value failed, ret = %d", static_cast<int32_t>(ret)),
87 return ge::GRAPH_FAILED);89 return ge::GRAPH_FAILED);
88 offset = static_cast<int64_t>(offsetShape.GetDim(1));90 offset = static_cast<int64_t>(offsetShape.GetDim(1));
89 if (offset % OFFSET_LIMIT != 0) {91 if (offset % OFFSET_LIMIT != 0) {
@@ -144,7 +146,7 @@ ge::graphStatus DropOutV3Tiling::UniqueProcess()
144 }146 }
145 default: {147 default: {
146 std::string valueStr = Ops::Base::ToString(pDescPtr->GetDataType());148 std::string valueStr = Ops::Base::ToString(pDescPtr->GetDataType());
147- std::string reasonMsg = "Unsupported p dtype";149+ std::string reasonMsg = "Unsupported p dtype, must be in [DT_FLOAT, DT_FLOAT16, DT_BF16, DT_DOUBLE]";
148 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input p", valueStr.c_str(),150 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input p", valueStr.c_str(),
149 reasonMsg.c_str());151 reasonMsg.c_str());
150 return ge::GRAPH_FAILED;152 return ge::GRAPH_FAILED;
@@ -48,7 +48,8 @@ static const std::initializer_list<op::DataType> MASK_DTYPE_SUPPORT_LIST = {op::
48static inline bool CheckSocVersion()48static inline bool CheckSocVersion()
49{49{
50 if (op::GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510) {50 if (op::GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510) {
51- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "aclnnDropoutV3Grad is not supported in current socversion.");51+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "aclnnDropoutV3Grad is not supported in current socversion %d.",
52+ static_cast<int>(op::GetCurrentPlatformInfo().GetCurNpuArch()));
52 return false;53 return false;
53 }54 }
54 return true;55 return true;
@@ -27,7 +27,7 @@ const int64_t TWENTY_FIVE_NUM = 25;
27 27 
28ge::graphStatus DropOutV3GradTilingFunc(gert::TilingContext* context)28ge::graphStatus DropOutV3GradTilingFunc(gert::TilingContext* context)
29{29{
30- OP_LOGD(context->GetNodeName(), "DropOutV3GradTiling running begin");30+ OP_LOGD(context->GetNodeName(), "DropOutV3GradTiling started");
31 auto compileInfo = context->GetCompileInfo<DropOutV3GradCompileInfo>();31 auto compileInfo = context->GetCompileInfo<DropOutV3GradCompileInfo>();
32 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);32 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
33 return DropOutV3GradTilingForAscendC(context);33 return DropOutV3GradTilingForAscendC(context);
@@ -47,14 +47,14 @@ ge::graphStatus TilingPrepareDropOutV3GradForAscendC(gert::TilingParseContext* c
47 uint64_t ubSize = 0;47 uint64_t ubSize = 0;
48 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);48 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
49 compileInfo->ubSize = static_cast<int64_t>(ubSize);49 compileInfo->ubSize = static_cast<int64_t>(ubSize);
50- OP_CHECK_IF((compileInfo->ubSize <= 0), OP_LOGE(context->GetNodeName(), "Invalid ub size."),50+ OP_CHECK_IF((compileInfo->ubSize <= 0),
51- return ge::GRAPH_FAILED);51+ OP_LOGE(context->GetNodeName(), "Invalid ub size %ld.", compileInfo->ubSize), return ge::GRAPH_FAILED);
52 return ge::GRAPH_SUCCESS;52 return ge::GRAPH_SUCCESS;
53}53}
54 54 
55ge::graphStatus TilingPrepareForDropOutV3Grad(gert::TilingParseContext* context)55ge::graphStatus TilingPrepareForDropOutV3Grad(gert::TilingParseContext* context)
56{56{
57- OP_LOGD(context->GetNodeName(), "TilingPrepareForDropOutV3Grad running begin");57+ OP_LOGD(context->GetNodeName(), "TilingPrepareForDropOutV3Grad started");
58 auto compileInfo = context->GetCompiledInfo<DropOutV3GradCompileInfo>();58 auto compileInfo = context->GetCompiledInfo<DropOutV3GradCompileInfo>();
59 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);59 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
60 return TilingPrepareDropOutV3GradForAscendC(context);60 return TilingPrepareDropOutV3GradForAscendC(context);
@@ -80,7 +80,7 @@ ge::graphStatus DropOutV3GradTiling::GetShapeAttrsInfo()
80 dType_ = gradYPtr->GetDataType();80 dType_ = gradYPtr->GetDataType();
81 if (DROP_SUPPORTED_DTYPE.find(dType_) == DROP_SUPPORTED_DTYPE.end()) {81 if (DROP_SUPPORTED_DTYPE.find(dType_) == DROP_SUPPORTED_DTYPE.end()) {
82 std::string valueStr = ToString(dType_);82 std::string valueStr = ToString(dType_);
83- std::string reasonMsg = "grad_y dtype only support float32, float16, bfloat16";83+ std::string reasonMsg = "grad_y dtype only supports float32, float16, bfloat16";
84 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor grad_y", valueStr.c_str(),84 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor grad_y", valueStr.c_str(),
85 reasonMsg.c_str());85 reasonMsg.c_str());
86 return ge::GRAPH_FAILED;86 return ge::GRAPH_FAILED;
@@ -92,7 +92,7 @@ ge::graphStatus DropOutV3GradTiling::GetShapeAttrsInfo()
92 bool dtypeInValid = (maskDtype != ge::DT_UINT8 && maskDtype != ge::DT_UINT1);92 bool dtypeInValid = (maskDtype != ge::DT_UINT8 && maskDtype != ge::DT_UINT1);
93 if (dtypeInValid) {93 if (dtypeInValid) {
94 std::string valueStr = ToString(maskDtype);94 std::string valueStr = ToString(maskDtype);
95- std::string reasonMsg = "mask dtype only support uint8, uint1 currently";95+ std::string reasonMsg = "mask dtype only supports uint8, uint1 currently";
96 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor mask", valueStr.c_str(),96 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor mask", valueStr.c_str(),
97 reasonMsg.c_str());97 reasonMsg.c_str());
98 return ge::GRAPH_FAILED;98 return ge::GRAPH_FAILED;
@@ -117,7 +117,8 @@ ge::graphStatus DropOutV3GradTiling::GetShapeAttrsInfo()
117ge::graphStatus DropOutV3GradTiling::DoOpTiling()117ge::graphStatus DropOutV3GradTiling::DoOpTiling()
118{118{
119 typeSize_ = ge::GetSizeByDataType(dType_);119 typeSize_ = ge::GetSizeByDataType(dType_);
120- OP_CHECK_IF(typeSize_ <= 0, OP_LOGE(context_->GetNodeName(), "get dataType size fail."), return ge::GRAPH_FAILED);120+ OP_CHECK_IF(typeSize_ <= 0, OP_LOGE(context_->GetNodeName(), "Failed to get dataType size."),
121+ return ge::GRAPH_FAILED);
121 // total: ub/db122 // total: ub/db
122 // used: grad_y*typesize (grad_y), grad_y/8 (mask), grad_y*typesize(grad_x)123 // used: grad_y*typesize (grad_y), grad_y/8 (mask), grad_y*typesize(grad_x)
123 int64_t ubBlock = GetUbBlockSize(context_);124 int64_t ubBlock = GetUbBlockSize(context_);
@@ -97,7 +97,7 @@ static const std::initializer_list<DataType>& GetOutDtypeSupportList()
97 } else if (IsRegBase()) {97 } else if (IsRegBase()) {
98 return ARCH3510_DTYPE_SUPPORT_LIST;98 return ARCH3510_DTYPE_SUPPORT_LIST;
99 } else {99 } else {
100- OP_LOGW("Unknown SocVersion.");100+ OP_LOGW("Unknown SocVersion %d.", static_cast<int>(socVersion));
101 return EMPTY_LIST;101 return EMPTY_LIST;
102 }102 }
103}103}
@@ -112,15 +112,12 @@ static const std::initializer_list<DataType>& GetProbDtypeSupportList()
112 } else if (IsRegBase()) {112 } else if (IsRegBase()) {
113 return ARCH3510_PROB_DTYPE_SUPPORT_LIST;113 return ARCH3510_PROB_DTYPE_SUPPORT_LIST;
114 } else {114 } else {
115- OP_LOGW("Unknown SocVersion.");115+ OP_LOGW("Unknown SocVersion %d.", static_cast<int>(socVersion));
116 return EMPTY_LIST;116 return EMPTY_LIST;
117 }117 }
118}118}
119 119 
120-static bool IsDoubleEqual(double f1, double f2)120+static bool IsDoubleEqual(double f1, double f2) { return std::abs(f1 - f2) <= std::numeric_limits<double>::epsilon(); }
121-{
122- return std::abs(f1 - f2) <= std::numeric_limits<double>::epsilon();
123-}
124 121 
125static bool CheckDtypeValidTensor(const aclTensor* self, const aclTensor* prob, const aclTensor* out)122static bool CheckDtypeValidTensor(const aclTensor* self, const aclTensor* prob, const aclTensor* out)
126{123{
@@ -167,7 +164,7 @@ static bool CheckProb(const aclScalar* prob)
167{164{
168 // 检查y的数据类型是否在支持列表内165 // 检查y的数据类型是否在支持列表内
169 if (prob->ToDouble() > 1 || prob->ToDouble() < 0) {166 if (prob->ToDouble() > 1 || prob->ToDouble() < 0) {
170- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "prob should be in range 0<=prob<=1 .");167+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "prob should be in range [0, 1], but got %f.", prob->ToDouble());
171 return false;168 return false;
172 }169 }
173 170 
@@ -178,9 +175,8 @@ static bool CheckFormat(const aclTensor* self)
178{175{
179 // 如果输入格式是私有格式,记录日志,直接报错176 // 如果输入格式是私有格式,记录日志,直接报错
180 if (op::IsPrivateFormat(self->GetStorageFormat())) {177 if (op::IsPrivateFormat(self->GetStorageFormat())) {
181- OP_LOGE(178+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Format only supports ND, NCHW, NHWC, HWCN, NDHWC, NCDHW, self [%s]",
182- ACLNN_ERR_PARAM_INVALID, "Format only support ND、NCHW、NHWC、HWCN、NDHWC、NCDHW, self [%s]",179+ ToString(self->GetStorageFormat()).GetString());
183- ToString(self->GetStorageFormat()).GetString());
184 return false;180 return false;
185 }181 }
186 return true;182 return true;
@@ -190,9 +186,9 @@ static bool CheckFormatTensor(const aclTensor* self, const aclTensor* prob)
190{186{
191 // 如果输入格式是私有格式,记录日志,直接报错187 // 如果输入格式是私有格式,记录日志,直接报错
192 if (op::IsPrivateFormat(self->GetStorageFormat()) || op::IsPrivateFormat(prob->GetStorageFormat())) {188 if (op::IsPrivateFormat(self->GetStorageFormat()) || op::IsPrivateFormat(prob->GetStorageFormat())) {
193- OP_LOGE(189+ OP_LOGE(ACLNN_ERR_PARAM_INVALID,
194- ACLNN_ERR_PARAM_INVALID, "Format only support ND、NCHW、NHWC、HWCN、NDHWC、NCDHW, self [%s], prob [%s]",190+ "Format only supports ND, NCHW, NHWC, HWCN, NDHWC, NCDHW, self [%s], prob [%s]",
195- ToString(self->GetStorageFormat()).GetString(), ToString(prob->GetStorageFormat()).GetString());191+ ToString(self->GetStorageFormat()).GetString(), ToString(prob->GetStorageFormat()).GetString());
196 return false;192 return false;
197 }193 }
198 return true;194 return true;
@@ -275,9 +271,8 @@ static inline int64_t InferDSAOutShapeV2(const aclIntArray* shape)
275 return (size + VEC_BIT_NUMBER - 1) / VEC_BIT_NUMBER * VEC_BIT_NUMBER / UINT8_BIT_NUMBER;271 return (size + VEC_BIT_NUMBER - 1) / VEC_BIT_NUMBER * VEC_BIT_NUMBER / UINT8_BIT_NUMBER;
276}272}
277 273 
278-aclnnStatus GetBernoulliByDSA(274+aclnnStatus GetBernoulliByDSA(const aclTensor* inputContiguous, const aclScalar* prob, int64_t seed, int64_t offset,
279- const aclTensor* inputContiguous, const aclScalar* prob, int64_t seed, int64_t offset, const aclTensor*& doMaskOut,275+ const aclTensor*& doMaskOut, aclOpExecutor* executor)
280- aclOpExecutor* executor)
281{276{
282 auto inputShape = op::ToShapeVector(inputContiguous->GetViewShape());277 auto inputShape = op::ToShapeVector(inputContiguous->GetViewShape());
283 auto dims = executor->ConvertToTensor(inputShape.data(), inputShape.size(), DataType::DT_INT64);278 auto dims = executor->ConvertToTensor(inputShape.data(), inputShape.size(), DataType::DT_INT64);
@@ -300,9 +295,9 @@ aclnnStatus GetBernoulliByDSA(
300 return ACLNN_SUCCESS;295 return ACLNN_SUCCESS;
301}296}
302 297 
303-aclnnStatus aclnnBernoulliTensorGetWorkspaceSize(298+aclnnStatus aclnnBernoulliTensorGetWorkspaceSize(const aclTensor* self, const aclTensor* prob, int64_t seed,
304- const aclTensor* self, const aclTensor* prob, int64_t seed, int64_t offset, aclTensor* out, uint64_t* workspaceSize,299+ int64_t offset, aclTensor* out, uint64_t* workspaceSize,
305- aclOpExecutor** executor)300+ aclOpExecutor** executor)
306{301{
307 OP_CHECK_COMM_INPUT(workspaceSize, executor);302 OP_CHECK_COMM_INPUT(workspaceSize, executor);
308 303 
@@ -351,9 +346,8 @@ aclnnStatus aclnnBernoulliTensorGetWorkspaceSize(
351 return ACLNN_SUCCESS;346 return ACLNN_SUCCESS;
352}347}
353 348 
354-aclnnStatus aclnnBernoulliGetWorkspaceSize(349+aclnnStatus aclnnBernoulliGetWorkspaceSize(const aclTensor* self, const aclScalar* prob, int64_t seed, int64_t offset,
355- const aclTensor* self, const aclScalar* prob, int64_t seed, int64_t offset, aclTensor* out, uint64_t* workspaceSize,350+ aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor)
356- aclOpExecutor** executor)
357{351{
358 OP_CHECK_COMM_INPUT(workspaceSize, executor);352 OP_CHECK_COMM_INPUT(workspaceSize, executor);
359 353 
@@ -387,8 +381,8 @@ aclnnStatus aclnnBernoulliGetWorkspaceSize(
387 } else if (IsDoubleEqual(prob->ToDouble(), 1)) {381 } else if (IsDoubleEqual(prob->ToDouble(), 1)) {
388 doMaskOut = l0op::OnesLike(inputContiguous, uniqueExecutor.get());382 doMaskOut = l0op::OnesLike(inputContiguous, uniqueExecutor.get());
389 } else {383 } else {
390- auto executeResult =384+ auto executeResult = GetBernoulliByDSA(inputContiguous, prob, seed, offset, doMaskOut,
391- GetBernoulliByDSA(inputContiguous, prob, seed, offset, doMaskOut, uniqueExecutor.get());385+ uniqueExecutor.get());
392 CHECK_RET(executeResult == ACLNN_SUCCESS, executeResult);386 CHECK_RET(executeResult == ACLNN_SUCCESS, executeResult);
393 }387 }
394 CHECK_RET(doMaskOut != nullptr, ACLNN_ERR_INNER_NULLPTR);388 CHECK_RET(doMaskOut != nullptr, ACLNN_ERR_INNER_NULLPTR);
@@ -417,17 +411,16 @@ aclnnStatus aclnnBernoulliGetWorkspaceSize(
417 return ACLNN_SUCCESS;411 return ACLNN_SUCCESS;
418}412}
419 413 
420-aclnnStatus aclnnInplaceBernoulliGetWorkspaceSize(414+aclnnStatus aclnnInplaceBernoulliGetWorkspaceSize(const aclTensor* selfRef, const aclScalar* prob, int64_t seed,
421- const aclTensor* selfRef, const aclScalar* prob, int64_t seed, int64_t offset, uint64_t* workspaceSize,415+ int64_t offset, uint64_t* workspaceSize, aclOpExecutor** executor)
422- aclOpExecutor** executor)
423{416{
424 auto out = const_cast<aclTensor*>(selfRef);417 auto out = const_cast<aclTensor*>(selfRef);
425 return aclnnBernoulliGetWorkspaceSize(selfRef, prob, seed, offset, out, workspaceSize, executor);418 return aclnnBernoulliGetWorkspaceSize(selfRef, prob, seed, offset, out, workspaceSize, executor);
426}419}
427 420 
428-aclnnStatus aclnnInplaceBernoulliTensorGetWorkspaceSize(421+aclnnStatus aclnnInplaceBernoulliTensorGetWorkspaceSize(const aclTensor* selfRef, const aclTensor* prob, int64_t seed,
429- const aclTensor* selfRef, const aclTensor* prob, int64_t seed, int64_t offset, uint64_t* workspaceSize,422+ int64_t offset, uint64_t* workspaceSize,
430- aclOpExecutor** executor)423+ aclOpExecutor** executor)
431{424{
432 auto out = const_cast<aclTensor*>(selfRef);425 auto out = const_cast<aclTensor*>(selfRef);
433 return aclnnBernoulliTensorGetWorkspaceSize(selfRef, prob, seed, offset, out, workspaceSize, executor);426 return aclnnBernoulliTensorGetWorkspaceSize(selfRef, prob, seed, offset, out, workspaceSize, executor);
@@ -447,8 +440,8 @@ aclnnStatus aclnnBernoulli(void* workspace, uint64_t workspaceSize, aclOpExecuto
447 return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);440 return CommonOpExecutorRun(workspace, workspaceSize, executor, stream);
448}441}
449 442 
450-aclnnStatus aclnnInplaceBernoulliTensor(443+aclnnStatus aclnnInplaceBernoulliTensor(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor,
451- void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream)444+ aclrtStream stream)
452{445{
453 L2_DFX_PHASE_2(aclnnInplaceBernoulliTensor);446 L2_DFX_PHASE_2(aclnnInplaceBernoulliTensor);
454 // 固定写法,调用框架能力,完成计算447 // 固定写法,调用框架能力,完成计算
@@ -11,8 +11,6 @@
11# ----------------------------------------------------------------------------11# ----------------------------------------------------------------------------
12 12 
13import torch13import torch
14-import torch_npu
15-import numpy as np
16import tensorflow as tf14import tensorflow as tf
17 15 
18from atk.configs.dataset_config import InputDataset16from atk.configs.dataset_config import InputDataset
@@ -22,16 +20,15 @@ from atk.tasks.api_execute.base_api import BaseApi
22from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi20from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi
23from atk.tasks.dataset.base_dataset import OpsDataset21from atk.tasks.dataset.base_dataset import OpsDataset
24 22 
23+ 
25def uniform_golden(torch_tensor, params):24def uniform_golden(torch_tensor, params):
26 seed = []25 seed = []
27 offset = [0]26 offset = [0]
28- start = params["from"]
29- end = params["to"]
30 seed.append(params["seed"])27 seed.append(params["seed"])
31 offset.append(params["offset"])28 offset.append(params["offset"])
32 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True29 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
33 if not is_contiguous:30 if not is_contiguous:
34- print("------------非连续操作-------------")31+ print("---- Non-contiguous case ----")
35 torch_tensor = torch.transpose(torch_tensor, 0, 1)32 torch_tensor = torch.transpose(torch_tensor, 0, 1)
36 matrix = torch_tensor33 matrix = torch_tensor
37 if params["dtype_input"][0] == torch.bfloat16:34 if params["dtype_input"][0] == torch.bfloat16:
@@ -41,17 +38,20 @@ def uniform_golden(torch_tensor, params):
41 matrix_shape = list(matrix.shape)38 matrix_shape = list(matrix.shape)
42 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])39 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])
43 40 
44- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1)41+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
45- 42+ shape=matrix_shape, key=seed, counter=offset, alg=1
43+ )
44+ 
46 output_data = tf.cast(uniform_data, dtype)45 output_data = tf.cast(uniform_data, dtype)
47 if output_data.shape == []:46 if output_data.shape == []:
48 output_data = torch.tensor(output_data.numpy())47 output_data = torch.tensor(output_data.numpy())
49 else:48 else:
50 output_data = torch.from_numpy(output_data.numpy())49 output_data = torch.from_numpy(output_data.numpy())
51- 50+ 
52 output_data = output_data.type(params["dtype_input"][0])51 output_data = output_data.type(params["dtype_input"][0])
53 return output_data52 return output_data
54 53 
54+ 
55@register("ascend_aclnn_dropout")55@register("ascend_aclnn_dropout")
56class MethodAclnnDropoutApi(BaseApi):56class MethodAclnnDropoutApi(BaseApi):
57 def __init__(self, task_result: TaskResult):57 def __init__(self, task_result: TaskResult):
@@ -60,9 +60,9 @@ class MethodAclnnDropoutApi(BaseApi):
60 self.change_flag = None60 self.change_flag = None
61 61 
62 def __call__(self, input_data: InputDataset, with_output: bool = False):62 def __call__(self, input_data: InputDataset, with_output: bool = False):
63- self.input = input_data.kwargs['input']63+ self.input = input_data.kwargs["input"]
64- self.p = input_data.kwargs['p']64+ self.p = input_data.kwargs["p"]
65- self.train = input_data.kwargs['train']65+ self.train = input_data.kwargs["train"]
66 self.seed = input_data.kwargs["seed"]66 self.seed = input_data.kwargs["seed"]
67 self.offset = input_data.kwargs["offset"]67 self.offset = input_data.kwargs["offset"]
68 self.shape = self.input.shape68 self.shape = self.input.shape
@@ -70,25 +70,35 @@ class MethodAclnnDropoutApi(BaseApi):
70 self.count = 170 self.count = 1
71 for item in self.shape:71 for item in self.shape:
72 self.count *= item72 self.count *= item
73- self.tensor = torch.ones([self.count], dtype = torch.float32)73+ self.tensor = torch.ones([self.count], dtype=torch.float32)
74 shape_x = self.input.shape74 shape_x = self.input.shape
75 75 
76 inputx = self.input.cpu()76 inputx = self.input.cpu()
77- params = {"from": 0.0, "to": 1.0, "seed": self.seed, "offset": self.offset, "is_contiguous": True,77+ params = {
78- "dtype_input": [torch.float32]}78+ "from": 0.0,
79+ "to": 1.0,
80+ "seed": self.seed,
81+ "offset": self.offset,
82+ "is_contiguous": True,
83+ "dtype_input": [torch.float32],
84+ }
79 85 
80 x = uniform_golden(self.tensor, params)86 x = uniform_golden(self.tensor, params)
81 output1 = x.to(torch.float32) >= torch.tensor([self.p], dtype=torch.float32)87 output1 = x.to(torch.float32) >= torch.tensor([self.p], dtype=torch.float32)
82 output1 = torch.tensor(output1, dtype=torch.float32).to(torch.uint8)88 output1 = torch.tensor(output1, dtype=torch.float32).to(torch.uint8)
83- output1[self.count:] = 089+ output1[self.count :] = 0
84 90 
85- mask = torch.zeros([int(int((self.count + 127) / 128) * 128 / 8)], dtype=torch.uint8)91+ mask = torch.zeros(
92+ [int(int((self.count + 127) / 128) * 128 / 8)], dtype=torch.uint8
93+ )
86 94 
87- mask_tensor = output1[:self.count].reshape(shape_x)95+ mask_tensor = output1[: self.count].reshape(shape_x)
88 96 
89 if self.input.dtype == torch.bfloat16:97 if self.input.dtype == torch.bfloat16:
90 keep_prob = torch.tensor(1.0 - self.p, dtype=torch.float32)98 keep_prob = torch.tensor(1.0 - self.p, dtype=torch.float32)
91- keep_prob_scalar_input_dtype = keep_prob.to(dtype=self.input.dtype).to(dtype=torch.float32)99+ keep_prob_scalar_input_dtype = keep_prob.to(dtype=self.input.dtype).to(
100+ dtype=torch.float32
101+ )
92 scale = 1.0 / keep_prob_scalar_input_dtype.to(dtype=torch.float32)102 scale = 1.0 / keep_prob_scalar_input_dtype.to(dtype=torch.float32)
93 103 
94 x_scaled = inputx.to(dtype=torch.float32) * scale104 x_scaled = inputx.to(dtype=torch.float32) * scale
@@ -106,10 +116,11 @@ class MethodAclnnDropoutApi(BaseApi):
106 116 
107 return output_data.to(dtype=self.input.dtype).to(torch.float32), mask117 return output_data.to(dtype=self.input.dtype).to(torch.float32), mask
108 118 
119+ 
109@register("aclnn_dropout")120@register("aclnn_dropout")
110class DropoutAclnnApi(AclnnBaseApi):121class DropoutAclnnApi(AclnnBaseApi):
111 def after_call(self, output_packages):122 def after_call(self, output_packages):
112 output1, output2 = super().after_call(output_packages)123 output1, output2 = super().after_call(output_packages)
113 output2.zero_()124 output2.zero_()
114 125 
115- return output1, output2126+ return output1, output2
@@ -11,7 +11,6 @@
11# ----------------------------------------------------------------------------11# ----------------------------------------------------------------------------
12 12 
13import torch13import torch
14-import torch_npu
15import tensorflow as tf14import tensorflow as tf
16import numpy as np15import numpy as np
17from atk.configs.dataset_config import InputDataset16from atk.configs.dataset_config import InputDataset
@@ -21,6 +20,7 @@ from atk.tasks.api_execute.base_api import BaseApi
21from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi20from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi
22from atk.tasks.dataset.base_dataset import OpsDataset21from atk.tasks.dataset.base_dataset import OpsDataset
23 22 
23+ 
24def revert_bit(n):24def revert_bit(n):
25 result = 025 result = 0
26 for i in range(8):26 for i in range(8):
@@ -29,12 +29,14 @@ def revert_bit(n):
29 n >>= 129 n >>= 1
30 return result30 return result
31 31 
32+ 
32def revert_array_bit(arr):33def revert_array_bit(arr):
33 res = []34 res = []
34 for item in np.array(arr).flatten():35 for item in np.array(arr).flatten():
35 res.append(revert_bit(item))36 res.append(revert_bit(item))
36 return np.array(res, dtype=np.uint8).reshape(np.array(arr).shape)37 return np.array(res, dtype=np.uint8).reshape(np.array(arr).shape)
37 38 
39+ 
38def bitmask_to_list(input_x, input_mask):40def bitmask_to_list(input_x, input_mask):
39 input_dtype = input_x.dtype41 input_dtype = input_x.dtype
40 shape_x = input_x.shape42 shape_x = input_x.shape
@@ -47,16 +49,15 @@ def bitmask_to_list(input_x, input_mask):
47 output = mask_tensor.to(dtype=torch.float32)49 output = mask_tensor.to(dtype=torch.float32)
48 return output.to(dtype=input_dtype)50 return output.to(dtype=input_dtype)
49 51 
52+ 
50def uniform_golden(torch_tensor, params):53def uniform_golden(torch_tensor, params):
51 seed = []54 seed = []
52 offset = [0]55 offset = [0]
53- start = params["from"]
54- end = params["to"]
55 seed.append(params["seed"])56 seed.append(params["seed"])
56 offset.append(params["offset"])57 offset.append(params["offset"])
57 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True58 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
58 if not is_contiguous:59 if not is_contiguous:
59- print("------------非连续操作-------------")60+ print("---- Non-contiguous case ----")
60 torch_tensor = torch.transpose(torch_tensor, 0, 1)61 torch_tensor = torch.transpose(torch_tensor, 0, 1)
61 matrix = torch_tensor62 matrix = torch_tensor
62 if params["dtype_input"][0] == torch.bfloat16:63 if params["dtype_input"][0] == torch.bfloat16:
@@ -66,17 +67,20 @@ def uniform_golden(torch_tensor, params):
66 matrix_shape = list(matrix.shape)67 matrix_shape = list(matrix.shape)
67 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])68 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])
68 69 
69- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1)70+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
71+ shape=matrix_shape, key=seed, counter=offset, alg=1
72+ )
70 73 
71 output_data = tf.cast(uniform_data, dtype)74 output_data = tf.cast(uniform_data, dtype)
72 if output_data.shape == []:75 if output_data.shape == []:
73 output_data = torch.tensor(output_data.numpy())76 output_data = torch.tensor(output_data.numpy())
74 else:77 else:
75 output_data = torch.from_numpy(output_data.numpy())78 output_data = torch.from_numpy(output_data.numpy())
76- 79+ 
77 output_data = output_data.type(params["dtype_input"][0])80 output_data = output_data.type(params["dtype_input"][0])
78 return output_data81 return output_data
79 82 
83+ 
80@register("ascend_aclnn_dropout_gen_mask")84@register("ascend_aclnn_dropout_gen_mask")
81class MethodAclnnDropoutGenMaskApi(BaseApi):85class MethodAclnnDropoutGenMaskApi(BaseApi):
82 def __init__(self, task_result: TaskResult):86 def __init__(self, task_result: TaskResult):
@@ -95,31 +99,42 @@ class MethodAclnnDropoutGenMaskApi(BaseApi):
95 self.count = 199 self.count = 1
96 for item in self.shape:100 for item in self.shape:
97 self.count *= item101 self.count *= item
98- self.tensor = torch.ones([self.count], dtype = torch.float32)102+ self.tensor = torch.ones([self.count], dtype=torch.float32)
99- 103+ 
100- params = {"from": 0.0, "to": 1.0, "seed": self.seed, "offset": self.offset, "is_contiguous": True,104+ params = {
101- "dtype_input": [torch.float32]}105+ "from": 0.0,
106+ "to": 1.0,
107+ "seed": self.seed,
108+ "offset": self.offset,
109+ "is_contiguous": True,
110+ "dtype_input": [torch.float32],
111+ }
102 112 
103 x = uniform_golden(self.tensor, params)113 x = uniform_golden(self.tensor, params)
104 output1 = x.to(torch.float32) >= torch.tensor([self.prob], dtype=torch.float32)114 output1 = x.to(torch.float32) >= torch.tensor([self.prob], dtype=torch.float32)
105 output1 = torch.tensor(output1, dtype=torch.float32).to(torch.uint8)115 output1 = torch.tensor(output1, dtype=torch.float32).to(torch.uint8)
106- output1[self.count:] = 0116+ output1[self.count :] = 0
107 117 
108 return output1.to(torch.uint8).contiguous()118 return output1.to(torch.uint8).contiguous()
109 119 
120+ 
110@register("aclnn_dropout_gen_mask")121@register("aclnn_dropout_gen_mask")
111class DropoutGenMaskAclnnApi(AclnnBaseApi):122class DropoutGenMaskAclnnApi(AclnnBaseApi):
112 def init_by_input_data(self, input_data: InputDataset):123 def init_by_input_data(self, input_data: InputDataset):
113 self.shape = input_data.kwargs["shape"]124 self.shape = input_data.kwargs["shape"]
114- self.tensor = torch.ones(self.shape, dtype = torch.float32).to("npu")125+ self.tensor = torch.ones(self.shape, dtype=torch.float32).to("npu")
115 input_args, output_packages = super().init_by_input_data(input_data)126 input_args, output_packages = super().init_by_input_data(input_data)
116 127 
117 self.count = 1128 self.count = 1
118 for item in self.shape:129 for item in self.shape:
119 self.count *= item130 self.count *= item
120- self.task_result.output_info_list[0].shape = [int(int((self.count + 127) / 128) * 128 / 8)]131+ self.task_result.output_info_list[0].shape = [
132+ int(int((self.count + 127) / 128) * 128 / 8)
133+ ]
121 self.task_result.output_info_list[0].stride = [1]134 self.task_result.output_info_list[0].stride = [1]
122- output = self.backend.convert_output_data(self.task_result.output_info_list[0], 0)135+ output = self.backend.convert_output_data(
136+ self.task_result.output_info_list[0], 0
137+ )
123 output_packages[0] = output[0]138 output_packages[0] = output[0]
124 input_args[-1] = output_packages[0]139 input_args[-1] = output_packages[0]
125 140 
@@ -127,7 +142,9 @@ class DropoutGenMaskAclnnApi(AclnnBaseApi):
127 142 
128 def after_call(self, output_packages):143 def after_call(self, output_packages):
129 output1 = super().after_call(output_packages)144 output1 = super().after_call(output_packages)
130- output1[0] = bitmask_to_list(self.tensor, output1[0].cpu()).to(torch.uint8).npu()145+ output1[0] = (
131- output1[0][self.count:] = 0146+ bitmask_to_list(self.tensor, output1[0].cpu()).to(torch.uint8).npu()
147+ )
148+ output1[0][self.count :] = 0
132 149 
133- return output1150+ return output1
@@ -11,7 +11,6 @@
11# ----------------------------------------------------------------------------11# ----------------------------------------------------------------------------
12 12 
13import torch13import torch
14-import torch_npu
15import tensorflow as tf14import tensorflow as tf
16import numpy as np15import numpy as np
17from atk.configs.dataset_config import InputDataset16from atk.configs.dataset_config import InputDataset
@@ -21,6 +20,7 @@ from atk.tasks.api_execute.base_api import BaseApi
21from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi20from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi
22from atk.tasks.dataset.base_dataset import OpsDataset21from atk.tasks.dataset.base_dataset import OpsDataset
23 22 
23+ 
24def revert_bit(n):24def revert_bit(n):
25 result = 025 result = 0
26 for i in range(8):26 for i in range(8):
@@ -29,12 +29,14 @@ def revert_bit(n):
29 n >>= 129 n >>= 1
30 return result30 return result
31 31 
32+ 
32def revert_array_bit(arr):33def revert_array_bit(arr):
33 res = []34 res = []
34 for item in np.array(arr).flatten():35 for item in np.array(arr).flatten():
35 res.append(revert_bit(item))36 res.append(revert_bit(item))
36 return np.array(res, dtype=np.uint8).reshape(np.array(arr).shape)37 return np.array(res, dtype=np.uint8).reshape(np.array(arr).shape)
37 38 
39+ 
38def bitmask_to_list(input_x, input_mask):40def bitmask_to_list(input_x, input_mask):
39 input_dtype = input_x.dtype41 input_dtype = input_x.dtype
40 shape_x = input_x.shape42 shape_x = input_x.shape
@@ -47,16 +49,15 @@ def bitmask_to_list(input_x, input_mask):
47 output = mask_tensor.to(dtype=torch.float32)49 output = mask_tensor.to(dtype=torch.float32)
48 return output.to(dtype=input_dtype)50 return output.to(dtype=input_dtype)
49 51 
52+ 
50def uniform_golden(torch_tensor, params):53def uniform_golden(torch_tensor, params):
51 seed = []54 seed = []
52 offset = [0]55 offset = [0]
53- start = params["from"]
54- end = params["to"]
55 seed.append(params["seed"])56 seed.append(params["seed"])
56 offset.append(params["offset"])57 offset.append(params["offset"])
57 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True58 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
58 if not is_contiguous:59 if not is_contiguous:
59- print("------------非连续操作-------------")60+ print("---- Non-contiguous case ----")
60 torch_tensor = torch.transpose(torch_tensor, 0, 1)61 torch_tensor = torch.transpose(torch_tensor, 0, 1)
61 matrix = torch_tensor62 matrix = torch_tensor
62 if params["dtype_input"][0] == torch.bfloat16:63 if params["dtype_input"][0] == torch.bfloat16:
@@ -66,23 +67,32 @@ def uniform_golden(torch_tensor, params):
66 matrix_shape = list(matrix.shape)67 matrix_shape = list(matrix.shape)
67 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])68 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])
68 if params["dtype_input"][0] == torch.bfloat16:69 if params["dtype_input"][0] == torch.bfloat16:
69- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1,70+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
70- dtype=tf.dtypes.bfloat16)71+ shape=matrix_shape,
72+ key=seed,
73+ counter=offset,
74+ alg=1,
75+ dtype=tf.dtypes.bfloat16,
76+ )
71 elif params["dtype_input"][0] == torch.float16:77 elif params["dtype_input"][0] == torch.float16:
72- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1,78+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
73- dtype=tf.dtypes.float16)79+ shape=matrix_shape, key=seed, counter=offset, alg=1, dtype=tf.dtypes.float16
80+ )
74 else:81 else:
75- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1)82+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
76- 83+ shape=matrix_shape, key=seed, counter=offset, alg=1
84+ )
85+ 
77 output_data = tf.cast(uniform_data, dtype)86 output_data = tf.cast(uniform_data, dtype)
78 if output_data.shape == []:87 if output_data.shape == []:
79 output_data = torch.tensor(output_data.numpy())88 output_data = torch.tensor(output_data.numpy())
80 else:89 else:
81 output_data = torch.from_numpy(output_data.numpy())90 output_data = torch.from_numpy(output_data.numpy())
82- 91+ 
83 output_data = output_data.type(params["dtype_input"][0])92 output_data = output_data.type(params["dtype_input"][0])
84 return output_data93 return output_data
85 94 
95+ 
86@register("ascend_aclnn_dropout_gen_mask_v2")96@register("ascend_aclnn_dropout_gen_mask_v2")
87class MethodAclnnDropoutGenMaskV2Api(BaseApi):97class MethodAclnnDropoutGenMaskV2Api(BaseApi):
88 def __init__(self, task_result: TaskResult):98 def __init__(self, task_result: TaskResult):
@@ -91,7 +101,13 @@ class MethodAclnnDropoutGenMaskV2Api(BaseApi):
91 self.change_flag = None101 self.change_flag = None
92 102 
93 def init_by_input_data(self, input_data: InputDataset):103 def init_by_input_data(self, input_data: InputDataset):
94- input_data.kwargs["prob"] = torch.tensor([input_data.kwargs["prob"]], dtype=input_data.kwargs["probDataType"]).to(torch.float32).numpy()[0]104+ input_data.kwargs["prob"] = (
105+ torch.tensor(
106+ [input_data.kwargs["prob"]], dtype=input_data.kwargs["probDataType"]
107+ )
108+ .to(torch.float32)
109+ .numpy()[0]
110+ )
95 111 
96 def __call__(self, input_data: InputDataset, with_output: bool = False):112 def __call__(self, input_data: InputDataset, with_output: bool = False):
97 self.shape = input_data.kwargs["shape"]113 self.shape = input_data.kwargs["shape"]
@@ -103,22 +119,35 @@ class MethodAclnnDropoutGenMaskV2Api(BaseApi):
103 self.count = 1119 self.count = 1
104 for item in self.shape:120 for item in self.shape:
105 self.count *= item121 self.count *= item
106- self.tensor = torch.ones([self.count], dtype = self.probDataType)122+ self.tensor = torch.ones([self.count], dtype=self.probDataType)
107- 123+ 
108- params = {"from": 0.0, "to": 1.0, "seed": self.seed, "offset": self.offset, "is_contiguous": True,124+ params = {
109- "dtype_input": [self.probDataType]}125+ "from": 0.0,
126+ "to": 1.0,
127+ "seed": self.seed,
128+ "offset": self.offset,
129+ "is_contiguous": True,
130+ "dtype_input": [self.probDataType],
131+ }
110 132 
111 x = uniform_golden(self.tensor, params)133 x = uniform_golden(self.tensor, params)
112 output1 = x.to(torch.float32) >= torch.tensor([self.prob], dtype=torch.float32)134 output1 = x.to(torch.float32) >= torch.tensor([self.prob], dtype=torch.float32)
113 output1 = torch.tensor(output1, dtype=torch.float32).to(torch.uint8)135 output1 = torch.tensor(output1, dtype=torch.float32).to(torch.uint8)
114- output1[self.count:] = 0136+ output1[self.count :] = 0
115 137 
116 return output1.to(torch.uint8).contiguous()138 return output1.to(torch.uint8).contiguous()
117 139 
140+ 
118@register("aclnn_dropout_gen_mask_v2")141@register("aclnn_dropout_gen_mask_v2")
119class DropoutGenMaskV2AclnnApi(AclnnBaseApi):142class DropoutGenMaskV2AclnnApi(AclnnBaseApi):
120 def init_by_input_data(self, input_data: InputDataset):143 def init_by_input_data(self, input_data: InputDataset):
121- input_data.kwargs["prob"] = torch.tensor([input_data.kwargs["prob"]], dtype=input_data.kwargs["probDataType"]).to(torch.float32).numpy()[0]144+ input_data.kwargs["prob"] = (
145+ torch.tensor(
146+ [input_data.kwargs["prob"]], dtype=input_data.kwargs["probDataType"]
147+ )
148+ .to(torch.float32)
149+ .numpy()[0]
150+ )
122 self.offset = input_data.kwargs["offset"]151 self.offset = input_data.kwargs["offset"]
123 self.seed = input_data.kwargs["seed"]152 self.seed = input_data.kwargs["seed"]
124 self.shape = input_data.kwargs["shape"]153 self.shape = input_data.kwargs["shape"]
@@ -128,10 +157,14 @@ class DropoutGenMaskV2AclnnApi(AclnnBaseApi):
128 self.count = 1157 self.count = 1
129 for item in self.shape:158 for item in self.shape:
130 self.count *= item159 self.count *= item
131- self.tensor = torch.ones([self.count], dtype = self.probDataType).to("npu")160+ self.tensor = torch.ones([self.count], dtype=self.probDataType).to("npu")
132- self.task_result.output_info_list[0].shape = [int(int((self.count + 127) / 128) * 128 / 8)]161+ self.task_result.output_info_list[0].shape = [
162+ int(int((self.count + 127) / 128) * 128 / 8)
163+ ]
133 self.task_result.output_info_list[0].stride = [1]164 self.task_result.output_info_list[0].stride = [1]
134- output = self.backend.convert_output_data(self.task_result.output_info_list[0], 0)165+ output = self.backend.convert_output_data(
166+ self.task_result.output_info_list[0], 0
167+ )
135 output_packages[0] = output[0]168 output_packages[0] = output[0]
136 input_args[-1] = output_packages[0]169 input_args[-1] = output_packages[0]
137 170 
@@ -139,7 +172,9 @@ class DropoutGenMaskV2AclnnApi(AclnnBaseApi):
139 172 
140 def after_call(self, output_packages):173 def after_call(self, output_packages):
141 output1 = super().after_call(output_packages)174 output1 = super().after_call(output_packages)
142- output1[0] = bitmask_to_list(self.tensor, output1[0].cpu()).to(torch.uint8).npu()175+ output1[0] = (
143- output1[0][self.count:] = 0176+ bitmask_to_list(self.tensor, output1[0].cpu()).to(torch.uint8).npu()
177+ )
178+ output1[0][self.count :] = 0
144 179 
145- return output1180+ return output1
@@ -53,7 +53,7 @@ static const std::initializer_list<DataType>& GetDtypeSupportList()
53static inline bool CheckDtypeValid(const aclTensor* self)53static inline bool CheckDtypeValid(const aclTensor* self)
54{54{
55 if (!CheckSocVersionIsSupportBf16() && (self->GetDataType() == op::DataType::DT_BF16)) {55 if (!CheckSocVersionIsSupportBf16() && (self->GetDataType() == op::DataType::DT_BF16)) {
56- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype DT_BF16 not support in current soc version.");56+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype DT_BF16 is not supported in current soc version.");
57 return false;57 return false;
58 }58 }
59 const auto& supportList = GetDtypeSupportList();59 const auto& supportList = GetDtypeSupportList();
@@ -19,8 +19,9 @@ from atk.tasks.api_execute.base_api import BaseApi
19from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi19from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi
20from atk.tasks.dataset.base_dataset import OpsDataset20from atk.tasks.dataset.base_dataset import OpsDataset
21 21 
22+ 
22def normal_golden(torch_tensor, params):23def normal_golden(torch_tensor, params):
23- if tf.__version__ >= '2':24+ if tf.__version__ >= "2":
24 tf.compat.v1.enable_eager_execution()25 tf.compat.v1.enable_eager_execution()
25 seed = []26 seed = []
26 offset = [0]27 offset = [0]
@@ -30,7 +31,7 @@ def normal_golden(torch_tensor, params):
30 offset.append(params["offset"])31 offset.append(params["offset"])
31 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True32 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
32 if not is_contiguous:33 if not is_contiguous:
33- print("------------非连续操作-------------")34+ print("---- Non-contiguous case ----")
34 torch_tensor = torch.transpose(torch_tensor, 0, 1)35 torch_tensor = torch.transpose(torch_tensor, 0, 1)
35 matrix = torch_tensor36 matrix = torch_tensor
36 if params["dtype_input"][0] == torch.bfloat16:37 if params["dtype_input"][0] == torch.bfloat16:
@@ -40,10 +41,12 @@ def normal_golden(torch_tensor, params):
40 else:41 else:
41 dtype = tf.dtypes.float3242 dtype = tf.dtypes.float32
42 matrix_shape = list(matrix.shape)43 matrix_shape = list(matrix.shape)
43- normal_data = tf.raw_ops.StatelessRandomNormalV2(shape=matrix_shape, key=seed, counter=offset, alg=1)44+ normal_data = tf.raw_ops.StatelessRandomNormalV2(
45+ shape=matrix_shape, key=seed, counter=offset, alg=1
46+ )
44 mul_data = tf.multiply(normal_data, std)47 mul_data = tf.multiply(normal_data, std)
45 add_data = tf.add(mul_data, mean)48 add_data = tf.add(mul_data, mean)
46- if tf.__version__ >= '2':49+ if tf.__version__ >= "2":
47 output_data = tf.cast(add_data, dtype).numpy()50 output_data = tf.cast(add_data, dtype).numpy()
48 output_data = torch.from_numpy(output_data)51 output_data = torch.from_numpy(output_data)
49 else:52 else:
@@ -52,7 +55,8 @@ def normal_golden(torch_tensor, params):
52 output_data = torch.from_numpy(output_data_tf)55 output_data = torch.from_numpy(output_data_tf)
53 output_data = output_data.to(params["dtype_input"][0])56 output_data = output_data.to(params["dtype_input"][0])
54 return output_data57 return output_data
55- 58+ 
59+ 
56@register("ascend_aclnn_inplace_normal")60@register("ascend_aclnn_inplace_normal")
57class MethodAclnnInplaceNormalApi(BaseApi):61class MethodAclnnInplaceNormalApi(BaseApi):
58 def __init__(self, task_result: TaskResult):62 def __init__(self, task_result: TaskResult):
@@ -62,30 +66,37 @@ class MethodAclnnInplaceNormalApi(BaseApi):
62 66 
63 def __call__(self, input_data: InputDataset, with_output: bool = False):67 def __call__(self, input_data: InputDataset, with_output: bool = False):
64 # 获取yaml中所需参数68 # 获取yaml中所需参数
65- self.tensor_ = input_data.kwargs['selfRef']69+ self.tensor_ = input_data.kwargs["selfRef"]
66- self.tensor_dtype_ = input_data.kwargs['selfRef'].dtype70+ self.tensor_dtype_ = input_data.kwargs["selfRef"].dtype
67- self.mean_ = input_data.kwargs['mean']71+ self.mean_ = input_data.kwargs["mean"]
68- self.std_ = input_data.kwargs['std']72+ self.std_ = input_data.kwargs["std"]
69- self.seed_ = input_data.kwargs['seed']73+ self.seed_ = input_data.kwargs["seed"]
70- self.offset_ = input_data.kwargs['offset']74+ self.offset_ = input_data.kwargs["offset"]
75+ 
76+ params = {
77+ "mean": self.mean_,
78+ "std": self.std_,
79+ "seed": self.seed_,
80+ "offset": self.offset_,
81+ "is_contiguous": True,
82+ "dtype_input": [self.tensor_dtype_],
83+ }
71 84 
72- params = {"mean": self.mean_, "std": self.std_, "seed": self.seed_, "offset": self.offset_, "is_contiguous": True,
73- "dtype_input": [self.tensor_dtype_]}
74-
75 x = normal_golden(self.tensor_, params)85 x = normal_golden(self.tensor_, params)
76- 86+ 
77 return x87 return x
78- 88+ 
89+ 
79@register("aclnn_inplace_normal")90@register("aclnn_inplace_normal")
80class InplaceNormalAclnnApi(AclnnBaseApi):91class InplaceNormalAclnnApi(AclnnBaseApi):
81 def __call__(self):92 def __call__(self):
82 super().__call__()93 super().__call__()
83- 94+ 
84 def init_by_input_data(self, input_data: InputDataset):95 def init_by_input_data(self, input_data: InputDataset):
85 input_args, output_packages = super().init_by_input_data(input_data)96 input_args, output_packages = super().init_by_input_data(input_data)
86 input_args.pop()97 input_args.pop()
87 output_packages[:] = [input_args[0]]98 output_packages[:] = [input_args[0]]
88- 99+ 
89 return input_args, output_packages100 return input_args, output_packages
90 101 
91 def after_call(self, output_packages):102 def after_call(self, output_packages):
@@ -94,4 +105,3 @@ class InplaceNormalAclnnApi(AclnnBaseApi):
94 output.append(self.acl_tensor_to_torch(output_pack))105 output.append(self.acl_tensor_to_torch(output_pack))
95 106 
96 return output107 return output
97- 
@@ -19,8 +19,9 @@ from atk.tasks.api_execute.base_api import BaseApi
19from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi19from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi
20from atk.tasks.dataset.base_dataset import OpsDataset20from atk.tasks.dataset.base_dataset import OpsDataset
21 21 
22+ 
22def normal_golden(torch_tensor, params):23def normal_golden(torch_tensor, params):
23- if tf.__version__ >= '2':24+ if tf.__version__ >= "2":
24 tf.compat.v1.enable_eager_execution()25 tf.compat.v1.enable_eager_execution()
25 seed = []26 seed = []
26 offset = [0]27 offset = [0]
@@ -30,7 +31,7 @@ def normal_golden(torch_tensor, params):
30 offset.append(params["offset"])31 offset.append(params["offset"])
31 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True32 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
32 if not is_contiguous:33 if not is_contiguous:
33- print("------------非连续操作-------------")34+ print("---- Non-contiguous case ----")
34 torch_tensor = torch.transpose(torch_tensor, 0, 1)35 torch_tensor = torch.transpose(torch_tensor, 0, 1)
35 matrix = torch_tensor36 matrix = torch_tensor
36 if params["dtype_input"][0] == torch.bfloat16:37 if params["dtype_input"][0] == torch.bfloat16:
@@ -40,10 +41,12 @@ def normal_golden(torch_tensor, params):
40 else:41 else:
41 dtype = tf.dtypes.float3242 dtype = tf.dtypes.float32
42 matrix_shape = list(matrix.shape)43 matrix_shape = list(matrix.shape)
43- normal_data = tf.raw_ops.StatelessRandomNormalV2(shape=matrix_shape, key=seed, counter=offset, alg=1)44+ normal_data = tf.raw_ops.StatelessRandomNormalV2(
45+ shape=matrix_shape, key=seed, counter=offset, alg=1
46+ )
44 mul_data = tf.multiply(normal_data, std)47 mul_data = tf.multiply(normal_data, std)
45 add_data = tf.add(mul_data, mean)48 add_data = tf.add(mul_data, mean)
46- if tf.__version__ >= '2':49+ if tf.__version__ >= "2":
47 output_data = tf.cast(add_data, dtype).numpy()50 output_data = tf.cast(add_data, dtype).numpy()
48 output_data = torch.from_numpy(output_data)51 output_data = torch.from_numpy(output_data)
49 else:52 else:
@@ -52,7 +55,8 @@ def normal_golden(torch_tensor, params):
52 output_data = torch.from_numpy(output_data_tf)55 output_data = torch.from_numpy(output_data_tf)
53 output_data = output_data.to(params["dtype_input"][0])56 output_data = output_data.to(params["dtype_input"][0])
54 return output_data57 return output_data
55- 58+ 
59+ 
56@register("ascend_aclnn_inplace_normal_tensor")60@register("ascend_aclnn_inplace_normal_tensor")
57class MethodAclnnInplaceNormalTensorApi(BaseApi):61class MethodAclnnInplaceNormalTensorApi(BaseApi):
58 def __init__(self, task_result: TaskResult):62 def __init__(self, task_result: TaskResult):
@@ -62,31 +66,38 @@ class MethodAclnnInplaceNormalTensorApi(BaseApi):
62 66 
63 def __call__(self, input_data: InputDataset, with_output: bool = False):67 def __call__(self, input_data: InputDataset, with_output: bool = False):
64 # 获取yaml中所需参数68 # 获取yaml中所需参数
65- self.tensor_ = input_data.kwargs['selfRef']69+ self.tensor_ = input_data.kwargs["selfRef"]
66- self.tensor_dtype_ = input_data.kwargs['selfRef'].dtype70+ self.tensor_dtype_ = input_data.kwargs["selfRef"].dtype
67- self.mean_ = input_data.kwargs['mean']71+ self.mean_ = input_data.kwargs["mean"]
68- self.std_ = input_data.kwargs['std']72+ self.std_ = input_data.kwargs["std"]
69- self.seed_ = input_data.kwargs['seedTensor'].cpu().numpy()73+ self.seed_ = input_data.kwargs["seedTensor"].cpu().numpy()
70- self.offset_ = input_data.kwargs['offsetTensor'].cpu().numpy()74+ self.offset_ = input_data.kwargs["offsetTensor"].cpu().numpy()
71- self.offset2_ = input_data.kwargs['offset']75+ self.offset2_ = input_data.kwargs["offset"]
76+ 
77+ params = {
78+ "mean": self.mean_,
79+ "std": self.std_,
80+ "seed": self.seed_[0],
81+ "offset": self.offset_[0] + self.offset2_,
82+ "is_contiguous": True,
83+ "dtype_input": [self.tensor_dtype_],
84+ }
72 85 
73- params = {"mean": self.mean_, "std": self.std_, "seed": self.seed_[0], "offset": self.offset_[0] + self.offset2_, "is_contiguous": True,
74- "dtype_input": [self.tensor_dtype_]}
75-
76 x = normal_golden(self.tensor_, params)86 x = normal_golden(self.tensor_, params)
77- 87+ 
78 return x88 return x
79- 89+ 
90+ 
80@register("aclnn_inplace_normal_tensor")91@register("aclnn_inplace_normal_tensor")
81class InplaceNormalTensorAclnnApi(AclnnBaseApi):92class InplaceNormalTensorAclnnApi(AclnnBaseApi):
82 def __call__(self):93 def __call__(self):
83 super().__call__()94 super().__call__()
84- 95+ 
85 def init_by_input_data(self, input_data: InputDataset):96 def init_by_input_data(self, input_data: InputDataset):
86 input_args, output_packages = super().init_by_input_data(input_data)97 input_args, output_packages = super().init_by_input_data(input_data)
87 input_args.pop()98 input_args.pop()
88 output_packages[:] = [input_args[0]]99 output_packages[:] = [input_args[0]]
89- 100+ 
90 return input_args, output_packages101 return input_args, output_packages
91 102 
92 def after_call(self, output_packages):103 def after_call(self, output_packages):
@@ -95,4 +106,3 @@ class InplaceNormalTensorAclnnApi(AclnnBaseApi):
95 output.append(self.acl_tensor_to_torch(output_pack))106 output.append(self.acl_tensor_to_torch(output_pack))
96 107 
97 return output108 return output
98- 
@@ -118,10 +118,10 @@ static Status ParseOpToGraphRandomuniform(const ge::Operator& op, ge::Graph& gra
118 }118 }
119 auto data0 = op::Const((prop.ori_name + "_data0").c_str()).set_attr_value(prop.shape);119 auto data0 = op::Const((prop.ori_name + "_data0").c_str()).set_attr_value(prop.shape);
120 // cast output to dst_dtype(onnx : Ascend)120 // cast output to dst_dtype(onnx : Ascend)
121- // float32, float16, int32, int64121+ // float32, float16, int32, uint8
122 std::map<int, int> kvlist = {{1, 0}, {10, 1}, {6, 3}, {2, 9}};122 std::map<int, int> kvlist = {{1, 0}, {10, 1}, {6, 3}, {2, 9}};
123 if (kvlist.find(prop.dtype) == kvlist.end()) {123 if (kvlist.find(prop.dtype) == kvlist.end()) {
124- OP_LOGE(GetOpName(op).c_str(), "only support float32/float16/int32/int64, but got %d", prop.dtype);124+ OP_LOGE(GetOpName(op).c_str(), "only float32/float16/int32/uint8 are supported, but got %d", prop.dtype);
125 return FAILED;125 return FAILED;
126 }126 }
127 ge::DataType temp_type = GetOmDtypeFromOnnxDtype(prop.dtype);127 ge::DataType temp_type = GetOmDtypeFromOnnxDtype(prop.dtype);
@@ -56,7 +56,8 @@ static aclnnStatus updateFrom(int64_t& from, op::DataType dtype)
56 digits = DOUBLE_DIGITS;56 digits = DOUBLE_DIGITS;
57 break;57 break;
58 default:58 default:
59- OP_LOGI("dtype must be bfloat16, float16, float32 or double.");59+ OP_LOGI("dtype must be bfloat16, float16, float32 or double, actual dtype: %d.",
60+ static_cast<int32_t>(dtype));
60 return ACLNN_SUCCESS;61 return ACLNN_SUCCESS;
61 }62 }
62 if (fromPlusOne < from) {63 if (fromPlusOne < from) {
@@ -103,7 +104,8 @@ static aclnnStatus updateTo(int64_t& to, op::DataType dtype)
103 digits = DOUBLE_DIGITS;104 digits = DOUBLE_DIGITS;
104 break;105 break;
105 default:106 default:
106- OP_LOGI("dtype must be bfloat16, float16, float32 or double.");107+ OP_LOGI("dtype must be bfloat16, float16, float32 or double, actual dtype: %d.",
108+ static_cast<int32_t>(dtype));
107 return ACLNN_SUCCESS;109 return ACLNN_SUCCESS;
108 }110 }
109 if (toMinusOne >= to) {111 if (toMinusOne >= to) {
@@ -53,7 +53,7 @@ static bool CheckDtypeValid(const aclTensor* self)
53 // 如果soc是310系列芯片,则不支持DT_BF16,需要校验拦截53 // 如果soc是310系列芯片,则不支持DT_BF16,需要校验拦截
54 if (!CheckSocVersionIsSupportBf16() && (self->GetDataType() == op::DataType::DT_BF16)) {54 if (!CheckSocVersionIsSupportBf16() && (self->GetDataType() == op::DataType::DT_BF16)) {
55 OP_LOGE(ACLNN_ERR_PARAM_INVALID,55 OP_LOGE(ACLNN_ERR_PARAM_INVALID,
56- "Input dtype of aclnnInplaceUniform is not support bfloat16 in current socversion.");56+ "Input dtype of aclnnInplaceUniform does not support bfloat16 in the current soc version.");
57 return false;57 return false;
58 }58 }
59 59 
@@ -102,7 +102,7 @@ static aclScalar* CreateScalar(float input, op::DataType dtype, aclOpExecutor* e
102 ratioBf16 = input;102 ratioBf16 = input;
103 return executor->AllocScalar(&ratioBf16.value, op::DataType::DT_BF16);103 return executor->AllocScalar(&ratioBf16.value, op::DataType::DT_BF16);
104 default:104 default:
105- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "invalid dtype, must be bfloat16 or float16.");105+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "invalid dtype %d, must be bfloat16 or float16.", static_cast<int>(dtype));
106 return nullptr;106 return nullptr;
107 }107 }
108}108}
@@ -30,11 +30,12 @@ type2digits[torch.bfloat16] = BF16_DIGITS
30type2digits[torch.float32] = FLOAT32_DIGITS30type2digits[torch.float32] = FLOAT32_DIGITS
31type2digits[torch.float64] = DOUBLE_DIGITS31type2digits[torch.float64] = DOUBLE_DIGITS
32 32 
33+ 
33def update_from(from_value, scalar_type=torch.float32):34def update_from(from_value, scalar_type=torch.float32):
34 # 判断是否为浮点数数据类型35 # 判断是否为浮点数数据类型
35 if not scalar_type.is_floating_point:36 if not scalar_type.is_floating_point:
36 return from_value37 return from_value
37- 38+ 
38 tmp_value = torch.tensor(from_value + 1, dtype=torch.int64)39 tmp_value = torch.tensor(from_value + 1, dtype=torch.int64)
39 from_plus_1 = tmp_value.to(scalar_type).to(torch.int64)40 from_plus_1 = tmp_value.to(scalar_type).to(torch.int64)
40 if from_plus_1 < from_value:41 if from_plus_1 < from_value:
@@ -44,17 +45,18 @@ def update_from(from_value, scalar_type=torch.float32):
44 while from_:45 while from_:
45 n += 146 n += 1
46 from_ >>= 147 from_ >>= 1
47- digits = type2digits[scalar_type] # 获取浮点类型的尾数位数48+ digits = type2digits[scalar_type] # 获取浮点类型的尾数位数
48 adjustment = 1 << (n - digits + 1)49 adjustment = 1 << (n - digits + 1)
49 from_value = int(from_plus_1 + adjustment)50 from_value = int(from_plus_1 + adjustment)
50- 51+ 
51 return from_value52 return from_value
52 53 
54+ 
53def update_to(to_value, scalar_type=torch.float32):55def update_to(to_value, scalar_type=torch.float32):
54 # 判断是否为浮点数数据类型56 # 判断是否为浮点数数据类型
55 if not scalar_type.is_floating_point:57 if not scalar_type.is_floating_point:
56 return to_value58 return to_value
57- 59+ 
58 tmp_value = torch.tensor(to_value - 1, dtype=torch.int64)60 tmp_value = torch.tensor(to_value - 1, dtype=torch.int64)
59 to_minus_1 = tmp_value.to(scalar_type).to(torch.int64)61 to_minus_1 = tmp_value.to(scalar_type).to(torch.int64)
60 if to_minus_1 >= to_value:62 if to_minus_1 >= to_value:
@@ -64,12 +66,13 @@ def update_to(to_value, scalar_type=torch.float32):
64 while to_:66 while to_:
65 n += 167 n += 1
66 to_ >>= 168 to_ >>= 1
67- digits = type2digits[scalar_type] # 获取浮点类型的尾数位数69+ digits = type2digits[scalar_type] # 获取浮点类型的尾数位数
68 adjustment = 1 << (n - digits + 1)70 adjustment = 1 << (n - digits + 1)
69 to_value = int(to_minus_1 - adjustment)71 to_value = int(to_minus_1 - adjustment)
70- 72+ 
71 return to_value73 return to_value
72 74 
75+ 
73def random_golden(torch_tensor, params):76def random_golden(torch_tensor, params):
74 seed = []77 seed = []
75 offset = [0]78 offset = [0]
@@ -78,18 +81,20 @@ def random_golden(torch_tensor, params):
78 seed.append(params["seed"])81 seed.append(params["seed"])
79 offset.append(params["offset"])82 offset.append(params["offset"])
80 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True83 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
81- 84+ 
82 start = update_from(start, params["dtype_input"][0])85 start = update_from(start, params["dtype_input"][0])
83 end = update_to(end, params["dtype_input"][0])86 end = update_to(end, params["dtype_input"][0])
84 if not is_contiguous:87 if not is_contiguous:
85- print("------------非连续操作-------------")88+ print("---- Non-contiguous case ----")
86 torch_tensor = torch.transpose(torch_tensor, 0, 1)89 torch_tensor = torch.transpose(torch_tensor, 0, 1)
87 matrix = torch_tensor90 matrix = torch_tensor
88 if params["dtype_input"][0] == torch.bfloat16:91 if params["dtype_input"][0] == torch.bfloat16:
89 matrix = matrix.float()92 matrix = matrix.float()
90 matrix = tf.constant(matrix)93 matrix = tf.constant(matrix)
91 matrix_shape = list(matrix.shape)94 matrix_shape = list(matrix.shape)
92- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1)95+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
96+ shape=matrix_shape, key=seed, counter=offset, alg=1
97+ )
93 mul_data = tf.multiply(uniform_data, (end - start))98 mul_data = tf.multiply(uniform_data, (end - start))
94 add_data = tf.add(mul_data, start)99 add_data = tf.add(mul_data, start)
95 if params["dtype_input"][0] == torch.bool:100 if params["dtype_input"][0] == torch.bool:
@@ -99,7 +104,8 @@ def random_golden(torch_tensor, params):
99 output_data = torch.from_numpy(output_data)104 output_data = torch.from_numpy(output_data)
100 output_data = output_data.to(dtype=params["dtype_input"][0])105 output_data = output_data.to(dtype=params["dtype_input"][0])
101 return output_data106 return output_data
102- 107+ 
108+ 
103@register("ascend_aclnn_inplace_random")109@register("ascend_aclnn_inplace_random")
104class MethodAclnnInplaceRandomApi(BaseApi):110class MethodAclnnInplaceRandomApi(BaseApi):
105 def __init__(self, task_result: TaskResult):111 def __init__(self, task_result: TaskResult):
@@ -109,35 +115,41 @@ class MethodAclnnInplaceRandomApi(BaseApi):
109 115 
110 def __call__(self, input_data: InputDataset, with_output: bool = False):116 def __call__(self, input_data: InputDataset, with_output: bool = False):
111 # 获取yaml中所需参数117 # 获取yaml中所需参数
112- self.Tensor = input_data.kwargs['selfRef']118+ self.Tensor = input_data.kwargs["selfRef"]
113- self.Tensor_dtype = input_data.kwargs['selfRef'].dtype119+ self.Tensor_dtype = input_data.kwargs["selfRef"].dtype
114- self.from_ = input_data.kwargs['from']120+ self.from_ = input_data.kwargs["from"]
115- self.to_ = input_data.kwargs['to']121+ self.to_ = input_data.kwargs["to"]
116- self.seed = input_data.kwargs['seed']122+ self.seed = input_data.kwargs["seed"]
117- self.offset = input_data.kwargs['offset']123+ self.offset = input_data.kwargs["offset"]
124+ 
125+ params = {
126+ "from": self.from_,
127+ "to": self.to_,
128+ "seed": self.seed,
129+ "offset": self.offset,
130+ "is_contiguous": True,
131+ "dtype_input": [self.Tensor_dtype],
132+ }
118 133 
119- params = {"from": self.from_, "to": self.to_, "seed": self.seed, "offset": self.offset, "is_contiguous": True,
120- "dtype_input": [self.Tensor_dtype]}
121-
122 x = random_golden(self.Tensor, params)134 x = random_golden(self.Tensor, params)
123- 135+ 
124 return x136 return x
125- 137+ 
138+ 
126@register("aclnn_inplace_random")139@register("aclnn_inplace_random")
127class InplaceRandomAclnnApi(AclnnBaseApi):140class InplaceRandomAclnnApi(AclnnBaseApi):
128 def __call__(self):141 def __call__(self):
129 super().__call__()142 super().__call__()
130- 143+ 
131 def init_by_input_data(self, input_data: InputDataset):144 def init_by_input_data(self, input_data: InputDataset):
132 input_args, output_packages = super().init_by_input_data(input_data)145 input_args, output_packages = super().init_by_input_data(input_data)
133 input_args.pop()146 input_args.pop()
134 output_packages[:] = [input_args[0]]147 output_packages[:] = [input_args[0]]
135- 148+ 
136 return input_args, output_packages149 return input_args, output_packages
137- 150+ 
138 def after_call(self, output_packages):151 def after_call(self, output_packages):
139 output = []152 output = []
140 for output_pack in output_packages:153 for output_pack in output_packages:
141 output.append(self.acl_tensor_to_torch(output_pack))154 output.append(self.acl_tensor_to_torch(output_pack))
142 return output155 return output
143- 
@@ -30,11 +30,12 @@ type2digits[torch.bfloat16] = BF16_DIGITS
30type2digits[torch.float32] = FLOAT32_DIGITS30type2digits[torch.float32] = FLOAT32_DIGITS
31type2digits[torch.float64] = DOUBLE_DIGITS31type2digits[torch.float64] = DOUBLE_DIGITS
32 32 
33+ 
33def update_from(from_value, scalar_type=torch.float32):34def update_from(from_value, scalar_type=torch.float32):
34 # 判断是否为浮点数数据类型35 # 判断是否为浮点数数据类型
35 if not scalar_type.is_floating_point:36 if not scalar_type.is_floating_point:
36 return from_value37 return from_value
37- 38+ 
38 tmp_value = torch.tensor(from_value + 1, dtype=torch.int64)39 tmp_value = torch.tensor(from_value + 1, dtype=torch.int64)
39 from_plus_1 = tmp_value.to(scalar_type).to(torch.int64)40 from_plus_1 = tmp_value.to(scalar_type).to(torch.int64)
40 if from_plus_1 < from_value:41 if from_plus_1 < from_value:
@@ -44,17 +45,18 @@ def update_from(from_value, scalar_type=torch.float32):
44 while from_:45 while from_:
45 n += 146 n += 1
46 from_ >>= 147 from_ >>= 1
47- digits = type2digits[scalar_type] # 获取浮点类型的尾数位数48+ digits = type2digits[scalar_type] # 获取浮点类型的尾数位数
48 adjustment = 1 << (n - digits + 1)49 adjustment = 1 << (n - digits + 1)
49 from_value = int(from_plus_1 + adjustment)50 from_value = int(from_plus_1 + adjustment)
50- 51+ 
51 return from_value52 return from_value
52 53 
54+ 
53def update_to(to_value, scalar_type=torch.float32):55def update_to(to_value, scalar_type=torch.float32):
54 # 判断是否为浮点数数据类型56 # 判断是否为浮点数数据类型
55 if not scalar_type.is_floating_point:57 if not scalar_type.is_floating_point:
56 return to_value58 return to_value
57- 59+ 
58 tmp_value = torch.tensor(to_value - 1, dtype=torch.int64)60 tmp_value = torch.tensor(to_value - 1, dtype=torch.int64)
59 to_minus_1 = tmp_value.to(scalar_type).to(torch.int64)61 to_minus_1 = tmp_value.to(scalar_type).to(torch.int64)
60 if to_minus_1 >= to_value:62 if to_minus_1 >= to_value:
@@ -64,12 +66,13 @@ def update_to(to_value, scalar_type=torch.float32):
64 while to_:66 while to_:
65 n += 167 n += 1
66 to_ >>= 168 to_ >>= 1
67- digits = type2digits[scalar_type] # 获取浮点类型的尾数位数69+ digits = type2digits[scalar_type] # 获取浮点类型的尾数位数
68 adjustment = 1 << (n - digits + 1)70 adjustment = 1 << (n - digits + 1)
69 to_value = int(to_minus_1 - adjustment)71 to_value = int(to_minus_1 - adjustment)
70- 72+ 
71 return to_value73 return to_value
72 74 
75+ 
73def random_golden(torch_tensor, params):76def random_golden(torch_tensor, params):
74 seed = []77 seed = []
75 offset = [0]78 offset = [0]
@@ -78,18 +81,20 @@ def random_golden(torch_tensor, params):
78 seed.append(params["seed"])81 seed.append(params["seed"])
79 offset.append(params["offset"])82 offset.append(params["offset"])
80 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True83 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
81- 84+ 
82 start = update_from(start, params["dtype_input"][0])85 start = update_from(start, params["dtype_input"][0])
83 end = update_to(end, params["dtype_input"][0])86 end = update_to(end, params["dtype_input"][0])
84 if not is_contiguous:87 if not is_contiguous:
85- print("------------非连续操作-------------")88+ print("---- Non-contiguous case ----")
86 torch_tensor = torch.transpose(torch_tensor, 0, 1)89 torch_tensor = torch.transpose(torch_tensor, 0, 1)
87 matrix = torch_tensor90 matrix = torch_tensor
88 if params["dtype_input"][0] == torch.bfloat16:91 if params["dtype_input"][0] == torch.bfloat16:
89 matrix = matrix.float()92 matrix = matrix.float()
90 matrix = tf.constant(matrix)93 matrix = tf.constant(matrix)
91 matrix_shape = list(matrix.shape)94 matrix_shape = list(matrix.shape)
92- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1)95+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
96+ shape=matrix_shape, key=seed, counter=offset, alg=1
97+ )
93 mul_data = tf.multiply(uniform_data, (end - start))98 mul_data = tf.multiply(uniform_data, (end - start))
94 add_data = tf.add(mul_data, start)99 add_data = tf.add(mul_data, start)
95 if params["dtype_input"][0] == torch.bool:100 if params["dtype_input"][0] == torch.bool:
@@ -99,7 +104,8 @@ def random_golden(torch_tensor, params):
99 output_data = torch.from_numpy(output_data)104 output_data = torch.from_numpy(output_data)
100 output_data = output_data.to(dtype=params["dtype_input"][0])105 output_data = output_data.to(dtype=params["dtype_input"][0])
101 return output_data106 return output_data
102- 107+ 
108+ 
103@register("ascend_aclnn_inplace_random_tensor")109@register("ascend_aclnn_inplace_random_tensor")
104class MethodAclnnInplaceRandomTensorApi(BaseApi):110class MethodAclnnInplaceRandomTensorApi(BaseApi):
105 def __init__(self, task_result: TaskResult):111 def __init__(self, task_result: TaskResult):
@@ -109,36 +115,42 @@ class MethodAclnnInplaceRandomTensorApi(BaseApi):
109 115 
110 def __call__(self, input_data: InputDataset, with_output: bool = False):116 def __call__(self, input_data: InputDataset, with_output: bool = False):
111 # 获取yaml中所需参数117 # 获取yaml中所需参数
112- self.Tensor = input_data.kwargs['selfRef']118+ self.Tensor = input_data.kwargs["selfRef"]
113- self.Tensor_dtype = input_data.kwargs['selfRef'].dtype119+ self.Tensor_dtype = input_data.kwargs["selfRef"].dtype
114- self.from_ = input_data.kwargs['from']120+ self.from_ = input_data.kwargs["from"]
115- self.to_ = input_data.kwargs['to']121+ self.to_ = input_data.kwargs["to"]
116- self.seed = input_data.kwargs['seedTensor'].cpu().numpy()122+ self.seed = input_data.kwargs["seedTensor"].cpu().numpy()
117- self.offset = input_data.kwargs['offsetTensor'].cpu().numpy()123+ self.offset = input_data.kwargs["offsetTensor"].cpu().numpy()
118- self.offset2 = input_data.kwargs['offset']124+ self.offset2 = input_data.kwargs["offset"]
125+ 
126+ params = {
127+ "from": self.from_,
128+ "to": self.to_,
129+ "seed": self.seed[0],
130+ "offset": self.offset[0] + self.offset2,
131+ "is_contiguous": True,
132+ "dtype_input": [self.Tensor_dtype],
133+ }
119 134 
120- params = {"from": self.from_, "to": self.to_, "seed": self.seed[0], "offset": self.offset[0] + self.offset2, "is_contiguous": True,
121- "dtype_input": [self.Tensor_dtype]}
122-
123 x = random_golden(self.Tensor, params)135 x = random_golden(self.Tensor, params)
124- 136+ 
125 return x137 return x
126- 138+ 
139+ 
127@register("aclnn_inplace_random_tensor")140@register("aclnn_inplace_random_tensor")
128class InplaceRandomTensorAclnnApi(AclnnBaseApi):141class InplaceRandomTensorAclnnApi(AclnnBaseApi):
129 def __call__(self):142 def __call__(self):
130 super().__call__()143 super().__call__()
131- 144+ 
132 def init_by_input_data(self, input_data: InputDataset):145 def init_by_input_data(self, input_data: InputDataset):
133 input_args, output_packages = super().init_by_input_data(input_data)146 input_args, output_packages = super().init_by_input_data(input_data)
134 input_args.pop()147 input_args.pop()
135 output_packages[:] = [input_args[0]]148 output_packages[:] = [input_args[0]]
136- 149+ 
137 return input_args, output_packages150 return input_args, output_packages
138- 151+ 
139 def after_call(self, output_packages):152 def after_call(self, output_packages):
140 output = []153 output = []
141 for output_pack in output_packages:154 for output_pack in output_packages:
142 output.append(self.acl_tensor_to_torch(output_pack))155 output.append(self.acl_tensor_to_torch(output_pack))
143 return output156 return output
144- 
@@ -11,7 +11,6 @@
11# ----------------------------------------------------------------------------11# ----------------------------------------------------------------------------
12 12 
13import torch13import torch
14-import numpy as np
15import tensorflow as tf14import tensorflow as tf
16 15 
17from atk.configs.dataset_config import InputDataset16from atk.configs.dataset_config import InputDataset
@@ -21,6 +20,7 @@ from atk.tasks.api_execute.base_api import BaseApi
21from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi20from atk.tasks.api_execute.aclnn_base_api import AclnnBaseApi
22from atk.tasks.dataset.base_dataset import OpsDataset21from atk.tasks.dataset.base_dataset import OpsDataset
23 22 
23+ 
24def uniform_golden(torch_tensor, params):24def uniform_golden(torch_tensor, params):
25 seed = []25 seed = []
26 offset = [0]26 offset = [0]
@@ -30,7 +30,7 @@ def uniform_golden(torch_tensor, params):
30 offset.append(params["offset"])30 offset.append(params["offset"])
31 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True31 is_contiguous = params["is_contiguous"] if "is_contiguous" in params else True
32 if not is_contiguous:32 if not is_contiguous:
33- print("------------非连续操作-------------")33+ print("---- Non-contiguous case ----")
34 torch_tensor = torch.transpose(torch_tensor, 0, 1)34 torch_tensor = torch.transpose(torch_tensor, 0, 1)
35 matrix = torch_tensor35 matrix = torch_tensor
36 if params["dtype_input"][0] == torch.bfloat16:36 if params["dtype_input"][0] == torch.bfloat16:
@@ -40,37 +40,46 @@ def uniform_golden(torch_tensor, params):
40 matrix_shape = list(matrix.shape)40 matrix_shape = list(matrix.shape)
41 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])41 # print("***********************params[\"dtype_input\"][0]=%s****************************" % params["dtype_input"][0])
42 if params["dtype_input"][0] == torch.bfloat16:42 if params["dtype_input"][0] == torch.bfloat16:
43- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1,43+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
44- dtype=tf.dtypes.bfloat16)44+ shape=matrix_shape,
45+ key=seed,
46+ counter=offset,
47+ alg=1,
48+ dtype=tf.dtypes.bfloat16,
49+ )
45 end = tf.constant(end)50 end = tf.constant(end)
46 end = tf.cast(end, tf.dtypes.bfloat16)51 end = tf.cast(end, tf.dtypes.bfloat16)
47 start = tf.constant(start)52 start = tf.constant(start)
48 start = tf.cast(start, tf.dtypes.bfloat16)53 start = tf.cast(start, tf.dtypes.bfloat16)
49 elif params["dtype_input"][0] == torch.float16:54 elif params["dtype_input"][0] == torch.float16:
50- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1,55+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
51- dtype=tf.dtypes.float16)56+ shape=matrix_shape, key=seed, counter=offset, alg=1, dtype=tf.dtypes.float16
57+ )
52 end = tf.constant(end)58 end = tf.constant(end)
53 end = tf.cast(end, tf.dtypes.float16)59 end = tf.cast(end, tf.dtypes.float16)
54 start = tf.constant(start)60 start = tf.constant(start)
55 start = tf.cast(start, tf.dtypes.float16)61 start = tf.cast(start, tf.dtypes.float16)
56 else:62 else:
57- uniform_data = tf.raw_ops.StatelessRandomUniformV2(shape=matrix_shape, key=seed, counter=offset, alg=1)63+ uniform_data = tf.raw_ops.StatelessRandomUniformV2(
64+ shape=matrix_shape, key=seed, counter=offset, alg=1
65+ )
58 end = tf.constant(end)66 end = tf.constant(end)
59 end = tf.cast(end, tf.dtypes.float32)67 end = tf.cast(end, tf.dtypes.float32)
60 start = tf.constant(start)68 start = tf.constant(start)
61 start = tf.cast(start, tf.dtypes.float32)69 start = tf.cast(start, tf.dtypes.float32)
62 mul_data = tf.multiply(uniform_data, (end - start))70 mul_data = tf.multiply(uniform_data, (end - start))
63 add_data = tf.add(mul_data, start)71 add_data = tf.add(mul_data, start)
64- 72+ 
65 output_data = tf.cast(add_data, dtype)73 output_data = tf.cast(add_data, dtype)
66 if output_data.shape == []:74 if output_data.shape == []:
67 output_data = torch.tensor(output_data.numpy())75 output_data = torch.tensor(output_data.numpy())
68 else:76 else:
69 output_data = torch.from_numpy(output_data.numpy())77 output_data = torch.from_numpy(output_data.numpy())
70- 78+ 
71 output_data = output_data.type(params["dtype_input"][0])79 output_data = output_data.type(params["dtype_input"][0])
72 return output_data80 return output_data
73- 81+ 
82+ 
74@register("ascend_aclnn_inplace_uniform")83@register("ascend_aclnn_inplace_uniform")
75class MethodAclnnInplaceUniformApi(BaseApi):84class MethodAclnnInplaceUniformApi(BaseApi):
76 def __init__(self, task_result: TaskResult):85 def __init__(self, task_result: TaskResult):
@@ -80,31 +89,38 @@ class MethodAclnnInplaceUniformApi(BaseApi):
80 89 
81 def __call__(self, input_data: InputDataset, with_output: bool = False):90 def __call__(self, input_data: InputDataset, with_output: bool = False):
82 # 获取yaml中所需参数91 # 获取yaml中所需参数
83- self.Tensor = input_data.kwargs['selfRef']92+ self.Tensor = input_data.kwargs["selfRef"]
84- self.Tensor_dtype = input_data.kwargs['selfRef'].dtype93+ self.Tensor_dtype = input_data.kwargs["selfRef"].dtype
85- self.from_ = input_data.kwargs['from']94+ self.from_ = input_data.kwargs["from"]
86- self.to_ = input_data.kwargs['to']95+ self.to_ = input_data.kwargs["to"]
87- self.seed = input_data.kwargs['seed']96+ self.seed = input_data.kwargs["seed"]
88- self.offset = input_data.kwargs['offset']97+ self.offset = input_data.kwargs["offset"]
98+ 
99+ params = {
100+ "from": self.from_,
101+ "to": self.to_,
102+ "seed": self.seed,
103+ "offset": self.offset,
104+ "is_contiguous": True,
105+ "dtype_input": [self.Tensor_dtype],
106+ }
89 107 
90- params = {"from": self.from_, "to": self.to_, "seed": self.seed, "offset": self.offset, "is_contiguous": True,
91- "dtype_input": [self.Tensor_dtype]}
92-
93 x = uniform_golden(self.Tensor, params)108 x = uniform_golden(self.Tensor, params)
94- 109+ 
95 return x110 return x
96- 111+ 
112+ 
97@register("aclnn_inplace_uniform")113@register("aclnn_inplace_uniform")
98class InplaceUniformAclnnApi(AclnnBaseApi):114class InplaceUniformAclnnApi(AclnnBaseApi):
99 def __call__(self):115 def __call__(self):
100 super().__call__()116 super().__call__()
101- 117+ 
102 def init_by_input_data(self, input_data: InputDataset):118 def init_by_input_data(self, input_data: InputDataset):
103 input_args, output_packages = super().init_by_input_data(input_data)119 input_args, output_packages = super().init_by_input_data(input_data)
104 input_args.pop()120 input_args.pop()
105 output_packages[:] = [input_args[0]]121 output_packages[:] = [input_args[0]]
106- 122+ 
107 return input_args, output_packages123 return input_args, output_packages
108- 124+ 
109 def after_call(self, output_packages):125 def after_call(self, output_packages):
110- return super().after_call(output_packages)126+ return super().after_call(output_packages)
@@ -15,8 +15,8 @@
15#include "random_graph_infer_base.h"15#include "random_graph_infer_base.h"
16namespace ops {16namespace ops {
17namespace GraphCommon {17namespace GraphCommon {
18-ge::graphStatus InferDataTypeByAttr(18+ge::graphStatus InferDataTypeByAttr(gert::InferDataTypeContext* context, const int32_t dtypeIndex,
19- gert::InferDataTypeContext* context, const int32_t dtypeIndex, ge::DataType& OutDtype)19+ ge::DataType& OutDtype)
20{20{
21 auto* attrs = context->GetAttrs();21 auto* attrs = context->GetAttrs();
22 OP_CHECK_NULL_WITH_CONTEXT(context, attrs);22 OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
@@ -27,9 +27,9 @@ ge::graphStatus InferDataTypeByAttr(
27 return ge::GRAPH_SUCCESS;27 return ge::GRAPH_SUCCESS;
28}28}
29 29 
30-ge::graphStatus CommonInferType(30+ge::graphStatus CommonInferType(gert::InferDataTypeContext* context, int32_t mode, int32_t dtypeIndex,
31- gert::InferDataTypeContext* context, int32_t mode, int32_t dtypeIndex,31+ const std::vector<OutputSpec>& extraOutputMap,
32- const std::vector<OutputSpec>& extraOutputMap, const std::set<ge::DataType>& supportDtype, bool isCheck)32+ const std::set<ge::DataType>& supportDtype, bool isCheck)
33{33{
34 if (context == nullptr) {34 if (context == nullptr) {
35 OP_LOGE(context, "Null context pointer");35 OP_LOGE(context, "Null context pointer");
@@ -60,10 +60,9 @@ ge::graphStatus CommonInferType(
60 return ge::GRAPH_FAILED;60 return ge::GRAPH_FAILED;
61 }61 }
62 62 
63- OP_CHECK_IF(63+ OP_CHECK_IF(isCheck && supportDtype.count(outDtype) == 0,
64- isCheck && supportDtype.count(outDtype) == 0,64+ OP_LOGE(context->GetNodeName(), "Unsupported dtype: %s", Ops::Base::ToString(outDtype).c_str()),
65- OP_LOGE(context->GetNodeName(), "Unsupported dtype: %s", Ops::Base::ToString(outDtype).c_str()),65+ return ge::GRAPH_FAILED);
66- return ge::GRAPH_FAILED);
67 66 
68 context->SetOutputDataType(0, outDtype);67 context->SetOutputDataType(0, outDtype);
69 68 
@@ -73,7 +72,7 @@ ge::graphStatus CommonInferType(
73 context->SetOutputDataType(extraOutputIndex, extraOutputType);72 context->SetOutputDataType(extraOutputIndex, extraOutputType);
74 }73 }
75 74 
76- OP_LOGD(context->GetNodeName(), "END to do infer data type.");75+ OP_LOGD(context->GetNodeName(), "End inferring data type.");
77 return ge::GRAPH_SUCCESS;76 return ge::GRAPH_SUCCESS;
78}77}
79} // namespace GraphCommon78} // namespace GraphCommon
@@ -224,7 +224,7 @@ ge::graphStatus CalcExecutionPoliciesForBlocks(RandomUnifiedSimtTilingDataStruct
224 224 
225ge::graphStatus RandomTilingParseArch35(gert::TilingParseContext* context, const std::string& operatorName)225ge::graphStatus RandomTilingParseArch35(gert::TilingParseContext* context, const std::string& operatorName)
226{226{
227- OP_LOGD(context, "Entering RandomTilingArch35 operator name : %s", operatorName.c_str());227+ OP_LOGD(context, "Entering RandomTilingArch35 operator name: %s", operatorName.c_str());
228 auto compileInfo = context->GetCompiledInfo<RandomOperatorCompileInfo>();228 auto compileInfo = context->GetCompiledInfo<RandomOperatorCompileInfo>();
229 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);229 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
230 auto platformInfo = context->GetPlatformInfo();230 auto platformInfo = context->GetPlatformInfo();
@@ -237,7 +237,7 @@ ge::graphStatus RandomTilingParseArch35(gert::TilingParseContext* context, const
237 OP_LOGE(context, "GetHardwareInfo Failed, vectorCoreNum:%ld, ubSize:%ld.", compileInfo->totalCoreNum,237 OP_LOGE(context, "GetHardwareInfo Failed, vectorCoreNum:%ld, ubSize:%ld.", compileInfo->totalCoreNum,
238 compileInfo->ubSize),238 compileInfo->ubSize),
239 return ge::GRAPH_FAILED);239 return ge::GRAPH_FAILED);
240- OP_LOGD(context, "Get totalCoreNum:%d, ubSize:%ld", compileInfo->totalCoreNum, compileInfo->ubSize);240+ OP_LOGD(context, "Get totalCoreNum:%ld, ubSize:%ld", compileInfo->totalCoreNum, compileInfo->ubSize);
241 return ge::GRAPH_SUCCESS;241 return ge::GRAPH_SUCCESS;
242}242}
243 243 
@@ -273,7 +273,7 @@ ge::graphStatus ExtractTensorValue(const gert::TilingContext* context, const int
273 ret = GetIntValue<int64_t>(context, constTensor, constShape);273 ret = GetIntValue<int64_t>(context, constTensor, constShape);
274 break;274 break;
275 default:275 default:
276- OP_LOGD(context->GetNodeName(), "ExtractTensorValue only support [int32, int64]. but is %s",276+ OP_LOGD(context->GetNodeName(), "ExtractTensorValue only supports [int32, int64], but got %s",
277 Ops::Base::ToString(constDtype).c_str());277 Ops::Base::ToString(constDtype).c_str());
278 return ge::GRAPH_FAILED;278 return ge::GRAPH_FAILED;
279 }279 }
@@ -308,7 +308,7 @@ ge::graphStatus RandomTilingArch35::DoTiling()
308 // 步骤3:前置处理(可选)308 // 步骤3:前置处理(可选)
309 ret = BeforeProcess();309 ret = BeforeProcess();
310 if (ret != ge::GRAPH_SUCCESS) {310 if (ret != ge::GRAPH_SUCCESS) {
311- OP_LOGE(opName_, "Before process failed");311+ OP_LOGE(opName_, "Before process failed");
312 return ret;312 return ret;
313 }313 }
314 314 
@@ -329,7 +329,7 @@ ge::graphStatus RandomTilingArch35::DoTiling()
329 // 步骤6:后置处理(可选)329 // 步骤6:后置处理(可选)
330 ret = UniqueProcess();330 ret = UniqueProcess();
331 if (ret != ge::GRAPH_SUCCESS) {331 if (ret != ge::GRAPH_SUCCESS) {
332- OP_LOGE(opName_, "Unique process failed");332+ OP_LOGE(opName_, "Unique process failed");
333 return ret;333 return ret;
334 }334 }
335 335 
@@ -416,7 +416,7 @@ ge::graphStatus RandomTilingArch35::GetPlatformInfo()
416 } else {416 } else {
417 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);417 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
418 auto aivNum = ascendcPlatform.GetCoreNumAiv();418 auto aivNum = ascendcPlatform.GetCoreNumAiv();
419- OP_CHECK_IF((aivNum <= 0), OP_LOGE(opName_, "RandomTilingArch35 fail to get coreNum."),419+ OP_CHECK_IF((aivNum <= 0), OP_LOGE(opName_, "RandomTilingArch35 fails to get coreNum."),
420 return ge::GRAPH_FAILED);420 return ge::GRAPH_FAILED);
421 totalCoreNum_ = aivNum;421 totalCoreNum_ = aivNum;
422 uint64_t ubSizePlatForm = 0;422 uint64_t ubSizePlatForm = 0;
@@ -425,9 +425,10 @@ ge::graphStatus RandomTilingArch35::GetPlatformInfo()
425 }425 }
426 ubSize_ -= config_.DcacheSize;426 ubSize_ -= config_.DcacheSize;
427 427 
428- OP_CHECK_IF((ubSize_ <= 0), OP_LOGE(opName_, "ub size less than Dcache Size. please check."),428+ OP_CHECK_IF((ubSize_ <= 0),
429+ OP_LOGE(opName_, "ubSize %ld is less than Dcache size %ld, please check.", ubSize_, config_.DcacheSize),
429 return ge::GRAPH_FAILED);430 return ge::GRAPH_FAILED);
430- OP_LOGI(opName_, "RandomTilingArch35::GetPlatformInfo ubSize_=%d, totalCoreNum_=%d", ubSize_, totalCoreNum_);431+ OP_LOGI(opName_, "RandomTilingArch35::GetPlatformInfo ubSize_=%ld, totalCoreNum_=%ld", ubSize_, totalCoreNum_);
431 return ge::GRAPH_SUCCESS;432 return ge::GRAPH_SUCCESS;
432}433}
433 434 
@@ -435,7 +436,7 @@ ge::graphStatus RandomTilingArch35::DoSimtBlockTiling()
435{436{
436 OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),437 OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),
437 return ge::GRAPH_FAILED);438 return ge::GRAPH_FAILED);
438- OP_CHECK_IF((config_.coreAlignSize == 0), OP_LOGE(opName_, "coreAlignSize is equal to 0. please check."),439+ OP_CHECK_IF((config_.coreAlignSize == 0), OP_LOGE(opName_, "coreAlignSize is equal to 0. please check."),
439 return ge::GRAPH_FAILED);440 return ge::GRAPH_FAILED);
440 441 
441 int64_t avgPerCore = Ops::Base::CeilDiv(simtTilingData_.outputSize, totalCoreNum_);442 int64_t avgPerCore = Ops::Base::CeilDiv(simtTilingData_.outputSize, totalCoreNum_);
@@ -452,7 +453,8 @@ ge::graphStatus RandomTilingArch35::FillUnifiedSimtTilingData()
452 return ret;453 return ret;
453 }454 }
454 455 
455- OP_CHECK_IF((simtTilingData_.outputSize < 0), OP_LOGE(opName_, "outputSize is less than 0. please check."),456+ OP_CHECK_IF((simtTilingData_.outputSize < 0),
457+ OP_LOGE(opName_, "outputSize is %ld, must not be less than 0.", simtTilingData_.outputSize),
456 return ge::GRAPH_FAILED);458 return ge::GRAPH_FAILED);
457 ret = config_.getSeedAndOffset(context_, simtTilingData_.seed, simtTilingData_.offset);459 ret = config_.getSeedAndOffset(context_, simtTilingData_.seed, simtTilingData_.offset);
458 if (ret != ge::GRAPH_SUCCESS) {460 if (ret != ge::GRAPH_SUCCESS) {
@@ -497,7 +499,8 @@ ge::graphStatus RandomTilingArch35::FillUnifiedTilingData()
497 return ret;499 return ret;
498 }500 }
499 501 
500- OP_CHECK_IF((tilingData_.outputSize <= 0), OP_LOGE(opName_, "outputSize is less than or equal to 0. please check."),502+ OP_CHECK_IF((tilingData_.outputSize <= 0),
503+ OP_LOGE(opName_, "outputSize is %ld, must be greater than 0.", tilingData_.outputSize),
501 return ge::GRAPH_FAILED);504 return ge::GRAPH_FAILED);
502 ret = config_.getKeyAndCounter(context_, tilingData_.key, tilingData_.counter);505 ret = config_.getKeyAndCounter(context_, tilingData_.key, tilingData_.counter);
503 if (ret != ge::GRAPH_SUCCESS) {506 if (ret != ge::GRAPH_SUCCESS) {
@@ -532,7 +535,7 @@ ge::graphStatus RandomTilingArch35::DoBlockTiling()
532 OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),535 OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),
533 return ge::GRAPH_FAILED);536 return ge::GRAPH_FAILED);
534 tilingData_.normalCoreProNum = Ops::Base::CeilDiv(tilingData_.outputSize, totalCoreNum_);537 tilingData_.normalCoreProNum = Ops::Base::CeilDiv(tilingData_.outputSize, totalCoreNum_);
535- OP_CHECK_IF((config_.coreAlignSize == 0), OP_LOGE(opName_, "coreAlignSize is equal to 0. please check."),538+ OP_CHECK_IF((config_.coreAlignSize == 0), OP_LOGE(opName_, "coreAlignSize is equal to 0. please check."),
536 return ge::GRAPH_FAILED);539 return ge::GRAPH_FAILED);
537 tilingData_.normalCoreProNum = (tilingData_.normalCoreProNum + config_.coreAlignSize - 1) / config_.coreAlignSize *540 tilingData_.normalCoreProNum = (tilingData_.normalCoreProNum + config_.coreAlignSize - 1) / config_.coreAlignSize *
538 config_.coreAlignSize;541 config_.coreAlignSize;
@@ -612,7 +615,11 @@ ge::graphStatus RandomTilingArch35::CheckTensor(const gert::CompileTimeTensorDes
612 // 校验dtype615 // 校验dtype
613 if (!rule.dtypeSet.empty() && rule.dtypeSet.count(tensorDesc->GetDataType()) == 0) {616 if (!rule.dtypeSet.empty() && rule.dtypeSet.count(tensorDesc->GetDataType()) == 0) {
614 std::string valueStr = Ops::Base::ToString(tensorDesc->GetDataType());617 std::string valueStr = Ops::Base::ToString(tensorDesc->GetDataType());
615- std::string reasonMsg = "dtype not in allowed set";618+ std::string reasonMsg = "dtype not in allowed set [";
619+ for (auto t : rule.dtypeSet) {
620+ reasonMsg += Ops::Base::ToString(t) + ", ";
621+ }
622+ reasonMsg += "]";
616 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), tensorName.c_str(), valueStr.c_str(),623 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), tensorName.c_str(), valueStr.c_str(),
617 reasonMsg.c_str());624 reasonMsg.c_str());
618 return ge::GRAPH_FAILED;625 return ge::GRAPH_FAILED;
@@ -632,7 +639,11 @@ ge::graphStatus RandomTilingArch35::CheckTensor(const gert::CompileTimeTensorDes
632 auto dimNum = tensorShape.GetDimNum();639 auto dimNum = tensorShape.GetDimNum();
633 if (!rule.dimNumSet.empty() && rule.dimNumSet.count(dimNum) == 0) {640 if (!rule.dimNumSet.empty() && rule.dimNumSet.count(dimNum) == 0) {
634 std::string valueStr = std::to_string(dimNum);641 std::string valueStr = std::to_string(dimNum);
635- std::string reasonMsg = "dim num not in allowed set";642+ std::string reasonMsg = "dim num not in allowed set [";
643+ for (auto d : rule.dimNumSet) {
644+ reasonMsg += std::to_string(d) + ", ";
645+ }
646+ reasonMsg += "]";
636 OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context_->GetNodeName(), tensorName.c_str(), valueStr.c_str(),647 OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context_->GetNodeName(), tensorName.c_str(), valueStr.c_str(),
637 reasonMsg.c_str());648 reasonMsg.c_str());
638 return ge::GRAPH_FAILED;649 return ge::GRAPH_FAILED;
@@ -640,7 +651,7 @@ ge::graphStatus RandomTilingArch35::CheckTensor(const gert::CompileTimeTensorDes
640 651 
641 // 自定义校验652 // 自定义校验
642 if (rule.customCheck && !rule.customCheck(context_)) {653 if (rule.customCheck && !rule.customCheck(context_)) {
643- std::string reasonMsg = "custom check failed";654+ std::string reasonMsg = "custom check failed, please check the custom check function";
644 OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_->GetNodeName(), tensorName.c_str(), "failed", reasonMsg.c_str());655 OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context_->GetNodeName(), tensorName.c_str(), "failed", reasonMsg.c_str());
645 return ge::GRAPH_FAILED;656 return ge::GRAPH_FAILED;
646 }657 }
@@ -13,7 +13,7 @@
13 * \brief13 * \brief
14 * 使用示例:14 * 使用示例:
15 * \code15 * \code
16- * // 1. 准备输入数据 16+ * // 1. 准备输入数据
17 * std::vector<ge::DataType> types = {ge::DT_FLOAT, ge::DT_INT32};17 * std::vector<ge::DataType> types = {ge::DT_FLOAT, ge::DT_INT32};
18 * std::vector<ge::Format> fmts = {ge::FORMAT_NCHW, ge::FORMAT_NHWC};18 * std::vector<ge::Format> fmts = {ge::FORMAT_NCHW, ge::FORMAT_NHWC};
19 * // 2. 初始化生成器19 * // 2. 初始化生成器
@@ -34,7 +34,6 @@
34 34 
35#include <algorithm>35#include <algorithm>
36#include <cstdint>36#include <cstdint>
37-#include <iostream>
38#include <limits>37#include <limits>
39#include <stdexcept>38#include <stdexcept>
40#include <string>39#include <string>
@@ -42,51 +41,44 @@
42#include <unordered_map>41#include <unordered_map>
43#include <utility>42#include <utility>
44#include <vector>43#include <vector>
45-#include <iomanip>
46 44 
47namespace randomdef {45namespace randomdef {
48namespace detail {46namespace detail {
49-#define GE_CASE(VAL) case ge::VAL: return #VAL47+#define GE_CASE(VAL) \
48+ case ge::VAL: \
49+ return #VAL
50inline std::string TypeToStr(ge::DataType type)50inline std::string TypeToStr(ge::DataType type)
51{51{
52 switch (type) {52 switch (type) {
53- GE_CASE(DT_FLOAT); GE_CASE(DT_FLOAT16); GE_CASE(DT_BF16);53+ GE_CASE(DT_FLOAT);
54- GE_CASE(DT_INT8); GE_CASE(DT_INT16); GE_CASE(DT_INT32); GE_CASE(DT_INT64);54+ GE_CASE(DT_FLOAT16);
55- GE_CASE(DT_UINT8); GE_CASE(DT_UINT16); GE_CASE(DT_UINT32); GE_CASE(DT_UINT64);55+ GE_CASE(DT_BF16);
56+ GE_CASE(DT_INT8);
57+ GE_CASE(DT_INT16);
58+ GE_CASE(DT_INT32);
59+ GE_CASE(DT_INT64);
60+ GE_CASE(DT_UINT8);
61+ GE_CASE(DT_UINT16);
62+ GE_CASE(DT_UINT32);
63+ GE_CASE(DT_UINT64);
56 GE_CASE(DT_BOOL);64 GE_CASE(DT_BOOL);
57- default: return "DT_" + std::to_string(type);65+ default:
66+ return "DT_" + std::to_string(type);
58 }67 }
59}68}
60 69 
61inline std::string TypeToStr(ge::Format fmt)70inline std::string TypeToStr(ge::Format fmt)
62{71{
63 switch (fmt) {72 switch (fmt) {
64- GE_CASE(FORMAT_NCHW); GE_CASE(FORMAT_NHWC); GE_CASE(FORMAT_ND);73+ GE_CASE(FORMAT_NCHW);
65- default: return "FMT_" + std::to_string(fmt);74+ GE_CASE(FORMAT_NHWC);
75+ GE_CASE(FORMAT_ND);
76+ default:
77+ return "FMT_" + std::to_string(fmt);
66 }78 }
67}79}
68-#undef GE_CASE 80+#undef GE_CASE
69 81 
70-template <typename T>
71-inline void PrintByColsCore(const std::vector<T>& seq, const char* varName, size_t cols)
72-{
73- static constexpr int COL_WIDTH = 15;
74- if (cols == 0) {
75- std::cerr << "[Warning] PrintByColsCore: cols is 0, doing nothing." << std::endl;
76- return;
77- }
78- std::cout << ">>> Sequence '" << varName << "' (Total: " << seq.size() << ", Cols: " << cols << "):" << std::endl;
79- for (size_t i = 0; i < seq.size(); ++i) {
80- std::cout << std::left << std::setw(COL_WIDTH) << TypeToStr(seq[i]);
81- if ((i + 1) % cols == 0) {
82- std::cout << std::endl;
83- }
84- }
85- if (seq.size() % cols != 0) {
86- std::cout << std::endl;
87- }
88- std::cout << "------------------------------------------------------------" << std::endl;
89-}
90} // namespace detail82} // namespace detail
91 83 
92struct InputOption {84struct InputOption {
@@ -95,8 +87,8 @@ struct InputOption {
95 template <typename T = ge::DataType>87 template <typename T = ge::DataType>
96 InputOption(std::string n, const std::vector<T>& v) : name(std::move(n))88 InputOption(std::string n, const std::vector<T>& v) : name(std::move(n))
97 {89 {
98- static_assert(90+ static_assert(std::is_arithmetic<T>::value || std::is_enum<T>::value,
99- std::is_arithmetic<T>::value || std::is_enum<T>::value, "InputOption data must be arithmetic or enum");91+ "InputOption data must be arithmetic or enum");
100 data.reserve(v.size());92 data.reserve(v.size());
101 for (const auto& item : v) {93 for (const auto& item : v) {
102 data.push_back(static_cast<int64_t>(item));94 data.push_back(static_cast<int64_t>(item));
@@ -104,7 +96,7 @@ struct InputOption {
104 }96 }
105};97};
106 98 
107-class RandomDtypeFmtGen {99+class RandomDtypeFmtGen {
108public:100public:
109 static constexpr size_t MAX_COMBINATIONS_LIMIT = 100000000;101 static constexpr size_t MAX_COMBINATIONS_LIMIT = 100000000;
110 102 
@@ -142,8 +134,8 @@ public:
142 throw std::overflow_error("Total combinations overflow size_t");134 throw std::overflow_error("Total combinations overflow size_t");
143 }135 }
144 if (currentStride * currentSize > MAX_COMBINATIONS_LIMIT) {136 if (currentStride * currentSize > MAX_COMBINATIONS_LIMIT) {
145- throw std::length_error(137+ throw std::length_error("Total combinations exceed safety limit (" +
146- "Total combinations exceed safety limit (" + std::to_string(MAX_COMBINATIONS_LIMIT) + ")");138+ std::to_string(MAX_COMBINATIONS_LIMIT) + ")");
147 }139 }
148 currentStride *= currentSize;140 currentStride *= currentSize;
149 }141 }
@@ -164,13 +156,6 @@ public:
164 return result;156 return result;
165 }157 }
166 158 
167- template <typename T = ge::DataType>
168- void Print(const std::string& name, size_t cols = 6) const
169- {
170- std::vector<T> seq = GetSequence<T>(name);
171- detail::PrintByColsCore<T>(seq, name.c_str(), cols);
172- }
173- 
174private:159private:
175 std::vector<std::vector<int64_t>> rawInputs_;160 std::vector<std::vector<int64_t>> rawInputs_;
176 std::vector<std::string> names_;161 std::vector<std::string> names_;
@@ -19,7 +19,7 @@ namespace randomCommon {
19template <typename T>19template <typename T>
20ge::graphStatus HandleShapeTensor(gert::Shape& outputShape, size_t xShapeSize, const T* xShapeData)20ge::graphStatus HandleShapeTensor(gert::Shape& outputShape, size_t xShapeSize, const T* xShapeData)
21{21{
22- std::cerr << "[DEBUG] HandleShapeTensor with type: " << typeid(T).name() << ", dims " << xShapeSize << std::endl;22+ OP_LOGD("RandomInferShape", "HandleShapeTensor with type: %s, dims %ld", typeid(T).name(), xShapeSize);
23 outputShape.SetDimNum(xShapeSize);23 outputShape.SetDimNum(xShapeSize);
24 for (size_t i = 0U; i < xShapeSize; i++) {24 for (size_t i = 0U; i < xShapeSize; i++) {
25 outputShape.SetDim(i, xShapeData[i]);25 outputShape.SetDim(i, xShapeData[i]);
@@ -27,9 +27,8 @@ ge::graphStatus HandleShapeTensor(gert::Shape& outputShape, size_t xShapeSize, c
27 return ge::GRAPH_SUCCESS;27 return ge::GRAPH_SUCCESS;
28}28}
29 29 
30-bool InferShapeForUnknow(30+bool InferShapeForUnknow(gert::InferShapeContext* context, const gert::Shape& inShape, gert::Shape& outShape,
31- gert::InferShapeContext* context, const gert::Shape& inShape, gert::Shape& outShape, int64_t& maskIndex,31+ int64_t& maskIndex, int64_t& offsetIndex)
32- int64_t& offsetIndex)
33{32{
34 if (Ops::Base::IsUnknownRank(inShape)) {33 if (Ops::Base::IsUnknownRank(inShape)) {
35 Ops::Base::SetUnknownRank(outShape);34 Ops::Base::SetUnknownRank(outShape);
@@ -65,7 +64,7 @@ bool DependencyMode(const gert::Tensor* inTensor, gert::Shape& outShape, size_t
65 if (shapeDtype == ge::DT_INT32) {64 if (shapeDtype == ge::DT_INT32) {
66 auto xShapeData = inTensor->GetData<int32_t>();65 auto xShapeData = inTensor->GetData<int32_t>();
67 if (xShapeData == nullptr) {66 if (xShapeData == nullptr) {
68- std::cerr << "[WARN] Empty DT_INT32 shape tensor, set 0-dim output" << std::endl;67+ OP_LOGW("RandomInferShape", "Empty DT_INT32 shape tensor, set 0-dim output");
69 Ops::Base::SetUnknownShape(xShapeSize, outShape);68 Ops::Base::SetUnknownShape(xShapeSize, outShape);
70 return true;69 return true;
71 }70 }
@@ -75,7 +74,7 @@ bool DependencyMode(const gert::Tensor* inTensor, gert::Shape& outShape, size_t
75 } else if (shapeDtype == ge::DT_INT64) {74 } else if (shapeDtype == ge::DT_INT64) {
76 auto xShapeData = inTensor->GetData<int64_t>();75 auto xShapeData = inTensor->GetData<int64_t>();
77 if (xShapeData == nullptr) {76 if (xShapeData == nullptr) {
78- std::cerr << "[WARN] Empty DT_INT64 shape tensor, set 0-dim output" << std::endl;77+ OP_LOGW("RandomInferShape", "Empty DT_INT64 shape tensor, set 0-dim output");
79 Ops::Base::SetUnknownShape(xShapeSize, outShape);78 Ops::Base::SetUnknownShape(xShapeSize, outShape);
80 return true;79 return true;
81 }80 }
@@ -83,14 +82,14 @@ bool DependencyMode(const gert::Tensor* inTensor, gert::Shape& outShape, size_t
83 return true;82 return true;
84 }83 }
85 }84 }
86- std::cerr << "[ERROR] Unsupported dtype: " << static_cast<int>(shapeDtype) << std::endl;85+ OP_LOGE("RandomInferShape", "Unsupported dtype: %d", static_cast<int>(shapeDtype));
87 return false;86 return false;
88}87}
89 88 
90-bool InputAndOutputCheck(89+bool InputAndOutputCheck(gert::InferShapeContext* context,
91- gert::InferShapeContext* context, const std::unordered_map<std::string, size_t>& requiredInputMap,90+ const std::unordered_map<std::string, size_t>& requiredInputMap,
92- const std::unordered_map<std::string, size_t>& outputMap, int64_t& maskIndex, int64_t& offsetIndex,91+ const std::unordered_map<std::string, size_t>& outputMap, int64_t& maskIndex,
93- const std::unordered_map<std::string, size_t>& optionalInputMap)92+ int64_t& offsetIndex, const std::unordered_map<std::string, size_t>& optionalInputMap)
94{93{
95 OP_LOGD(context->GetNodeName(), "InputAndOutputCheck start");94 OP_LOGD(context->GetNodeName(), "InputAndOutputCheck start");
96 for (const auto& item : requiredInputMap) {95 for (const auto& item : requiredInputMap) {
@@ -119,15 +118,15 @@ bool InputAndOutputCheck(
119 offsetIndex = outputIndex;118 offsetIndex = outputIndex;
120 }119 }
121 }120 }
122- OP_LOGD(121+ OP_LOGD(context->GetNodeName(), "InputAndOutputCheck end, maskIndex = %ld, offsetIndex = %ld", maskIndex,
123- context->GetNodeName(), "InputAndOutputCheck end, maskIndex = %ld, offsetIndex = %ld", maskIndex, offsetIndex);122+ offsetIndex);
124 return true;123 return true;
125}124}
126 125 
127-ge::graphStatus CommonInferShape(126+ge::graphStatus CommonInferShape(gert::InferShapeContext* context,
128- gert::InferShapeContext* context, const std::unordered_map<std::string, size_t>& requiredInputMap,127+ const std::unordered_map<std::string, size_t>& requiredInputMap,
129- const std::unordered_map<std::string, size_t>& outputMap, int32_t mode,128+ const std::unordered_map<std::string, size_t>& outputMap, int32_t mode,
130- const std::unordered_map<std::string, size_t>& optionalInputMap)129+ const std::unordered_map<std::string, size_t>& optionalInputMap)
131{130{
132 if (context == nullptr) {131 if (context == nullptr) {
133 return ge::GRAPH_FAILED;132 return ge::GRAPH_FAILED;
@@ -36,87 +36,78 @@ using namespace ge;
36using std::map;36using std::map;
37using std::string;37using std::string;
38using std::vector;38using std::vector;
39-#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \39+#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
40- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \40+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
41- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \41+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
42- TensorDesc placeholder##intputIndex##_desc = \42+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
43- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \43+ intputDtype); \
44- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \44+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
45- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \45+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
46- Tensor tensor_placeholder##intputIndex; \46+ Tensor tensor_placeholder##intputIndex; \
47- ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, \47+ ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
48- tensor_placeholder##intputIndex, \48+ placeholder##intputIndex##_desc, value); \
49- placeholder##intputIndex##_desc, \49+ if (ret != SUCCESS) { \
50- value); \50+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
51- if (ret != SUCCESS) { \51+ return FAILED; \
52- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \52+ } \
53- return FAILED; \53+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
54- } \54+ input.push_back(tensor_placeholder##intputIndex); \
55- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \55+ graph.AddOp(placeholder##intputIndex); \
56- input.push_back(tensor_placeholder##intputIndex); \56+ add1.set_input_##intputName(placeholder##intputIndex); \
57- graph.AddOp(placeholder##intputIndex); \
58- add1.set_input_##intputName(placeholder##intputIndex); \
59 inputs.push_back(placeholder##intputIndex)57 inputs.push_back(placeholder##intputIndex)
60 58 
61-#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \59+#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
62- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \60+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
63- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \61+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
64- TensorDesc placeholder##intputIndex##_desc = \62+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
65- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \63+ intputDtype); \
66- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \64+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
67- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \65+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
68- Tensor tensor_placeholder##intputIndex; \66+ Tensor tensor_placeholder##intputIndex; \
69- ret = GenOnesData(placeholder##intputIndex##_shape, \67+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
70- tensor_placeholder##intputIndex, \68+ placeholder##intputIndex##_desc, intputDtype, value); \
71- placeholder##intputIndex##_desc, \69+ if (ret != SUCCESS) { \
72- intputDtype, \70+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
73- value); \71+ return FAILED; \
74- if (ret != SUCCESS) { \72+ } \
75- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \73+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
76- return FAILED; \74+ input.push_back(tensor_placeholder##intputIndex); \
77- } \75+ graph.AddOp(placeholder##intputIndex); \
78- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \76+ add1.set_input_##intputName(placeholder##intputIndex); \
79- input.push_back(tensor_placeholder##intputIndex); \
80- graph.AddOp(placeholder##intputIndex); \
81- add1.set_input_##intputName(placeholder##intputIndex); \
82 inputs.push_back(placeholder##intputIndex)77 inputs.push_back(placeholder##intputIndex)
83 78 
84-#define ADD_INPUT_ATTR(attrName, attrValue) \79+#define ADD_INPUT_ATTR(attrName, attrValue) add1.set_attr_##attrName(attrValue)
85- add1.set_attr_##attrName(attrValue)
86 80 
87-#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \81+#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \
88- TensorDesc outputName##outputIndex##_desc = \82+ TensorDesc outputName##outputIndex##_desc = TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
89- TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
90 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)83 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)
91 84 
92-#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \85+#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \
93- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \86+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
94- auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \87+ auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \
95- TensorDesc placeholder##intputIndex##_desc = \88+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
96- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \89+ intputDtype); \
97- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \90+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
98- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \91+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
99- Tensor tensor_placeholder##intputIndex; \92+ Tensor tensor_placeholder##intputIndex; \
100- ret = GenOnesData(placeholder##intputIndex##_shape, \93+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
101- tensor_placeholder##intputIndex, \94+ placeholder##intputIndex##_desc, intputDtype, 1); \
102- placeholder##intputIndex##_desc, \95+ if (ret != SUCCESS) { \
103- intputDtype, \96+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
104- 1); \97+ return FAILED; \
105- if (ret != SUCCESS) { \98+ } \
106- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \99+ placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \
107- return FAILED; \100+ \
108- } \101+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
109- placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \ 102+ graph.AddOp(placeholder##intputIndex); \
110- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \103+ add1.set_input_##intputName(placeholder##intputIndex); \
111- graph.AddOp(placeholder##intputIndex); \104+ add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
112- add1.set_input_##intputName(placeholder##intputIndex); \
113- add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
114 inputs.push_back(placeholder##intputIndex)105 inputs.push_back(placeholder##intputIndex)
115 106 
116-#define LOG_PRINT(message, ...) \107+#define LOG_PRINT(message, ...) \
117- do { \108+ do { \
118- printf(message, ##__VA_ARGS__); \109+ printf(message, ##__VA_ARGS__); \
119- } while (0)110+ } while (0)
120 111 
121string GetTime()112string GetTime()
122{113{
@@ -159,7 +150,7 @@ uint32_t GetDataTypeSize(DataType dt)
159 return dilation;150 return dilation;
160}151}
161 152 
162-int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, float value)153+int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, float value)
163{154{
164 input_tensor_desc.SetRealDimCnt(shapes.size());155 input_tensor_desc.SetRealDimCnt(shapes.size());
165 size_t size = 1;156 size_t size = 1;
@@ -168,18 +159,18 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorD
168 }159 }
169 uint32_t byteSizeFloat32 = 4;160 uint32_t byteSizeFloat32 = 4;
170 uint32_t data_len = size * byteSizeFloat32;161 uint32_t data_len = size * byteSizeFloat32;
171- float *pData = new (std::nothrow) float[size];162+ float* pData = new (std::nothrow) float[size];
172 163 
173 for (size_t i = 0; i < size; ++i) {164 for (size_t i = 0; i < size; ++i) {
174 *(pData + i) = value;165 *(pData + i) = value;
175 }166 }
176- input_tensor = Tensor(input_tensor_desc, (uint8_t *)pData, data_len);167+ input_tensor = Tensor(input_tensor_desc, (uint8_t*)pData, data_len);
177 delete[] pData;168 delete[] pData;
178 return SUCCESS;169 return SUCCESS;
179}170}
180 171 
181-int32_t GenOnesData(172+int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type,
182- vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, DataType data_type, int value)173+ int value)
183{174{
184 input_tensor_desc.SetRealDimCnt(shapes.size());175 input_tensor_desc.SetRealDimCnt(shapes.size());
185 size_t size = 1;176 size_t size = 1;
@@ -187,25 +178,25 @@ int32_t GenOnesData(
187 size *= shapes[i];178 size *= shapes[i];
188 }179 }
189 uint32_t data_len = size * GetDataTypeSize(data_type);180 uint32_t data_len = size * GetDataTypeSize(data_type);
190- int64_t *pData = new (std::nothrow) int64_t[size];181+ int64_t* pData = new (std::nothrow) int64_t[size];
191 for (uint32_t i = 0; i < size; ++i) {182 for (uint32_t i = 0; i < size; ++i) {
192 pData[i] = static_cast<int64_t>(value);183 pData[i] = static_cast<int64_t>(value);
193 }184 }
194- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);185+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
195 delete[] pData;186 delete[] pData;
196 return SUCCESS;187 return SUCCESS;
197}188}
198 189 
199-int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t *inputData)190+int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
200{191{
201- FILE *fp = fopen(bin_file.c_str(), "w");192+ FILE* fp = fopen(bin_file.c_str(), "w");
202 fwrite(inputData, sizeof(uint8_t), data_size, fp);193 fwrite(inputData, sizeof(uint8_t), data_size, fp);
203 fclose(fp);194 fclose(fp);
204 return SUCCESS;195 return SUCCESS;
205}196}
206 197 
207-int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vector<Operator> &inputs,198+int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs,
208- std::vector<Operator> &outputs, Graph &graph)199+ std::vector<Operator>& outputs, Graph& graph)
209{200{
210 Status ret = SUCCESS;201 Status ret = SUCCESS;
211 // 自定义代码:添加单算子定义到图中202 // 自定义代码:添加单算子定义到图中
@@ -219,7 +210,7 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
219 ADD_INPUT_ATTR(dtype, 0);210 ADD_INPUT_ATTR(dtype, 0);
220 ADD_INPUT_ATTR(seed, 10);211 ADD_INPUT_ATTR(seed, 10);
221 ADD_INPUT_ATTR(seed2, 5);212 ADD_INPUT_ATTR(seed2, 5);
222- 213+ 
223 ADD_OUTPUT(1, y, ge::DT_FLOAT, outShape);214 ADD_OUTPUT(1, y, ge::DT_FLOAT, outShape);
224 ADD_OUTPUT(2, offset, ge::DT_INT64, offsetShape);215 ADD_OUTPUT(2, offset, ge::DT_INT64, offsetShape);
225 216 
@@ -228,9 +219,9 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
228 return SUCCESS;219 return SUCCESS;
229}220}
230 221 
231-int main(int argc, char *argv[])222+int main(int argc, char* argv[])
232{223{
233- const char *graph_name = "tc_ge_irrun_test";224+ const char* graph_name = "tc_ge_irrun_test";
234 Graph graph(graph_name);225 Graph graph(graph_name);
235 std::vector<ge::Tensor> input;226 std::vector<ge::Tensor> input;
236 227 
@@ -238,7 +229,7 @@ int main(int argc, char *argv[])
238 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};229 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
239 Status ret = ge::GEInitialize(global_options);230 Status ret = ge::GEInitialize(global_options);
240 if (ret != SUCCESS) {231 if (ret != SUCCESS) {
241- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());232+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
242 return FAILED;233 return FAILED;
243 }234 }
244 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());235 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -247,7 +238,7 @@ int main(int argc, char *argv[])
247 std::vector<Operator> outputs{};238 std::vector<Operator> outputs{};
248 239 
249 std::cout << argv[1] << std::endl;240 std::cout << argv[1] << std::endl;
250- char *endptr;241+ char* endptr;
251 242 
252 DataType inDtype = DT_INT64;243 DataType inDtype = DT_INT64;
253 std::cout << inDtype << std::endl;244 std::cout << inDtype << std::endl;
@@ -266,7 +257,7 @@ int main(int argc, char *argv[])
266 257 
267 };258 };
268 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());259 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());
269- ge::Session *session = new Session(build_options);260+ ge::Session* session = new Session(build_options);
270 261 
271 if (session == nullptr) {262 if (session == nullptr) {
272 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());263 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());
@@ -289,7 +280,7 @@ int main(int argc, char *argv[])
289 std::vector<ge::Tensor> output;280 std::vector<ge::Tensor> output;
290 ret = session->RunGraph(graph_id, input, output);281 ret = session->RunGraph(graph_id, input, output);
291 if (ret != SUCCESS) {282 if (ret != SUCCESS) {
292- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());283+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
293 delete session;284 delete session;
294 GEFinalize();285 GEFinalize();
295 return FAILED;286 return FAILED;
@@ -300,23 +291,23 @@ int main(int argc, char *argv[])
300 for (int i = 0; i < input_num; i++) {291 for (int i = 0; i < input_num; i++) {
301 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;292 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;
302 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";293 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
303- uint8_t *input_data_i = input[i].GetData();294+ uint8_t* input_data_i = input[i].GetData();
304 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();295 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
305- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;296+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
306 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());297 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
307- WriteDataToFile((const char *)input_file.c_str(), data_size, input_data_i);298+ WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
308 }299 }
309 300 
310 int output_num = output.size();301 int output_num = output.size();
311 for (int i = 0; i < output_num; i++) {302 for (int i = 0; i < output_num; i++) {
312 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;303 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;
313 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";304 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
314- uint8_t *output_data_i = output[i].GetData();305+ uint8_t* output_data_i = output[i].GetData();
315 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();306 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
316- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;307+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
317 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());308 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
318- WriteDataToFile((const char *)output_file.c_str(), data_size, output_data_i);309+ WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
319- float *resultData = (float*)output_data_i;310+ float* resultData = (float*)output_data_i;
320 for (int64_t j = 0; j < output_shape; j++) {311 for (int64_t j = 0; j < output_shape; j++) {
321 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);312 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);
322 }313 }
@@ -331,7 +322,7 @@ int main(int argc, char *argv[])
331 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());322 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
332 ret = ge::GEFinalize();323 ret = ge::GEFinalize();
333 if (ret != SUCCESS) {324 if (ret != SUCCESS) {
334- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());325+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
335 return FAILED;326 return FAILED;
336 }327 }
337 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());328 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -123,7 +123,7 @@ bool RandomStandardNormalFusionPass::MeetRequirements(const std::unique_ptr<Matc
123 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);123 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);
124 }124 }
125 if (version < GE_COMPILER_VERSION_900) {125 if (version < GE_COMPILER_VERSION_900) {
126- OP_LOGD(kPassName.c_str(), "GE runtime version %d < 90000000, skip pass.", version);126+ OP_LOGD(kPassName.c_str(), "GE runtime version %d < 9.0.0, skip pass.", version);
127 return false;127 return false;
128 }128 }
129 129 
@@ -34,15 +34,16 @@ OpTilingConfig RandomStandardNormalV2Tiling::BuildOpConfig()
34 config.inputCheckRules = {34 config.inputCheckRules = {
35 // 输入索引: dtype列表,shapeSize,dim_num35 // 输入索引: dtype列表,shapeSize,dim_num
36 {0, {{ge::DT_INT32, ge::DT_INT64}, -1, {1}, nullptr}}, // shape36 {0, {{ge::DT_INT32, ge::DT_INT64}, -1, {1}, nullptr}}, // shape
37- {1, {{ge::DT_INT64}, 1, {}, nullptr}}, // offset37+ {1, {{ge::DT_INT64}, 1, {}, nullptr}}, // offset
38 };38 };
39 config.DcacheSize = DCACHE_SIZE;39 config.DcacheSize = DCACHE_SIZE;
40- config.outputCheckRules = {// 输出索引: dtype列表,shapeSize,dim_num40+ config.outputCheckRules = {
41- {0, {{ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16}, -1, {1,2,3,4,5,6,7,8}, nullptr}}}; // y41+ // 输出索引: dtype列表,shapeSize,dim_num
42+ {0, {{ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16}, -1, {1, 2, 3, 4, 5, 6, 7, 8}, nullptr}}}; // y
42 43 
43 // 获取output_size:输入0(shape)的shapeSize44 // 获取output_size:输入0(shape)的shapeSize
44 config.getOutputSize = [](gert::TilingContext* ctx, int64_t& shapeSize) -> ge::graphStatus {45 config.getOutputSize = [](gert::TilingContext* ctx, int64_t& shapeSize) -> ge::graphStatus {
45- return RandomUtils::GetAndCheckOutputSize<0,0>(ctx, shapeSize);46+ return RandomUtils::GetAndCheckOutputSize<0, 0>(ctx, shapeSize);
46 };47 };
47 48 
48 // 获取key[2]:从attr1(seed) counter[4] attr(seed2)49 // 获取key[2]:从attr1(seed) counter[4] attr(seed2)
@@ -78,13 +79,11 @@ static ge::graphStatus TilingPrepare4RandomStandardNormalV2Tiling(gert::TilingPa
78 uint64_t ubSizePlatForm;79 uint64_t ubSizePlatForm;
79 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);80 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
80 compileInfo->ubSize = static_cast<int64_t>(ubSizePlatForm);81 compileInfo->ubSize = static_cast<int64_t>(ubSizePlatForm);
81- OP_CHECK_IF(82+ OP_CHECK_IF((compileInfo->totalCoreNum <= 0 || compileInfo->ubSize <= 0),
82- (compileInfo->totalCoreNum <= 0 || compileInfo->ubSize <= 0),83+ OP_LOGE(context, "RandomStandardNormalV2 GetHardwareInfo Failed, vectorCoreNum:%ld, ubSize:%ld.",
83- OP_LOGE(84+ compileInfo->totalCoreNum, compileInfo->ubSize),
84- context, "RandomStandardNormalV2 GetHardwareInfo Failed, vectorCoreNum:%ld, ubSize:%ld.",85+ return ge::GRAPH_FAILED);
85- compileInfo->totalCoreNum, compileInfo->ubSize),86+ OP_LOGD(context, "Get totalCoreNum:%ld, ubSize:%ld", compileInfo->totalCoreNum, compileInfo->ubSize);
86- return ge::GRAPH_FAILED);
87- OP_LOGD(context, "Get totalCoreNum:%d, ubSize:%ld", compileInfo->totalCoreNum, compileInfo->ubSize);
88 return ge::GRAPH_SUCCESS;87 return ge::GRAPH_SUCCESS;
89}88}
90 89 
@@ -36,87 +36,78 @@ using namespace ge;
36using std::map;36using std::map;
37using std::string;37using std::string;
38using std::vector;38using std::vector;
39-#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \39+#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
40- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \40+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
41- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \41+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
42- TensorDesc placeholder##intputIndex##_desc = \42+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
43- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \43+ intputDtype); \
44- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \44+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
45- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \45+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
46- Tensor tensor_placeholder##intputIndex; \46+ Tensor tensor_placeholder##intputIndex; \
47- ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, \47+ ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
48- tensor_placeholder##intputIndex, \48+ placeholder##intputIndex##_desc, value); \
49- placeholder##intputIndex##_desc, \49+ if (ret != SUCCESS) { \
50- value); \50+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
51- if (ret != SUCCESS) { \51+ return FAILED; \
52- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \52+ } \
53- return FAILED; \53+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
54- } \54+ input.push_back(tensor_placeholder##intputIndex); \
55- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \55+ graph.AddOp(placeholder##intputIndex); \
56- input.push_back(tensor_placeholder##intputIndex); \56+ add1.set_input_##intputName(placeholder##intputIndex); \
57- graph.AddOp(placeholder##intputIndex); \
58- add1.set_input_##intputName(placeholder##intputIndex); \
59 inputs.push_back(placeholder##intputIndex)57 inputs.push_back(placeholder##intputIndex)
60 58 
61-#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \59+#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
62- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \60+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
63- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \61+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
64- TensorDesc placeholder##intputIndex##_desc = \62+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
65- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \63+ intputDtype); \
66- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \64+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
67- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \65+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
68- Tensor tensor_placeholder##intputIndex; \66+ Tensor tensor_placeholder##intputIndex; \
69- ret = GenOnesData(placeholder##intputIndex##_shape, \67+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
70- tensor_placeholder##intputIndex, \68+ placeholder##intputIndex##_desc, intputDtype, value); \
71- placeholder##intputIndex##_desc, \69+ if (ret != SUCCESS) { \
72- intputDtype, \70+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
73- value); \71+ return FAILED; \
74- if (ret != SUCCESS) { \72+ } \
75- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \73+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
76- return FAILED; \74+ input.push_back(tensor_placeholder##intputIndex); \
77- } \75+ graph.AddOp(placeholder##intputIndex); \
78- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \76+ add1.set_input_##intputName(placeholder##intputIndex); \
79- input.push_back(tensor_placeholder##intputIndex); \
80- graph.AddOp(placeholder##intputIndex); \
81- add1.set_input_##intputName(placeholder##intputIndex); \
82 inputs.push_back(placeholder##intputIndex)77 inputs.push_back(placeholder##intputIndex)
83 78 
84-#define ADD_INPUT_ATTR(attrName, attrValue) \79+#define ADD_INPUT_ATTR(attrName, attrValue) add1.set_attr_##attrName(attrValue)
85- add1.set_attr_##attrName(attrValue)
86 80 
87-#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \81+#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \
88- TensorDesc outputName##outputIndex##_desc = \82+ TensorDesc outputName##outputIndex##_desc = TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
89- TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
90 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)83 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)
91 84 
92-#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \85+#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \
93- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \86+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
94- auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \87+ auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \
95- TensorDesc placeholder##intputIndex##_desc = \88+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
96- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \89+ intputDtype); \
97- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \90+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
98- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \91+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
99- Tensor tensor_placeholder##intputIndex; \92+ Tensor tensor_placeholder##intputIndex; \
100- ret = GenOnesData(placeholder##intputIndex##_shape, \93+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
101- tensor_placeholder##intputIndex, \94+ placeholder##intputIndex##_desc, intputDtype, 1); \
102- placeholder##intputIndex##_desc, \95+ if (ret != SUCCESS) { \
103- intputDtype, \96+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
104- 1); \97+ return FAILED; \
105- if (ret != SUCCESS) { \98+ } \
106- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \99+ placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \
107- return FAILED; \100+ \
108- } \101+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
109- placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \ 102+ graph.AddOp(placeholder##intputIndex); \
110- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \103+ add1.set_input_##intputName(placeholder##intputIndex); \
111- graph.AddOp(placeholder##intputIndex); \104+ add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
112- add1.set_input_##intputName(placeholder##intputIndex); \
113- add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
114 inputs.push_back(placeholder##intputIndex)105 inputs.push_back(placeholder##intputIndex)
115 106 
116-#define LOG_PRINT(message, ...) \107+#define LOG_PRINT(message, ...) \
117- do { \108+ do { \
118- printf(message, ##__VA_ARGS__); \109+ printf(message, ##__VA_ARGS__); \
119- } while (0)110+ } while (0)
120 111 
121string GetTime()112string GetTime()
122{113{
@@ -153,7 +144,7 @@ uint32_t GetDataTypeSize(DataType dt)
153 return dilation;144 return dilation;
154}145}
155 146 
156-int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, float value)147+int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, float value)
157{148{
158 input_tensor_desc.SetRealDimCnt(shapes.size());149 input_tensor_desc.SetRealDimCnt(shapes.size());
159 size_t size = 1;150 size_t size = 1;
@@ -162,18 +153,18 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorD
162 }153 }
163 uint32_t byteSizeFloat32 = 4;154 uint32_t byteSizeFloat32 = 4;
164 uint32_t data_len = size * byteSizeFloat32;155 uint32_t data_len = size * byteSizeFloat32;
165- float *pData = new (std::nothrow) float[size];156+ float* pData = new (std::nothrow) float[size];
166 157 
167 for (size_t i = 0; i < size; ++i) {158 for (size_t i = 0; i < size; ++i) {
168 *(pData + i) = value;159 *(pData + i) = value;
169 }160 }
170- input_tensor = Tensor(input_tensor_desc, (uint8_t *)pData, data_len);161+ input_tensor = Tensor(input_tensor_desc, (uint8_t*)pData, data_len);
171 delete[] pData;162 delete[] pData;
172 return SUCCESS;163 return SUCCESS;
173}164}
174 165 
175-int32_t GenOnesData(166+int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type,
176- vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, DataType data_type, int value)167+ int value)
177{168{
178 input_tensor_desc.SetRealDimCnt(shapes.size());169 input_tensor_desc.SetRealDimCnt(shapes.size());
179 size_t size = 1;170 size_t size = 1;
@@ -181,25 +172,25 @@ int32_t GenOnesData(
181 size *= shapes[i];172 size *= shapes[i];
182 }173 }
183 uint32_t data_len = size * GetDataTypeSize(data_type);174 uint32_t data_len = size * GetDataTypeSize(data_type);
184- int64_t *pData = new (std::nothrow) int64_t[size];175+ int64_t* pData = new (std::nothrow) int64_t[size];
185 for (uint32_t i = 0; i < size; ++i) {176 for (uint32_t i = 0; i < size; ++i) {
186 pData[i] = static_cast<int64_t>(value);177 pData[i] = static_cast<int64_t>(value);
187 }178 }
188- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);179+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
189 delete[] pData;180 delete[] pData;
190 return SUCCESS;181 return SUCCESS;
191}182}
192 183 
193-int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t *inputData)184+int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
194{185{
195- FILE *fp = fopen(bin_file.c_str(), "w");186+ FILE* fp = fopen(bin_file.c_str(), "w");
196 fwrite(inputData, sizeof(uint8_t), data_size, fp);187 fwrite(inputData, sizeof(uint8_t), data_size, fp);
197 fclose(fp);188 fclose(fp);
198 return SUCCESS;189 return SUCCESS;
199}190}
200 191 
201-int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vector<Operator> &inputs,192+int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs,
202- std::vector<Operator> &outputs, Graph &graph)193+ std::vector<Operator>& outputs, Graph& graph)
203{194{
204 Status ret = SUCCESS;195 Status ret = SUCCESS;
205 // 自定义代码:添加单算子定义到图中196 // 自定义代码:添加单算子定义到图中
@@ -216,7 +207,7 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
216 207 
217 ADD_INPUT_ATTR(seed, 10);208 ADD_INPUT_ATTR(seed, 10);
218 ADD_INPUT_ATTR(seed2, 5);209 ADD_INPUT_ATTR(seed2, 5);
219- 210+ 
220 ADD_OUTPUT(1, y, ge::DT_INT64, outShape);211 ADD_OUTPUT(1, y, ge::DT_INT64, outShape);
221 ADD_OUTPUT(2, offset, ge::DT_INT64, offsetShape);212 ADD_OUTPUT(2, offset, ge::DT_INT64, offsetShape);
222 213 
@@ -225,9 +216,9 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
225 return SUCCESS;216 return SUCCESS;
226}217}
227 218 
228-int main(int argc, char *argv[])219+int main(int argc, char* argv[])
229{220{
230- const char *graph_name = "tc_ge_irrun_test";221+ const char* graph_name = "tc_ge_irrun_test";
231 Graph graph(graph_name);222 Graph graph(graph_name);
232 std::vector<ge::Tensor> input;223 std::vector<ge::Tensor> input;
233 224 
@@ -235,7 +226,7 @@ int main(int argc, char *argv[])
235 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};226 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
236 Status ret = ge::GEInitialize(global_options);227 Status ret = ge::GEInitialize(global_options);
237 if (ret != SUCCESS) {228 if (ret != SUCCESS) {
238- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());229+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
239 return FAILED;230 return FAILED;
240 }231 }
241 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());232 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -244,7 +235,7 @@ int main(int argc, char *argv[])
244 std::vector<Operator> outputs{};235 std::vector<Operator> outputs{};
245 236 
246 std::cout << argv[1] << std::endl;237 std::cout << argv[1] << std::endl;
247- char *endptr;238+ char* endptr;
248 239 
249 DataType inDtype = DT_INT64;240 DataType inDtype = DT_INT64;
250 std::cout << inDtype << std::endl;241 std::cout << inDtype << std::endl;
@@ -263,7 +254,7 @@ int main(int argc, char *argv[])
263 254 
264 };255 };
265 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());256 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());
266- ge::Session *session = new Session(build_options);257+ ge::Session* session = new Session(build_options);
267 258 
268 if (session == nullptr) {259 if (session == nullptr) {
269 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());260 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());
@@ -286,7 +277,7 @@ int main(int argc, char *argv[])
286 std::vector<ge::Tensor> output;277 std::vector<ge::Tensor> output;
287 ret = session->RunGraph(graph_id, input, output);278 ret = session->RunGraph(graph_id, input, output);
288 if (ret != SUCCESS) {279 if (ret != SUCCESS) {
289- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());280+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
290 delete session;281 delete session;
291 GEFinalize();282 GEFinalize();
292 return FAILED;283 return FAILED;
@@ -297,25 +288,25 @@ int main(int argc, char *argv[])
297 for (int i = 0; i < input_num; i++) {288 for (int i = 0; i < input_num; i++) {
298 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;289 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;
299 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";290 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
300- uint8_t *input_data_i = input[i].GetData();291+ uint8_t* input_data_i = input[i].GetData();
301 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();292 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
302- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;293+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
303 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());294 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
304- WriteDataToFile((const char *)input_file.c_str(), data_size, input_data_i);295+ WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
305 }296 }
306 297 
307 int output_num = output.size();298 int output_num = output.size();
308 for (int i = 0; i < output_num; i++) {299 for (int i = 0; i < output_num; i++) {
309 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;300 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;
310 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";301 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
311- uint8_t *output_data_i = output[i].GetData();302+ uint8_t* output_data_i = output[i].GetData();
312 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();303 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
313- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;304+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
314 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());305 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
315- WriteDataToFile((const char *)output_file.c_str(), data_size, output_data_i);306+ WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
316- int64_t *resultData = (int64_t*)output_data_i;307+ int64_t* resultData = (int64_t*)output_data_i;
317 for (int64_t j = 0; j < output_shape; j++) {308 for (int64_t j = 0; j < output_shape; j++) {
318- LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);309+ LOG_PRINT("result[%ld] is: %ld\n", j, resultData[j]);
319 }310 }
320 }311 }
321 312 
@@ -328,7 +319,7 @@ int main(int argc, char *argv[])
328 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());319 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
329 ret = ge::GEFinalize();320 ret = ge::GEFinalize();
330 if (ret != SUCCESS) {321 if (ret != SUCCESS) {
331- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());322+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
332 return FAILED;323 return FAILED;
333 }324 }
334 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());325 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -19,17 +19,17 @@
19#include "platform/platform_ascendc.h"19#include "platform/platform_ascendc.h"
20#include "op_common/op_host/util/platform_util.h"20#include "op_common/op_host/util/platform_util.h"
21#include "random_uniform_int_v2_tiling_arch35.h"21#include "random_uniform_int_v2_tiling_arch35.h"
22-#include "random/random_common/op_host/arch35/random_tiling_base.h"22+#include "random/random_common/op_host/arch35/random_tiling_base.h"
23#include "op_host/math_tiling_templates_registry.h"23#include "op_host/math_tiling_templates_registry.h"
24#include "register/op_def_registry.h"24#include "register/op_def_registry.h"
25 25 
26namespace optiling {26namespace optiling {
27 27 
28template <typename T>28template <typename T>
29-ge::graphStatus RandomUniformIntV2Tiling::GetIntValue(const gert::Tensor *constTensor, gert::Shape &constShape)29+ge::graphStatus RandomUniformIntV2Tiling::GetIntValue(const gert::Tensor* constTensor, gert::Shape& constShape)
30{30{
31 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetIntValue begin.");31 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetIntValue begin.");
32- const T *constValue = constTensor->GetData<T>();32+ const T* constValue = constTensor->GetData<T>();
33 OP_CHECK_NULL_WITH_CONTEXT(context_, constValue);33 OP_CHECK_NULL_WITH_CONTEXT(context_, constValue);
34 const size_t constNum = constTensor->GetShapeSize();34 const size_t constNum = constTensor->GetShapeSize();
35 constShape.SetDimNum(0);35 constShape.SetDimNum(0);
@@ -40,7 +40,7 @@ ge::graphStatus RandomUniformIntV2Tiling::GetIntValue(const gert::Tensor *constT
40 return ge::GRAPH_SUCCESS;40 return ge::GRAPH_SUCCESS;
41}41}
42 42 
43-ge::graphStatus RandomUniformIntV2Tiling::GetIntValueByDtype(const gert::Tensor *constTensor, gert::Shape &constShape,43+ge::graphStatus RandomUniformIntV2Tiling::GetIntValueByDtype(const gert::Tensor* constTensor, gert::Shape& constShape,
44 ge::DataType dType)44 ge::DataType dType)
45{45{
46 ge::graphStatus ret = ge::GRAPH_SUCCESS;46 ge::graphStatus ret = ge::GRAPH_SUCCESS;
@@ -59,8 +59,8 @@ ge::graphStatus RandomUniformIntV2Tiling::GetMinAndMaxValue()
59 OP_CHECK_NULL_WITH_CONTEXT(context_, minDesc);59 OP_CHECK_NULL_WITH_CONTEXT(context_, minDesc);
60 minDtype_ = minDesc->GetDataType();60 minDtype_ = minDesc->GetDataType();
61 if ((minDtype_ != ge::DataType::DT_INT32) && (minDtype_ != ge::DataType::DT_INT64)) {61 if ((minDtype_ != ge::DataType::DT_INT32) && (minDtype_ != ge::DataType::DT_INT64)) {
62- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "min",62+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "min", Ops::Base::ToString(minDtype_).c_str(),
63- Ops::Base::ToString(minDtype_).c_str(), "dtype must be in [DT_INT32, DT_INT64]");63+ "dtype must be in [DT_INT32, DT_INT64]");
64 return ge::GRAPH_FAILED;64 return ge::GRAPH_FAILED;
65 }65 }
66 66 
@@ -68,23 +68,22 @@ ge::graphStatus RandomUniformIntV2Tiling::GetMinAndMaxValue()
68 OP_CHECK_NULL_WITH_CONTEXT(context_, minTensor);68 OP_CHECK_NULL_WITH_CONTEXT(context_, minTensor);
69 auto minTensorSize = static_cast<int64_t>(minTensor->GetShapeSize());69 auto minTensorSize = static_cast<int64_t>(minTensor->GetShapeSize());
70 if (minTensorSize != 1) {70 if (minTensorSize != 1) {
71- OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "min data shape_size",71+ OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "min data shape_size", std::to_string(minTensorSize).c_str(),
72- std::to_string(minTensorSize).c_str(), "shape_size must be 1");72+ "shape_size must be 1");
73 return ge::GRAPH_FAILED;73 return ge::GRAPH_FAILED;
74 }74 }
75 gert::Shape minShape;75 gert::Shape minShape;
76 auto ret = GetIntValueByDtype(minTensor, minShape, minDtype_);76 auto ret = GetIntValueByDtype(minTensor, minShape, minDtype_);
77- OP_CHECK_IF(ret != ge::GRAPH_SUCCESS,77+ OP_CHECK_IF(ret != ge::GRAPH_SUCCESS, OP_LOGE(opName_, "min GetIntValueByDtype failed."), return ge::GRAPH_FAILED);
78- OP_LOGE(opName_, "min GetIntValueByDtype failed."), return ge::GRAPH_FAILED);
79 lo_ = static_cast<int64_t>(minShape.GetDim((0)));78 lo_ = static_cast<int64_t>(minShape.GetDim((0)));
80 79 
81 auto maxDesc = context_->GetRequiredInputDesc(IN_MAX_IDX);80 auto maxDesc = context_->GetRequiredInputDesc(IN_MAX_IDX);
82 OP_CHECK_NULL_WITH_CONTEXT(context_, maxDesc);81 OP_CHECK_NULL_WITH_CONTEXT(context_, maxDesc);
83 auto maxDtype = maxDesc->GetDataType();82 auto maxDtype = maxDesc->GetDataType();
84 if (maxDtype != minDtype_) {83 if (maxDtype != minDtype_) {
85- OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(opName_, "max, min",84+ OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(
86- (Ops::Base::ToString(maxDtype) + ", " + Ops::Base::ToString(minDtype_)).c_str(),85+ opName_, "max, min", (Ops::Base::ToString(maxDtype) + ", " + Ops::Base::ToString(minDtype_)).c_str(),
87- "max dtype must be same as min dtype");86+ "max dtype must be the same as min dtype");
88 return ge::GRAPH_FAILED;87 return ge::GRAPH_FAILED;
89 }88 }
90 89 
@@ -92,19 +91,18 @@ ge::graphStatus RandomUniformIntV2Tiling::GetMinAndMaxValue()
92 OP_CHECK_NULL_WITH_CONTEXT(context_, maxTensor);91 OP_CHECK_NULL_WITH_CONTEXT(context_, maxTensor);
93 auto maxTensorSize = static_cast<int64_t>(maxTensor->GetShapeSize());92 auto maxTensorSize = static_cast<int64_t>(maxTensor->GetShapeSize());
94 if (maxTensorSize != 1) {93 if (maxTensorSize != 1) {
95- OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "max data shape_size",94+ OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "max data shape_size", std::to_string(maxTensorSize).c_str(),
96- std::to_string(maxTensorSize).c_str(), "shape_size must be 1");95+ "shape_size must be 1");
97 return ge::GRAPH_FAILED;96 return ge::GRAPH_FAILED;
98 }97 }
99 gert::Shape maxShape;98 gert::Shape maxShape;
100 ret = GetIntValueByDtype(maxTensor, maxShape, maxDtype);99 ret = GetIntValueByDtype(maxTensor, maxShape, maxDtype);
101- OP_CHECK_IF(ret != ge::GRAPH_SUCCESS,100+ OP_CHECK_IF(ret != ge::GRAPH_SUCCESS, OP_LOGE(opName_, "max GetIntValueByDtype failed."), return ge::GRAPH_FAILED);
102- OP_LOGE(opName_, "max GetIntValueByDtype failed."), return ge::GRAPH_FAILED);
103 const int64_t maxTensorValue = static_cast<int64_t>(maxShape.GetDim((0)));101 const int64_t maxTensorValue = static_cast<int64_t>(maxShape.GetDim((0)));
104 if (maxTensorValue <= lo_) {102 if (maxTensorValue <= lo_) {
105 OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(opName_, "max, min",103 OP_LOGE_FOR_INVALID_VALUES_WITH_REASON(opName_, "max, min",
106- (std::to_string(maxTensorValue) + ", " + std::to_string(lo_)).c_str(),104+ (std::to_string(maxTensorValue) + ", " + std::to_string(lo_)).c_str(),
107- "max value must be greater than min value");105+ "max value must be greater than min value");
108 return ge::GRAPH_FAILED;106 return ge::GRAPH_FAILED;
109 }107 }
110 range_ = static_cast<uint64_t>(maxTensorValue) - static_cast<uint64_t>(lo_);108 range_ = static_cast<uint64_t>(maxTensorValue) - static_cast<uint64_t>(lo_);
@@ -122,8 +120,10 @@ ge::graphStatus RandomUniformIntV2Tiling::GetPlatformInfo()
122 120 
123 totalCoreNum_ = static_cast<int64_t>(compileInfo->totalCoreNum);121 totalCoreNum_ = static_cast<int64_t>(compileInfo->totalCoreNum);
124 ubSize_ = compileInfo->ubSize;122 ubSize_ = compileInfo->ubSize;
125- OP_CHECK_IF((ubSize_ <= 0), OP_LOGE(opName_, "ub size is invalid."), return ge::GRAPH_FAILED);123+ OP_CHECK_IF((ubSize_ <= 0), OP_LOGE(opName_, "ub size %ld is invalid, must be greater than 0.", ubSize_),
126- OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetPlatformInfo ubSize_=%d, totalCoreNum_=%d", ubSize_, totalCoreNum_);124+ return ge::GRAPH_FAILED);
125+ OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetPlatformInfo ubSize_=%ld, totalCoreNum_=%ld", ubSize_,
126+ totalCoreNum_);
127 return ge::GRAPH_SUCCESS;127 return ge::GRAPH_SUCCESS;
128}128}
129 129 
@@ -131,21 +131,18 @@ ge::graphStatus RandomUniformIntV2Tiling::GetPlatformInfo()
131ge::graphStatus RandomUniformIntV2Tiling::GetShapeAttrsInfo()131ge::graphStatus RandomUniformIntV2Tiling::GetShapeAttrsInfo()
132{132{
133 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetShapeAttrsInfo begin.");133 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetShapeAttrsInfo begin.");
134- OP_CHECK_IF(GetInputInfo(), 134+ OP_CHECK_IF(GetInputInfo(), OP_LOGE(opName_, "GetInputInfo failed!"), return ge::GRAPH_FAILED);
135- OP_LOGE(opName_, "GetInputInfo failed!"), return ge::GRAPH_FAILED);135+ 
136- 136+ OP_CHECK_IF(GetOutputInfo(), OP_LOGE(opName_, "GetOutputInfo failed!"), return ge::GRAPH_FAILED);
137- OP_CHECK_IF(GetOutputInfo(), 137+ 
138- OP_LOGE(opName_, "GetOutputInfo failed!"), return ge::GRAPH_FAILED);
139-
140 if (shapeSize_ != outputSize_) {138 if (shapeSize_ != outputSize_) {
141- OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(opName_, "input shape, output",139+ OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(
142- (std::to_string(shapeSize_) + ", " + std::to_string(outputSize_)).c_str(),140+ opName_, "input shape, output", (std::to_string(shapeSize_) + ", " + std::to_string(outputSize_)).c_str(),
143 "input shape size must be equal to output size");141 "input shape size must be equal to output size");
144 return ge::GRAPH_FAILED;142 return ge::GRAPH_FAILED;
145 }143 }
146 144 
147- OP_CHECK_IF(GetAttrInfo(), 145+ OP_CHECK_IF(GetAttrInfo(), OP_LOGE(opName_, "GetAttrInfo failed!"), return ge::GRAPH_FAILED);
148- OP_LOGE(opName_, "GetAttrInfo failed!"), return ge::GRAPH_FAILED);
149 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetShapeAttrsInfo end.");146 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetShapeAttrsInfo end.");
150 return ge::GRAPH_SUCCESS;147 return ge::GRAPH_SUCCESS;
151}148}
@@ -157,8 +154,8 @@ ge::graphStatus RandomUniformIntV2Tiling::GetInputInfo()
157 OP_CHECK_NULL_WITH_CONTEXT(context_, shapeDesc);154 OP_CHECK_NULL_WITH_CONTEXT(context_, shapeDesc);
158 auto shapeDtype = shapeDesc->GetDataType();155 auto shapeDtype = shapeDesc->GetDataType();
159 if ((shapeDtype != ge::DataType::DT_INT32) && (shapeDtype != ge::DataType::DT_INT64)) {156 if ((shapeDtype != ge::DataType::DT_INT32) && (shapeDtype != ge::DataType::DT_INT64)) {
160- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "input shape",157+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "input shape", Ops::Base::ToString(shapeDtype).c_str(),
161- Ops::Base::ToString(shapeDtype).c_str(), "dtype must be in [DT_INT32, DT_INT64]");158+ "dtype must be in [DT_INT32, DT_INT64]");
162 return ge::GRAPH_FAILED;159 return ge::GRAPH_FAILED;
163 }160 }
164 161 
@@ -166,8 +163,8 @@ ge::graphStatus RandomUniformIntV2Tiling::GetInputInfo()
166 OP_CHECK_NULL_WITH_CONTEXT(context_, input1Shape);163 OP_CHECK_NULL_WITH_CONTEXT(context_, input1Shape);
167 uint32_t shapeDimNum = input1Shape->GetStorageShape().GetDimNum();164 uint32_t shapeDimNum = input1Shape->GetStorageShape().GetDimNum();
168 if (shapeDimNum != 1) {165 if (shapeDimNum != 1) {
169- OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(opName_, "input shape",166+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(opName_, "input shape", std::to_string(shapeDimNum).c_str(),
170- std::to_string(shapeDimNum).c_str(), "must be 1D tensor");167+ "must be 1D tensor");
171 return ge::GRAPH_FAILED;168 return ge::GRAPH_FAILED;
172 }169 }
173 170 
@@ -175,8 +172,8 @@ ge::graphStatus RandomUniformIntV2Tiling::GetInputInfo()
175 OP_CHECK_NULL_WITH_CONTEXT(context_, shapeTensor);172 OP_CHECK_NULL_WITH_CONTEXT(context_, shapeTensor);
176 gert::Shape constShape;173 gert::Shape constShape;
177 auto ret = GetIntValueByDtype(shapeTensor, constShape, shapeDtype);174 auto ret = GetIntValueByDtype(shapeTensor, constShape, shapeDtype);
178- OP_CHECK_IF(ret != ge::GRAPH_SUCCESS,175+ OP_CHECK_IF(ret != ge::GRAPH_SUCCESS, OP_LOGE(opName_, "input shape GetIntValueByDtype failed."),
179- OP_LOGE(opName_, "input shape GetIntValueByDtype failed."), return ge::GRAPH_FAILED);176+ return ge::GRAPH_FAILED);
180 OP_LOGD(opName_, "RandomUniformIntV2Tiling::GetInputInfo get shapeTensor end.");177 OP_LOGD(opName_, "RandomUniformIntV2Tiling::GetInputInfo get shapeTensor end.");
181 178 
182 uint32_t shapeRank = constShape.GetDimNum();179 uint32_t shapeRank = constShape.GetDimNum();
@@ -184,8 +181,8 @@ ge::graphStatus RandomUniformIntV2Tiling::GetInputInfo()
184 shapeSize_ *= static_cast<int64_t>(constShape.GetDim(idx));181 shapeSize_ *= static_cast<int64_t>(constShape.GetDim(idx));
185 }182 }
186 if (shapeSize_ == 0) {183 if (shapeSize_ == 0) {
187- OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "input shape",184+ OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "input shape", std::to_string(shapeSize_).c_str(),
188- std::to_string(shapeSize_).c_str(), "shape_size must not be 0");185+ "shape_size must not be 0");
189 return ge::GRAPH_FAILED;186 return ge::GRAPH_FAILED;
190 }187 }
191 188 
@@ -193,17 +190,17 @@ ge::graphStatus RandomUniformIntV2Tiling::GetInputInfo()
193 OP_CHECK_NULL_WITH_CONTEXT(context_, offsetDesc);190 OP_CHECK_NULL_WITH_CONTEXT(context_, offsetDesc);
194 auto offsetDtype = offsetDesc->GetDataType();191 auto offsetDtype = offsetDesc->GetDataType();
195 if (offsetDtype != ge::DataType::DT_INT64) {192 if (offsetDtype != ge::DataType::DT_INT64) {
196- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "input offset",193+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "input offset", Ops::Base::ToString(offsetDtype).c_str(),
197- Ops::Base::ToString(offsetDtype).c_str(), "dtype must be DT_INT64");194+ "dtype must be DT_INT64");
198 return ge::GRAPH_FAILED;195 return ge::GRAPH_FAILED;
199 }196 }
200- 197+ 
201 auto offsetTensor = context_->GetInputTensor(IN_OFFSET_IDX);198 auto offsetTensor = context_->GetInputTensor(IN_OFFSET_IDX);
202 OP_CHECK_NULL_WITH_CONTEXT(context_, offsetTensor);199 OP_CHECK_NULL_WITH_CONTEXT(context_, offsetTensor);
203- auto offsetTensorSize = static_cast<int64_t>(offsetTensor->GetShapeSize()); // 验证200+ auto offsetTensorSize = static_cast<int64_t>(offsetTensor->GetShapeSize()); // 验证
204 if (offsetTensorSize != 1) {201 if (offsetTensorSize != 1) {
205 OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "input offset shape_size",202 OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "input offset shape_size",
206- std::to_string(offsetTensorSize).c_str(), "shape_size must be 1");203+ std::to_string(offsetTensorSize).c_str(), "shape_size must be 1");
207 return ge::GRAPH_FAILED;204 return ge::GRAPH_FAILED;
208 }205 }
209 206 
@@ -221,9 +218,9 @@ ge::graphStatus RandomUniformIntV2Tiling::GetOutputInfo()
221 OP_CHECK_NULL_WITH_CONTEXT(context_, outDesc);218 OP_CHECK_NULL_WITH_CONTEXT(context_, outDesc);
222 outDtype_ = outDesc->GetDataType();219 outDtype_ = outDesc->GetDataType();
223 if (outDtype_ != minDtype_) {220 if (outDtype_ != minDtype_) {
224- OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(opName_, "out, min",221+ OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(
225- (Ops::Base::ToString(outDtype_) + ", " + Ops::Base::ToString(minDtype_)).c_str(),222+ opName_, "out, min", (Ops::Base::ToString(outDtype_) + ", " + Ops::Base::ToString(minDtype_)).c_str(),
226- "out dtype must be same as min dtype");223+ "out dtype must be the same as min dtype");
227 return ge::GRAPH_FAILED;224 return ge::GRAPH_FAILED;
228 }225 }
229 226 
@@ -232,8 +229,8 @@ ge::graphStatus RandomUniformIntV2Tiling::GetOutputInfo()
232 auto outTensor = outputShape->GetStorageShape();229 auto outTensor = outputShape->GetStorageShape();
233 outputSize_ = outTensor.GetShapeSize();230 outputSize_ = outTensor.GetShapeSize();
234 if (outputSize_ == 0) {231 if (outputSize_ == 0) {
235- OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "output shape_size",232+ OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(opName_, "output shape_size", std::to_string(outputSize_).c_str(),
236- std::to_string(outputSize_).c_str(), "shape_size must not be 0");233+ "shape_size must not be 0");
237 return ge::GRAPH_FAILED;234 return ge::GRAPH_FAILED;
238 }235 }
239 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetOutputInfo end.");236 OP_LOGI(opName_, "RandomUniformIntV2Tiling::GetOutputInfo end.");
@@ -249,7 +246,7 @@ ge::graphStatus RandomUniformIntV2Tiling::GetAttrInfo()
249 OP_CHECK_NULL_WITH_CONTEXT(context_, seedAttr);246 OP_CHECK_NULL_WITH_CONTEXT(context_, seedAttr);
250 const auto* seed2Attr = attrs->GetAttrPointer<int64_t>(ATTR_SEED2_IDX);247 const auto* seed2Attr = attrs->GetAttrPointer<int64_t>(ATTR_SEED2_IDX);
251 OP_CHECK_NULL_WITH_CONTEXT(context_, seed2Attr);248 OP_CHECK_NULL_WITH_CONTEXT(context_, seed2Attr);
252- 249+ 
253 seed_ = *seedAttr;250 seed_ = *seedAttr;
254 seed2_ = *seed2Attr;251 seed2_ = *seed2Attr;
255 if (seed_ == 0 && seed2_ == 0) {252 if (seed_ == 0 && seed2_ == 0) {
@@ -260,10 +257,7 @@ ge::graphStatus RandomUniformIntV2Tiling::GetAttrInfo()
260 return ge::GRAPH_SUCCESS;257 return ge::GRAPH_SUCCESS;
261}258}
262 259 
263-bool RandomUniformIntV2Tiling::IsCapable()260+bool RandomUniformIntV2Tiling::IsCapable() { return true; }
264-{
265- return true;
266-}
267 261 
268void RandomUniformIntV2Tiling::SetTilingData()262void RandomUniformIntV2Tiling::SetTilingData()
269{263{
@@ -279,7 +273,7 @@ void RandomUniformIntV2Tiling::SetTilingData()
279 tilingData->lo = lo_;273 tilingData->lo = lo_;
280}274}
281 275 
282-void RandomUniformIntV2Tiling::DoBlockTiling() 276+void RandomUniformIntV2Tiling::DoBlockTiling()
283{277{
284 outputDtypeSize_ = ge::GetSizeByDataType(outDtype_);278 outputDtypeSize_ = ge::GetSizeByDataType(outDtype_);
285 if (outputDtypeSize_ == 0) {279 if (outputDtypeSize_ == 0) {
@@ -296,14 +290,14 @@ void RandomUniformIntV2Tiling::DoBlockTiling()
296 return;290 return;
297}291}
298 292 
299-void RandomUniformIntV2Tiling::UbTiling() 293+void RandomUniformIntV2Tiling::UbTiling()
300{294{
301- // quarterUbSize: 2 for double buffer; coefVal for temp RNG, philox temp buff need uint32 to int32/int64295+ // quarterUbSize: 2 for double buffer; coefVal for temp RNG, philox temp buff need uint32 to int32/int64
302- int64_t coefVal = DOUBLE_BUFFER;296+ int64_t coefVal = DOUBLE_BUFFER;
303- auto quarterUbSize = (ubSize_ - DCACHE_SIZE) / (DOUBLE_BUFFER + coefVal);297+ auto quarterUbSize = (ubSize_ - DCACHE_SIZE) / (DOUBLE_BUFFER + coefVal);
304- auto ubBlockSize = static_cast<int32_t>(Ops::Base::GetUbBlockSize(context_));298+ auto ubBlockSize = static_cast<int32_t>(Ops::Base::GetUbBlockSize(context_));
305- auto alignFactor = ubBlockSize / outputDtypeSize_;299+ auto alignFactor = ubBlockSize / outputDtypeSize_;
306- singleUbSize_ = (quarterUbSize / outputDtypeSize_ / alignFactor) * alignFactor;300+ singleUbSize_ = (quarterUbSize / outputDtypeSize_ / alignFactor) * alignFactor;
307}301}
308 302 
309// 3、计算数据切分TilingData303// 3、计算数据切分TilingData
@@ -317,10 +311,7 @@ ge::graphStatus RandomUniformIntV2Tiling::DoOpTiling()
317}311}
318 312 
319// 4、计算高阶API的TilingData313// 4、计算高阶API的TilingData
320-ge::graphStatus RandomUniformIntV2Tiling::DoLibApiTiling()314+ge::graphStatus RandomUniformIntV2Tiling::DoLibApiTiling() { return ge::GRAPH_SUCCESS; }
321-{
322- return ge::GRAPH_SUCCESS;
323-}
324 315 
325// 5、计算TilingKey316// 5、计算TilingKey
326uint64_t RandomUniformIntV2Tiling::GetTilingKey() const317uint64_t RandomUniformIntV2Tiling::GetTilingKey() const
@@ -377,13 +368,13 @@ static ge::graphStatus TilingPrepare4RandomUniformIntV2Tiling(gert::TilingParseC
377 uint64_t ubSizePlatForm;368 uint64_t ubSizePlatForm;
378 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);369 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
379 compileInfo->ubSize = static_cast<int64_t>(ubSizePlatForm);370 compileInfo->ubSize = static_cast<int64_t>(ubSizePlatForm);
380- OP_CHECK_IF(371+ OP_CHECK_IF((compileInfo->totalCoreNum <= 0 || compileInfo->ubSize <= 0),
381- (compileInfo->totalCoreNum <= 0 || compileInfo->ubSize <= 0),372+ OP_LOGE(context,
382- OP_LOGE(373+ "RandomUniformIntV2 GetHardwareInfo failed, vectorCoreNum and ubSize should be greater than 0, "
383- context, "RandomUniformIntV2 GetHardwareInfo Failed, vectorCoreNum:%ld, ubSize:%ld.", compileInfo->totalCoreNum,374+ "vectorCoreNum:%ld, ubSize:%ld.",
384- compileInfo->ubSize),375+ compileInfo->totalCoreNum, compileInfo->ubSize),
385- return ge::GRAPH_FAILED);376+ return ge::GRAPH_FAILED);
386- OP_LOGD(context, "Get totalCoreNum:%d, ubSize:%ld", compileInfo->totalCoreNum, compileInfo->ubSize);377+ OP_LOGD(context, "Get totalCoreNum:%ld, ubSize:%ld", compileInfo->totalCoreNum, compileInfo->ubSize);
387 return ge::GRAPH_SUCCESS;378 return ge::GRAPH_SUCCESS;
388}379}
389 380 
@@ -399,4 +390,4 @@ IMPL_OP_OPTILING(RandomUniformIntV2)
399 .TilingParse<RandomUniformIntV2CompileInfo>(TilingPrepare4RandomUniformIntV2Tiling)390 .TilingParse<RandomUniformIntV2CompileInfo>(TilingPrepare4RandomUniformIntV2Tiling)
400 .TilingInputsDataDependency({IN_SHAPE_IDX, IN_MIN_IDX, IN_MAX_IDX});391 .TilingInputsDataDependency({IN_SHAPE_IDX, IN_MIN_IDX, IN_MAX_IDX});
401 392 
402-} // namespace optiling393+} // namespace optiling
@@ -37,87 +37,78 @@ using namespace ge;
37using std::map;37using std::map;
38using std::string;38using std::string;
39using std::vector;39using std::vector;
40-#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \40+#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
41- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \41+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
42- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \42+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
43- TensorDesc placeholder##intputIndex##_desc = \43+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
44- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \44+ intputDtype); \
45- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \45+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
46- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \46+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
47- Tensor tensor_placeholder##intputIndex; \47+ Tensor tensor_placeholder##intputIndex; \
48- ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, \48+ ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
49- tensor_placeholder##intputIndex, \49+ placeholder##intputIndex##_desc, value); \
50- placeholder##intputIndex##_desc, \50+ if (ret != SUCCESS) { \
51- value); \51+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
52- if (ret != SUCCESS) { \52+ return FAILED; \
53- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \53+ } \
54- return FAILED; \54+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
55- } \55+ input.push_back(tensor_placeholder##intputIndex); \
56- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \56+ graph.AddOp(placeholder##intputIndex); \
57- input.push_back(tensor_placeholder##intputIndex); \57+ add1.set_input_##intputName(placeholder##intputIndex); \
58- graph.AddOp(placeholder##intputIndex); \
59- add1.set_input_##intputName(placeholder##intputIndex); \
60 inputs.push_back(placeholder##intputIndex)58 inputs.push_back(placeholder##intputIndex)
61 59 
62-#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \60+#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
63- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \61+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
64- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \62+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
65- TensorDesc placeholder##intputIndex##_desc = \63+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
66- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \64+ intputDtype); \
67- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \65+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
68- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \66+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
69- Tensor tensor_placeholder##intputIndex; \67+ Tensor tensor_placeholder##intputIndex; \
70- ret = GenOnesData(placeholder##intputIndex##_shape, \68+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
71- tensor_placeholder##intputIndex, \69+ placeholder##intputIndex##_desc, intputDtype, value); \
72- placeholder##intputIndex##_desc, \70+ if (ret != SUCCESS) { \
73- intputDtype, \71+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
74- value); \72+ return FAILED; \
75- if (ret != SUCCESS) { \73+ } \
76- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \74+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
77- return FAILED; \75+ input.push_back(tensor_placeholder##intputIndex); \
78- } \76+ graph.AddOp(placeholder##intputIndex); \
79- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \77+ add1.set_input_##intputName(placeholder##intputIndex); \
80- input.push_back(tensor_placeholder##intputIndex); \
81- graph.AddOp(placeholder##intputIndex); \
82- add1.set_input_##intputName(placeholder##intputIndex); \
83 inputs.push_back(placeholder##intputIndex)78 inputs.push_back(placeholder##intputIndex)
84 79 
85-#define ADD_INPUT_ATTR(attrName, attrValue) \80+#define ADD_INPUT_ATTR(attrName, attrValue) add1.set_attr_##attrName(attrValue)
86- add1.set_attr_##attrName(attrValue)
87 81 
88-#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \82+#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \
89- TensorDesc outputName##outputIndex##_desc = \83+ TensorDesc outputName##outputIndex##_desc = TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
90- TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
91 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)84 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)
92 85 
93-#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \86+#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \
94- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \87+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
95- auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \88+ auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \
96- TensorDesc placeholder##intputIndex##_desc = \89+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
97- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \90+ intputDtype); \
98- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \91+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
99- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \92+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
100- Tensor tensor_placeholder##intputIndex; \93+ Tensor tensor_placeholder##intputIndex; \
101- ret = GenOnesData(placeholder##intputIndex##_shape, \94+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
102- tensor_placeholder##intputIndex, \95+ placeholder##intputIndex##_desc, intputDtype, 1); \
103- placeholder##intputIndex##_desc, \96+ if (ret != SUCCESS) { \
104- intputDtype, \97+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
105- 1); \98+ return FAILED; \
106- if (ret != SUCCESS) { \99+ } \
107- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \100+ placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \
108- return FAILED; \101+ \
109- } \102+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
110- placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \ 103+ graph.AddOp(placeholder##intputIndex); \
111- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \104+ add1.set_input_##intputName(placeholder##intputIndex); \
112- graph.AddOp(placeholder##intputIndex); \105+ add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
113- add1.set_input_##intputName(placeholder##intputIndex); \
114- add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
115 inputs.push_back(placeholder##intputIndex)106 inputs.push_back(placeholder##intputIndex)
116 107 
117-#define LOG_PRINT(message, ...) \108+#define LOG_PRINT(message, ...) \
118- do { \109+ do { \
119- printf(message, ##__VA_ARGS__); \110+ printf(message, ##__VA_ARGS__); \
120- } while (0)111+ } while (0)
121 112 
122string GetTime()113string GetTime()
123{114{
@@ -160,7 +151,7 @@ uint32_t GetDataTypeSize(DataType dt)
160 return dilation;151 return dilation;
161}152}
162 153 
163-int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, float value)154+int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, float value)
164{155{
165 input_tensor_desc.SetRealDimCnt(shapes.size());156 input_tensor_desc.SetRealDimCnt(shapes.size());
166 size_t size = 1;157 size_t size = 1;
@@ -169,17 +160,17 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorD
169 }160 }
170 uint32_t byteSizeFloat32 = 4;161 uint32_t byteSizeFloat32 = 4;
171 uint32_t data_len = size * byteSizeFloat32;162 uint32_t data_len = size * byteSizeFloat32;
172- float *pData = new (std::nothrow) float[size];163+ float* pData = new (std::nothrow) float[size];
173 164 
174 for (size_t i = 0; i < size; ++i) {165 for (size_t i = 0; i < size; ++i) {
175 *(pData + i) = value;166 *(pData + i) = value;
176 }167 }
177- input_tensor = Tensor(input_tensor_desc, (uint8_t *)pData, data_len);168+ input_tensor = Tensor(input_tensor_desc, (uint8_t*)pData, data_len);
178 return SUCCESS;169 return SUCCESS;
179}170}
180 171 
181-int32_t GenOnesData(172+int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type,
182- vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, DataType data_type, int value)173+ int value)
183{174{
184 input_tensor_desc.SetRealDimCnt(shapes.size());175 input_tensor_desc.SetRealDimCnt(shapes.size());
185 size_t size = 1;176 size_t size = 1;
@@ -187,24 +178,24 @@ int32_t GenOnesData(
187 size *= shapes[i];178 size *= shapes[i];
188 }179 }
189 uint32_t data_len = size * GetDataTypeSize(data_type);180 uint32_t data_len = size * GetDataTypeSize(data_type);
190- int32_t *pData = new (std::nothrow) int32_t[data_len];181+ int32_t* pData = new (std::nothrow) int32_t[data_len];
191 for (uint32_t i = 0; i < size; ++i) {182 for (uint32_t i = 0; i < size; ++i) {
192 *(pData + i) = value;183 *(pData + i) = value;
193 }184 }
194- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);185+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
195 return SUCCESS;186 return SUCCESS;
196}187}
197 188 
198-int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t *inputData)189+int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
199{190{
200- FILE *fp = fopen(bin_file.c_str(), "w");191+ FILE* fp = fopen(bin_file.c_str(), "w");
201 fwrite(inputData, sizeof(uint8_t), data_size, fp);192 fwrite(inputData, sizeof(uint8_t), data_size, fp);
202 fclose(fp);193 fclose(fp);
203 return SUCCESS;194 return SUCCESS;
204}195}
205 196 
206-int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vector<Operator> &inputs,197+int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs,
207- std::vector<Operator> &outputs, Graph &graph)198+ std::vector<Operator>& outputs, Graph& graph)
208{199{
209 Status ret = SUCCESS;200 Status ret = SUCCESS;
210 // 自定义代码:添加单算子定义到图中201 // 自定义代码:添加单算子定义到图中
@@ -218,7 +209,7 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
218 ADD_INPUT_ATTR(dtype, 0);209 ADD_INPUT_ATTR(dtype, 0);
219 ADD_INPUT_ATTR(seed, 10);210 ADD_INPUT_ATTR(seed, 10);
220 ADD_INPUT_ATTR(seed2, 5);211 ADD_INPUT_ATTR(seed2, 5);
221- 212+ 
222 ADD_OUTPUT(1, y, ge::DT_FLOAT, outShape);213 ADD_OUTPUT(1, y, ge::DT_FLOAT, outShape);
223 ADD_OUTPUT(2, offset, ge::DT_INT64, offsetShape);214 ADD_OUTPUT(2, offset, ge::DT_INT64, offsetShape);
224 215 
@@ -227,9 +218,9 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
227 return SUCCESS;218 return SUCCESS;
228}219}
229 220 
230-int main(int argc, char *argv[])221+int main(int argc, char* argv[])
231{222{
232- const char *graph_name = "tc_ge_irrun_test";223+ const char* graph_name = "tc_ge_irrun_test";
233 Graph graph(graph_name);224 Graph graph(graph_name);
234 std::vector<ge::Tensor> input;225 std::vector<ge::Tensor> input;
235 226 
@@ -237,7 +228,7 @@ int main(int argc, char *argv[])
237 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};228 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
238 Status ret = ge::GEInitialize(global_options);229 Status ret = ge::GEInitialize(global_options);
239 if (ret != SUCCESS) {230 if (ret != SUCCESS) {
240- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());231+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
241 return FAILED;232 return FAILED;
242 }233 }
243 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());234 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -246,7 +237,7 @@ int main(int argc, char *argv[])
246 std::vector<Operator> outputs{};237 std::vector<Operator> outputs{};
247 238 
248 std::cout << argv[1] << std::endl;239 std::cout << argv[1] << std::endl;
249- char *endptr;240+ char* endptr;
250 241 
251 DataType inDtype = DT_INT64;242 DataType inDtype = DT_INT64;
252 std::cout << inDtype << std::endl;243 std::cout << inDtype << std::endl;
@@ -265,7 +256,7 @@ int main(int argc, char *argv[])
265 256 
266 };257 };
267 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());258 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());
268- ge::Session *session = new Session(build_options);259+ ge::Session* session = new Session(build_options);
269 260 
270 if (session == nullptr) {261 if (session == nullptr) {
271 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());262 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());
@@ -288,7 +279,7 @@ int main(int argc, char *argv[])
288 std::vector<ge::Tensor> output;279 std::vector<ge::Tensor> output;
289 ret = session->RunGraph(graph_id, input, output);280 ret = session->RunGraph(graph_id, input, output);
290 if (ret != SUCCESS) {281 if (ret != SUCCESS) {
291- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());282+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
292 delete session;283 delete session;
293 GEFinalize();284 GEFinalize();
294 return FAILED;285 return FAILED;
@@ -299,23 +290,23 @@ int main(int argc, char *argv[])
299 for (int i = 0; i < input_num; i++) {290 for (int i = 0; i < input_num; i++) {
300 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;291 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;
301 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";292 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
302- uint8_t *input_data_i = input[i].GetData();293+ uint8_t* input_data_i = input[i].GetData();
303 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();294 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
304- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;295+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
305 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());296 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
306- WriteDataToFile((const char *)input_file.c_str(), data_size, input_data_i);297+ WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
307 }298 }
308 299 
309 int output_num = output.size();300 int output_num = output.size();
310 for (int i = 0; i < output_num; i++) {301 for (int i = 0; i < output_num; i++) {
311 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;302 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;
312 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";303 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
313- uint8_t *output_data_i = output[i].GetData();304+ uint8_t* output_data_i = output[i].GetData();
314 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();305 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
315- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;306+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
316 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());307 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
317- WriteDataToFile((const char *)output_file.c_str(), data_size, output_data_i);308+ WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
318- float *resultData = (float*)output_data_i;309+ float* resultData = (float*)output_data_i;
319 for (int64_t j = 0; j < output_shape; j++) {310 for (int64_t j = 0; j < output_shape; j++) {
320 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);311 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);
321 }312 }
@@ -330,7 +321,7 @@ int main(int argc, char *argv[])
330 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());321 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
331 ret = ge::GEFinalize();322 ret = ge::GEFinalize();
332 if (ret != SUCCESS) {323 if (ret != SUCCESS) {
333- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());324+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
334 return FAILED;325 return FAILED;
335 }326 }
336 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());327 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -91,7 +91,7 @@ bool RandomUniformFusionPass::MeetRequirements(const std::unique_ptr<MatchResult
91 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);91 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);
92 }92 }
93 if (version < GE_COMPILER_VERSION_900) {93 if (version < GE_COMPILER_VERSION_900) {
94- OP_LOGD(kPassName.c_str(), "GE runtime version %d < 90000000, skip pass.", version);94+ OP_LOGD(kPassName.c_str(), "GE runtime version %d < 9.0.0, skip pass.", version);
95 return false;95 return false;
96 }96 }
97 97 
@@ -37,47 +37,45 @@ using namespace ge;
37using std::map;37using std::map;
38using std::string;38using std::string;
39using std::vector;39using std::vector;
40-#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape) \40+#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape) \
41- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \41+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
42- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \42+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
43- TensorDesc placeholder##intputIndex##_desc = \43+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
44- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \44+ intputDtype); \
45- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \45+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
46- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \46+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
47- Tensor tensor_placeholder##intputIndex; \47+ Tensor tensor_placeholder##intputIndex; \
48- ret = GenOnesData( \48+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
49- placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, placeholder##intputIndex##_desc, \49+ placeholder##intputIndex##_desc, intputDtype, 2); \
50- intputDtype, 2); \50+ if (ret != SUCCESS) { \
51- if (ret != SUCCESS) { \51+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
52- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \52+ return FAILED; \
53- return FAILED; \53+ } \
54- } \54+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
55- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \55+ input.push_back(tensor_placeholder##intputIndex); \
56- input.push_back(tensor_placeholder##intputIndex); \56+ graph.AddOp(placeholder##intputIndex); \
57- graph.AddOp(placeholder##intputIndex); \57+ add1.set_input_##intputName(placeholder##intputIndex); \
58- add1.set_input_##intputName(placeholder##intputIndex); \
59 inputs.push_back(placeholder##intputIndex);58 inputs.push_back(placeholder##intputIndex);
60 59 
61-#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \60+#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \
62- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \61+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
63- auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \62+ auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \
64- TensorDesc placeholder##intputIndex##_desc = \63+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
65- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \64+ intputDtype); \
66- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \65+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
67- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \66+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
68- Tensor tensor_placeholder##intputIndex; \67+ Tensor tensor_placeholder##intputIndex; \
69- ret = GenOnesData( \68+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
70- placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, placeholder##intputIndex##_desc, \69+ placeholder##intputIndex##_desc, intputDtype, 2); \
71- intputDtype, 2); \70+ if (ret != SUCCESS) { \
72- if (ret != SUCCESS) { \71+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
73- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \72+ return FAILED; \
74- return FAILED; \73+ } \
75- } \74+ placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \
76- placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \75+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
77- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \76+ graph.AddOp(placeholder##intputIndex); \
78- graph.AddOp(placeholder##intputIndex); \77+ add1.set_input_##intputName(placeholder##intputIndex); \
79- add1.set_input_##intputName(placeholder##intputIndex); \78+ add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
80- add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
81 inputs.push_back(placeholder##intputIndex);79 inputs.push_back(placeholder##intputIndex);
82 80 
83#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \81#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \
@@ -143,8 +141,8 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorD
143 return SUCCESS;141 return SUCCESS;
144}142}
145 143 
146-int32_t GenOnesData(144+int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type,
147- vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type, int value)145+ int value)
148{146{
149 input_tensor_desc.SetRealDimCnt(shapes.size());147 input_tensor_desc.SetRealDimCnt(shapes.size());
150 size_t size = 1;148 size_t size = 1;
@@ -168,9 +166,8 @@ int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
168 return SUCCESS;166 return SUCCESS;
169}167}
170 168 
171-int CreateOppInGraph(169+int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs,
172- DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs, std::vector<Operator>& outputs,170+ std::vector<Operator>& outputs, Graph& graph)
173- Graph& graph)
174{171{
175 Status ret = SUCCESS;172 Status ret = SUCCESS;
176 // 自定义代码:添加单算子定义到图中173 // 自定义代码:添加单算子定义到图中
@@ -195,7 +192,7 @@ int main(int argc, char* argv[])
195 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};192 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
196 Status ret = ge::GEInitialize(global_options);193 Status ret = ge::GEInitialize(global_options);
197 if (ret != SUCCESS) {194 if (ret != SUCCESS) {
198- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());195+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
199 return FAILED;196 return FAILED;
200 }197 }
201 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());198 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -247,7 +244,7 @@ int main(int argc, char* argv[])
247 std::vector<ge::Tensor> output;244 std::vector<ge::Tensor> output;
248 ret = session->RunGraph(graph_id, input, output);245 ret = session->RunGraph(graph_id, input, output);
249 if (ret != SUCCESS) {246 if (ret != SUCCESS) {
250- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());247+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
251 delete session;248 delete session;
252 GEFinalize();249 GEFinalize();
253 return FAILED;250 return FAILED;
@@ -260,7 +257,7 @@ int main(int argc, char* argv[])
260 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";257 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
261 uint8_t* input_data_i = input[i].GetData();258 uint8_t* input_data_i = input[i].GetData();
262 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();259 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
263- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;260+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
264 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());261 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
265 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);262 WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
266 }263 }
@@ -271,7 +268,7 @@ int main(int argc, char* argv[])
271 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";268 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
272 uint8_t* output_data_i = output[i].GetData();269 uint8_t* output_data_i = output[i].GetData();
273 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();270 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
274- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;271+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
275 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());272 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
276 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);273 WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
277 }274 }
@@ -286,7 +283,7 @@ int main(int argc, char* argv[])
286 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());283 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
287 ret = ge::GEFinalize();284 ret = ge::GEFinalize();
288 if (ret != SUCCESS) {285 if (ret != SUCCESS) {
289- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());286+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
290 return FAILED;287 return FAILED;
291 }288 }
292 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());289 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -66,7 +66,7 @@ void SimThreadExponentialTiling::PrintInfo()
66 OP_LOGD(nodeName, "range = %f.", tiling.get_range());66 OP_LOGD(nodeName, "range = %f.", tiling.get_range());
67 OP_LOGD(nodeName, "handleNumLoop = %u.", tiling.get_handleNumLoop());67 OP_LOGD(nodeName, "handleNumLoop = %u.", tiling.get_handleNumLoop());
68 OP_LOGD(nodeName, "handleNumTail = %u.", tiling.get_handleNumTail());68 OP_LOGD(nodeName, "handleNumTail = %u.", tiling.get_handleNumTail());
69- OP_LOGD(nodeName, "state = %u.", tiling.get_state());69+ OP_LOGD(nodeName, "state = %lu.", tiling.get_state());
70 OP_LOGD(nodeName, "start = %f.", tiling.get_start());70 OP_LOGD(nodeName, "start = %f.", tiling.get_start());
71 OP_LOGD(nodeName, "end = %f.", tiling.get_end());71 OP_LOGD(nodeName, "end = %f.", tiling.get_end());
72 OP_LOGD(nodeName, "lambda = %f.", tiling.get_lambda());72 OP_LOGD(nodeName, "lambda = %f.", tiling.get_lambda());
@@ -162,7 +162,8 @@ ge::graphStatus SimThreadExponentialTiling::GetInputTensorInfo()
162 selfDType = selfDesc->GetDataType();162 selfDType = selfDesc->GetDataType();
163 GetDataTypeKey(selfDType);163 GetDataTypeKey(selfDType);
164 OP_CHECK_IF(GetDataTypeKey(selfDType) == false,164 OP_CHECK_IF(GetDataTypeKey(selfDType) == false,
165- OP_LOGE(nodeName, "The dtype of input self must be in [float32, float16, bfloat16]."),165+ OP_LOGE(nodeName, "The dtype %d of input self must be in [float32, float16, bfloat16].",
166+ static_cast<int>(selfDType)),
166 return ge::GRAPH_FAILED);167 return ge::GRAPH_FAILED);
167 168 
168 return ge::GRAPH_SUCCESS;169 return ge::GRAPH_SUCCESS;
@@ -195,7 +196,7 @@ ge::graphStatus SimThreadExponentialTiling::Tiling4Block()
195 // 分核计算196 // 分核计算
196 useCoreNum = static_cast<int64_t>(197 useCoreNum = static_cast<int64_t>(
197 Ops::Base::CeilDiv(batchNumTotal, Ops::Base::CeilDiv(batchNumTotal, totalCoreNum)));198 Ops::Base::CeilDiv(batchNumTotal, Ops::Base::CeilDiv(batchNumTotal, totalCoreNum)));
198- OP_CHECK_IF(useCoreNum == 0, OP_LOGE(nodeName, "useCoreNum %u must be not equal to 0.", useCoreNum),199+ OP_CHECK_IF(useCoreNum == 0, OP_LOGE(nodeName, "useCoreNum %u must not be equal to 0.", useCoreNum),
199 return ge::GRAPH_FAILED);200 return ge::GRAPH_FAILED);
200 // useCoreNum = static_cast<int64_t>(CeilDiv(batchNumTotal, CeilDiv(batchNumTotal, totalCoreNum)));201 // useCoreNum = static_cast<int64_t>(CeilDiv(batchNumTotal, CeilDiv(batchNumTotal, totalCoreNum)));
201 batchNumPerCore = (batchNumTotal + useCoreNum - 1) / useCoreNum;202 batchNumPerCore = (batchNumTotal + useCoreNum - 1) / useCoreNum;
@@ -223,7 +224,7 @@ ge::graphStatus SimThreadExponentialTiling::SetAttrParams()
223 OP_CHECK_NULL_WITH_CONTEXT(context, lambdaPtr);224 OP_CHECK_NULL_WITH_CONTEXT(context, lambdaPtr);
224 lambda = static_cast<float>(*lambdaPtr);225 lambda = static_cast<float>(*lambdaPtr);
225 OP_CHECK_IF(lambda == 0,226 OP_CHECK_IF(lambda == 0,
226- OP_LOGE(context->GetNodeName(), "lambda is the denominator and cannot be zero, but get %f.", lambda),227+ OP_LOGE(context->GetNodeName(), "lambda is the denominator and cannot be zero, but got %f.", lambda),
227 return ge::GRAPH_FAILED);228 return ge::GRAPH_FAILED);
228 const int64_t* seedPtr = attrs->GetAttrPointer<int64_t>(ATTR_2);229 const int64_t* seedPtr = attrs->GetAttrPointer<int64_t>(ATTR_2);
229 OP_CHECK_NULL_WITH_CONTEXT(context, seedPtr);230 OP_CHECK_NULL_WITH_CONTEXT(context, seedPtr);
@@ -278,7 +279,7 @@ ge::graphStatus SimThreadExponentialTiling::DoTiling()
278static ge::graphStatus Tiling4SimThreadExponential(gert::TilingContext* context)279static ge::graphStatus Tiling4SimThreadExponential(gert::TilingContext* context)
279{280{
280 auto nodeName = context->GetNodeName();281 auto nodeName = context->GetNodeName();
281- OP_LOGD(nodeName, "Tiling4SimThreadExponential running begin.");282+ OP_LOGD(nodeName, "Tiling4SimThreadExponential started.");
282 283 
283 SimThreadExponentialTiling tilingObj(context);284 SimThreadExponentialTiling tilingObj(context);
284 return tilingObj.DoTiling();285 return tilingObj.DoTiling();
@@ -287,7 +288,7 @@ static ge::graphStatus Tiling4SimThreadExponential(gert::TilingContext* context)
287ge::graphStatus TilingPrepare4SimThreadExponential(gert::TilingParseContext* context)288ge::graphStatus TilingPrepare4SimThreadExponential(gert::TilingParseContext* context)
288{289{
289 auto nodeName = context->GetNodeName();290 auto nodeName = context->GetNodeName();
290- OP_LOGD(nodeName, "TilingPrepare4SimThreadExponential running end.");291+ OP_LOGD(nodeName, "TilingPrepare4SimThreadExponential finished.");
291 292 
292 return ge::GRAPH_SUCCESS;293 return ge::GRAPH_SUCCESS;
293}294}
@@ -10,14 +10,8 @@
10# See LICENSE in the root of the software repository for the full text of the License.10# See LICENSE in the root of the software repository for the full text of the License.
11# ----------------------------------------------------------------------------11# ----------------------------------------------------------------------------
12 12 
13-import numpy as np
14 13 
15- 14+__golden__ = {"kernel": {"sim_thread_exponential": "sim_thread_exponential_golden"}}
16-__golden__ = {
17- "kernel": {
18- "sim_thread_exponential": "sim_thread_exponential_golden"
19- }
20-}
21 15 
22 16 
23class MaxPool3DGradGoldenGpuClient:17class MaxPool3DGradGoldenGpuClient:
@@ -35,10 +29,11 @@ class MaxPool3DGradGoldenGpuClient:
35 import struct29 import struct
36 import torch30 import torch
37 import numpy as np31 import numpy as np
32+ 
38 self._deps_loaded = True33 self._deps_loaded = True
39 34 
40 def _recv_all(self, sock, n):35 def _recv_all(self, sock, n):
41- data = b''36+ data = b""
42 while len(data) < n:37 while len(data) < n:
43 packet = sock.recv(n - len(data))38 packet = sock.recv(n - len(data))
44 if not packet:39 if not packet:
@@ -48,19 +43,24 @@ class MaxPool3DGradGoldenGpuClient:
48 43 
49 def _send_msg(self, sock, msg):44 def _send_msg(self, sock, msg):
50 msg = pickle.dumps(msg)45 msg = pickle.dumps(msg)
51- msg = struct.pack('>I', len(msg)) + msg46+ msg = struct.pack(">I", len(msg)) + msg
52 sock.sendall(msg)47 sock.sendall(msg)
53 48 
54 def _recv_msg(self, sock):49 def _recv_msg(self, sock):
55 raw_msglen = self._recv_all(sock, 4)50 raw_msglen = self._recv_all(sock, 4)
56 if not raw_msglen:51 if not raw_msglen:
57 return None52 return None
58- msglen = struct.unpack('>I', raw_msglen)[0]53+ msglen = struct.unpack(">I", raw_msglen)[0]
59 return pickle.loads(self._recv_all(sock, msglen))54 return pickle.loads(self._recv_all(sock, msglen))
60 55 
61 def compute_on_gpu(self, attr_count, attr_seed, attr_offset, attr_lambd, dtype):56 def compute_on_gpu(self, attr_count, attr_seed, attr_offset, attr_lambd, dtype):
62- request = {"attr_count": attr_count, "attr_seed": attr_seed, "attr_offset": attr_offset,57+ request = {
63- "attr_lambd": attr_lambd, "dtype": dtype}58+ "attr_count": attr_count,
59+ "attr_seed": attr_seed,
60+ "attr_offset": attr_offset,
61+ "attr_lambd": attr_lambd,
62+ "dtype": dtype,
63+ }
64 try:64 try:
65 with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:65 with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
66 s.settimeout(3000)66 s.settimeout(3000)
@@ -69,18 +69,18 @@ class MaxPool3DGradGoldenGpuClient:
69 result = self._recv_msg(s)69 result = self._recv_msg(s)
70 return result70 return result
71 except Exception as e:71 except Exception as e:
72- print(f"连接错误: {e}")72+ print(f"Connection error: {e}")
73 73 
74 74 
75def sim_thread_exponential_golden(self, count, lambd=1.0, seed=0, offset=0, **kwargs):75def sim_thread_exponential_golden(self, count, lambd=1.0, seed=0, offset=0, **kwargs):
76- '''76+ """
77 Kernel golden for sim_thread_exponential.77 Kernel golden for sim_thread_exponential.
78 All the parameters follow @sim_thread_exponential_def.cpp without outputs.78 All the parameters follow @sim_thread_exponential_def.cpp without outputs.
79 All the input Tensors are numpy.ndarray.79 All the input Tensors are numpy.ndarray.
80 kwargs may contain: short_soc_version, input_ori_shapes, output_ori_shapes,80 kwargs may contain: short_soc_version, input_ori_shapes, output_ori_shapes,
81 input_formats, output_formats, input_ori_formats, output_ori_formats,81 input_formats, output_formats, input_ori_formats, output_ori_formats,
82 input_dtypes, output_dtypes.82 input_dtypes, output_dtypes.
83- '''83+ """
84 input_dtypes = kwargs.get("input_dtypes", [])84 input_dtypes = kwargs.get("input_dtypes", [])
85 dtype = input_dtypes[0] if input_dtypes else "float32"85 dtype = input_dtypes[0] if input_dtypes else "float32"
86 86 
@@ -127,7 +127,7 @@ bool BernoulliFusionPass::MeetRequirements(const std::unique_ptr<MatchResult>& m
127 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);127 aclsysGetVersionNum(const_cast<char*>("ge_compiler"), &version);
128 }128 }
129 if (version < GE_COMPILER_VERSION_900) {129 if (version < GE_COMPILER_VERSION_900) {
130- OP_LOGD(kPassName.c_str(), "GE runtime version %d < 90000000, skip pass.", version);130+ OP_LOGD(kPassName.c_str(), "GE runtime version %d < 9.0.0, skip pass.", version);
131 return false;131 return false;
132 }132 }
133 133 
@@ -216,7 +216,7 @@ std::unique_ptr<Graph> BernoulliFusionPass::Replacement(const std::unique_ptr<Ma
216 matchResult->ToSubgraphBoundary()->GetAllInputs(subgraphInputs);216 matchResult->ToSubgraphBoundary()->GetAllInputs(subgraphInputs);
217 GraphUniqPtr replaceGraph = replaceGraphBuilder.BuildAndReset({output});217 GraphUniqPtr replaceGraph = replaceGraphBuilder.BuildAndReset({output});
218 if (InferShape(replaceGraph, subgraphInputs) != SUCCESS) {218 if (InferShape(replaceGraph, subgraphInputs) != SUCCESS) {
219- OP_LOGE(kPassName.c_str(), "Infershape failed.");219+ OP_LOGE(kPassName.c_str(), "InferShape failed.");
220 return nullptr;220 return nullptr;
221 }221 }
222 return replaceGraph;222 return replaceGraph;
@@ -42,9 +42,12 @@ OpTilingConfig StatelessBernoulliTiling::BuildOpConfig()
42 {INPUT_IDX_OFFSET, {{ge::DT_INT64}, 1, {}, nullptr}},42 {INPUT_IDX_OFFSET, {{ge::DT_INT64}, 1, {}, nullptr}},
43 };43 };
44 config.outputCheckRules = {44 config.outputCheckRules = {
45- {OUTPUT_IDX_Y, {{ge::DT_INT8, ge::DT_UINT8, ge::DT_INT16, ge::DT_UINT16,45+ {OUTPUT_IDX_Y,
46- ge::DT_INT32, ge::DT_UINT32, ge::DT_INT64, ge::DT_UINT64, ge::DT_FLOAT,46+ {{ge::DT_INT8, ge::DT_UINT8, ge::DT_INT16, ge::DT_UINT16, ge::DT_INT32, ge::DT_UINT32, ge::DT_INT64,
47- ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BOOL}, -1, {}, nullptr}},47+ ge::DT_UINT64, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BOOL},
48+ -1,
49+ {},
50+ nullptr}},
48 };51 };
49 52 
50 config.getOutputSize = [](gert::TilingContext* ctx, int64_t& size) {53 config.getOutputSize = [](gert::TilingContext* ctx, int64_t& size) {
@@ -65,7 +68,8 @@ OpTilingConfig StatelessBernoulliTiling::BuildOpConfig()
65 if (offset % OFFSET_MULTIPLE != 0) {68 if (offset % OFFSET_MULTIPLE != 0) {
66 std::string valueStr = std::to_string(offset);69 std::string valueStr = std::to_string(offset);
67 std::string reasonMsg = "offset value must be a multiple of 4";70 std::string reasonMsg = "offset value must be a multiple of 4";
68- OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(ctx->GetNodeName(), "input offset", valueStr.c_str(), reasonMsg.c_str());71+ OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(ctx->GetNodeName(), "input offset", valueStr.c_str(),
72+ reasonMsg.c_str());
69 return ge::GRAPH_FAILED;73 return ge::GRAPH_FAILED;
70 }74 }
71 return ge::GRAPH_SUCCESS;75 return ge::GRAPH_SUCCESS;
@@ -79,8 +83,8 @@ OpTilingConfig StatelessBernoulliTiling::BuildOpConfig()
79 83 
80ge::graphStatus StatelessBernoulliTiling::DoSimtBlockTiling()84ge::graphStatus StatelessBernoulliTiling::DoSimtBlockTiling()
81{85{
82- OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),86+ OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is %ld, must be greater than 0.", totalCoreNum_),
83- return ge::GRAPH_FAILED);87+ return ge::GRAPH_FAILED);
84 88 
85 auto probTensor = context_->GetRequiredInputTensor(INPUT_IDX_PROB);89 auto probTensor = context_->GetRequiredInputTensor(INPUT_IDX_PROB);
86 OP_CHECK_NULL_WITH_CONTEXT(context_, probTensor);90 OP_CHECK_NULL_WITH_CONTEXT(context_, probTensor);
@@ -220,7 +220,7 @@ int main(int argc, char* argv[])
220 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};220 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
221 Status ret = ge::GEInitialize(global_options);221 Status ret = ge::GEInitialize(global_options);
222 if (ret != SUCCESS) {222 if (ret != SUCCESS) {
223- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());223+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
224 return FAILED;224 return FAILED;
225 }225 }
226 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());226 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -265,7 +265,7 @@ int main(int argc, char* argv[])
265 std::vector<ge::Tensor> output;265 std::vector<ge::Tensor> output;
266 ret = session->RunGraph(graph_id, input, output);266 ret = session->RunGraph(graph_id, input, output);
267 if (ret != SUCCESS) {267 if (ret != SUCCESS) {
268- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());268+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
269 delete session;269 delete session;
270 GEFinalize();270 GEFinalize();
271 return FAILED;271 return FAILED;
@@ -409,7 +409,7 @@ int main(int argc, char* argv[])
409 delete session;409 delete session;
410 ret = ge::GEFinalize();410 ret = ge::GEFinalize();
411 if (ret != SUCCESS) {411 if (ret != SUCCESS) {
412- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());412+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
413 return FAILED;413 return FAILED;
414 }414 }
415 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());415 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -37,87 +37,78 @@ using namespace ge;
37using std::map;37using std::map;
38using std::string;38using std::string;
39using std::vector;39using std::vector;
40-#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \40+#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
41- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \41+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
42- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \42+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
43- TensorDesc placeholder##intputIndex##_desc = \43+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
44- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \44+ intputDtype); \
45- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \45+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
46- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \46+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
47- Tensor tensor_placeholder##intputIndex; \47+ Tensor tensor_placeholder##intputIndex; \
48- ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, \48+ ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
49- tensor_placeholder##intputIndex, \49+ placeholder##intputIndex##_desc, value); \
50- placeholder##intputIndex##_desc, \50+ if (ret != SUCCESS) { \
51- value); \51+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
52- if (ret != SUCCESS) { \52+ return FAILED; \
53- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \53+ } \
54- return FAILED; \54+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
55- } \55+ input.push_back(tensor_placeholder##intputIndex); \
56- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \56+ graph.AddOp(placeholder##intputIndex); \
57- input.push_back(tensor_placeholder##intputIndex); \57+ add1.set_input_##intputName(placeholder##intputIndex); \
58- graph.AddOp(placeholder##intputIndex); \
59- add1.set_input_##intputName(placeholder##intputIndex); \
60 inputs.push_back(placeholder##intputIndex)58 inputs.push_back(placeholder##intputIndex)
61 59 
62-#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \60+#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
63- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \61+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
64- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \62+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
65- TensorDesc placeholder##intputIndex##_desc = \63+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
66- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \64+ intputDtype); \
67- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \65+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
68- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \66+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
69- Tensor tensor_placeholder##intputIndex; \67+ Tensor tensor_placeholder##intputIndex; \
70- ret = GenOnesData(placeholder##intputIndex##_shape, \68+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
71- tensor_placeholder##intputIndex, \69+ placeholder##intputIndex##_desc, intputDtype, value); \
72- placeholder##intputIndex##_desc, \70+ if (ret != SUCCESS) { \
73- intputDtype, \71+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
74- value); \72+ return FAILED; \
75- if (ret != SUCCESS) { \73+ } \
76- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \74+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
77- return FAILED; \75+ input.push_back(tensor_placeholder##intputIndex); \
78- } \76+ graph.AddOp(placeholder##intputIndex); \
79- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \77+ add1.set_input_##intputName(placeholder##intputIndex); \
80- input.push_back(tensor_placeholder##intputIndex); \
81- graph.AddOp(placeholder##intputIndex); \
82- add1.set_input_##intputName(placeholder##intputIndex); \
83 inputs.push_back(placeholder##intputIndex)78 inputs.push_back(placeholder##intputIndex)
84 79 
85-#define ADD_INPUT_ATTR(attrName, attrValue) \80+#define ADD_INPUT_ATTR(attrName, attrValue) add1.set_attr_##attrName(attrValue);
86- add1.set_attr_##attrName(attrValue);
87 81 
88-#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \82+#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \
89- TensorDesc outputName##outputIndex##_desc = \83+ TensorDesc outputName##outputIndex##_desc = TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
90- TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
91 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)84 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)
92 85 
93-#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \86+#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape) \
94- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \87+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
95- auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \88+ auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \
96- TensorDesc placeholder##intputIndex##_desc = \89+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
97- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \90+ intputDtype); \
98- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \91+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
99- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \92+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
100- Tensor tensor_placeholder##intputIndex; \93+ Tensor tensor_placeholder##intputIndex; \
101- ret = GenOnesData(placeholder##intputIndex##_shape, \94+ ret = GenOnesData(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
102- tensor_placeholder##intputIndex, \95+ placeholder##intputIndex##_desc, intputDtype, 1); \
103- placeholder##intputIndex##_desc, \96+ if (ret != SUCCESS) { \
104- intputDtype, \97+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
105- 1); \98+ return FAILED; \
106- if (ret != SUCCESS) { \99+ } \
107- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \100+ placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \
108- return FAILED; \101+ \
109- } \102+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
110- placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \ 103+ graph.AddOp(placeholder##intputIndex); \
111- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \104+ add1.set_input_##intputName(placeholder##intputIndex); \
112- graph.AddOp(placeholder##intputIndex); \105+ add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
113- add1.set_input_##intputName(placeholder##intputIndex); \
114- add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
115 inputs.push_back(placeholder##intputIndex)106 inputs.push_back(placeholder##intputIndex)
116 107 
117-#define LOG_PRINT(message, ...) \108+#define LOG_PRINT(message, ...) \
118- do { \109+ do { \
119- printf(message, ##__VA_ARGS__); \110+ printf(message, ##__VA_ARGS__); \
120- } while (0)111+ } while (0)
121 112 
122string GetTime()113string GetTime()
123{114{
@@ -160,7 +151,7 @@ uint32_t GetDataTypeSize(DataType dt)
160 return dilation;151 return dilation;
161}152}
162 153 
163-int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, float value)154+int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, float value)
164{155{
165 input_tensor_desc.SetRealDimCnt(shapes.size());156 input_tensor_desc.SetRealDimCnt(shapes.size());
166 size_t size = 1;157 size_t size = 1;
@@ -169,17 +160,17 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorD
169 }160 }
170 uint32_t byteSizeFloat32 = 4;161 uint32_t byteSizeFloat32 = 4;
171 uint32_t data_len = size * byteSizeFloat32;162 uint32_t data_len = size * byteSizeFloat32;
172- float *pData = new (std::nothrow) float[size];163+ float* pData = new (std::nothrow) float[size];
173 164 
174 for (size_t i = 0; i < size; ++i) {165 for (size_t i = 0; i < size; ++i) {
175 *(pData + i) = value;166 *(pData + i) = value;
176 }167 }
177- input_tensor = Tensor(input_tensor_desc, (uint8_t *)pData, data_len);168+ input_tensor = Tensor(input_tensor_desc, (uint8_t*)pData, data_len);
178 return SUCCESS;169 return SUCCESS;
179}170}
180 171 
181-int32_t GenOnesData(172+int32_t GenOnesData(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, DataType data_type,
182- vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, DataType data_type, int value)173+ int value)
183{174{
184 input_tensor_desc.SetRealDimCnt(shapes.size());175 input_tensor_desc.SetRealDimCnt(shapes.size());
185 size_t size = 1;176 size_t size = 1;
@@ -187,24 +178,24 @@ int32_t GenOnesData(
187 size *= shapes[i];178 size *= shapes[i];
188 }179 }
189 uint32_t data_len = size * GetDataTypeSize(data_type);180 uint32_t data_len = size * GetDataTypeSize(data_type);
190- int32_t *pData = new (std::nothrow) int32_t[data_len];181+ int32_t* pData = new (std::nothrow) int32_t[data_len];
191 for (uint32_t i = 0; i < size; ++i) {182 for (uint32_t i = 0; i < size; ++i) {
192 *(pData + i) = value;183 *(pData + i) = value;
193 }184 }
194- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);185+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
195 return SUCCESS;186 return SUCCESS;
196}187}
197 188 
198-int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t *inputData)189+int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
199{190{
200- FILE *fp = fopen(bin_file.c_str(), "w");191+ FILE* fp = fopen(bin_file.c_str(), "w");
201 fwrite(inputData, sizeof(uint8_t), data_size, fp);192 fwrite(inputData, sizeof(uint8_t), data_size, fp);
202 fclose(fp);193 fclose(fp);
203 return SUCCESS;194 return SUCCESS;
204}195}
205 196 
206-int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vector<Operator> &inputs,197+int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor>& input, std::vector<Operator>& inputs,
207- std::vector<Operator> &outputs, Graph &graph)198+ std::vector<Operator>& outputs, Graph& graph)
208{199{
209 Status ret = SUCCESS;200 Status ret = SUCCESS;
210 // 自定义代码:添加单算子定义到图中201 // 自定义代码:添加单算子定义到图中
@@ -213,7 +204,7 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
213 std::vector<int64_t> countShape = {1};204 std::vector<int64_t> countShape = {1};
214 std::vector<int64_t> seedShape = {1};205 std::vector<int64_t> seedShape = {1};
215 std::vector<int64_t> offsetShape = {1};206 std::vector<int64_t> offsetShape = {1};
216- std::vector<int64_t> yShape = {2,1};207+ std::vector<int64_t> yShape = {2, 1};
217 std::vector<int64_t> maskShape = {2};208 std::vector<int64_t> maskShape = {2};
218 209 
219 ADD_INPUT(1, x, ge::DT_BOOL, xShape);210 ADD_INPUT(1, x, ge::DT_BOOL, xShape);
@@ -228,9 +219,9 @@ int CreateOppInGraph(DataType inDtype, std::vector<ge::Tensor> &input, std::vect
228 return SUCCESS;219 return SUCCESS;
229}220}
230 221 
231-int main(int argc, char *argv[])222+int main(int argc, char* argv[])
232{223{
233- const char *graph_name = "tc_ge_irrun_test";224+ const char* graph_name = "tc_ge_irrun_test";
234 Graph graph(graph_name);225 Graph graph(graph_name);
235 std::vector<ge::Tensor> input;226 std::vector<ge::Tensor> input;
236 227 
@@ -238,7 +229,7 @@ int main(int argc, char *argv[])
238 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};229 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
239 Status ret = ge::GEInitialize(global_options);230 Status ret = ge::GEInitialize(global_options);
240 if (ret != SUCCESS) {231 if (ret != SUCCESS) {
241- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());232+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
242 return FAILED;233 return FAILED;
243 }234 }
244 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());235 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -247,7 +238,7 @@ int main(int argc, char *argv[])
247 std::vector<Operator> outputs{};238 std::vector<Operator> outputs{};
248 239 
249 std::cout << argv[1] << std::endl;240 std::cout << argv[1] << std::endl;
250- char *endptr;241+ char* endptr;
251 242 
252 DataType inDtype = DT_INT64;243 DataType inDtype = DT_INT64;
253 std::cout << inDtype << std::endl;244 std::cout << inDtype << std::endl;
@@ -266,7 +257,7 @@ int main(int argc, char *argv[])
266 257 
267 };258 };
268 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());259 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());
269- ge::Session *session = new Session(build_options);260+ ge::Session* session = new Session(build_options);
270 261 
271 if (session == nullptr) {262 if (session == nullptr) {
272 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());263 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());
@@ -289,7 +280,7 @@ int main(int argc, char *argv[])
289 std::vector<ge::Tensor> output;280 std::vector<ge::Tensor> output;
290 ret = session->RunGraph(graph_id, input, output);281 ret = session->RunGraph(graph_id, input, output);
291 if (ret != SUCCESS) {282 if (ret != SUCCESS) {
292- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());283+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
293 delete session;284 delete session;
294 GEFinalize();285 GEFinalize();
295 return FAILED;286 return FAILED;
@@ -300,23 +291,23 @@ int main(int argc, char *argv[])
300 for (int i = 0; i < input_num; i++) {291 for (int i = 0; i < input_num; i++) {
301 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;292 std::cout << "input " << i << " dtype : " << input[i].GetTensorDesc().GetDataType() << std::endl;
302 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";293 string input_file = "./tc_ge_irrun_test_0008_npu_input_" + std::to_string(i) + ".bin";
303- uint8_t *input_data_i = input[i].GetData();294+ uint8_t* input_data_i = input[i].GetData();
304 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();295 int64_t input_shape = input[i].GetTensorDesc().GetShape().GetShapeSize();
305- std::cout << "this is " << i << "th input, input shape size =" << input_shape << std::endl;296+ std::cout << "this is input " << i << ", input shape size =" << input_shape << std::endl;
306 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());297 uint32_t data_size = input_shape * GetDataTypeSize(input[i].GetTensorDesc().GetDataType());
307- WriteDataToFile((const char *)input_file.c_str(), data_size, input_data_i);298+ WriteDataToFile((const char*)input_file.c_str(), data_size, input_data_i);
308 }299 }
309 300 
310 int output_num = output.size();301 int output_num = output.size();
311 for (int i = 0; i < output_num; i++) {302 for (int i = 0; i < output_num; i++) {
312 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;303 std::cout << "output " << i << " dtype : " << output[i].GetTensorDesc().GetDataType() << std::endl;
313 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";304 string output_file = "./tc_ge_irrun_test_0008_npu_output_" + std::to_string(i) + ".bin";
314- uint8_t *output_data_i = output[i].GetData();305+ uint8_t* output_data_i = output[i].GetData();
315 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();306 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
316- std::cout << "this is " << i << "th output, output shape size =" << output_shape << std::endl;307+ std::cout << "this is output " << i << ", output shape size =" << output_shape << std::endl;
317 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());308 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
318- WriteDataToFile((const char *)output_file.c_str(), data_size, output_data_i);309+ WriteDataToFile((const char*)output_file.c_str(), data_size, output_data_i);
319- float *resultData = (float*)output_data_i;310+ float* resultData = (float*)output_data_i;
320 for (int64_t j = 0; j < output_shape; j++) {311 for (int64_t j = 0; j < output_shape; j++) {
321 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);312 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);
322 }313 }
@@ -331,7 +322,7 @@ int main(int argc, char *argv[])
331 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());322 printf("%s - INFO - [XIR]: Start to finalize ir graph session\n", GetTime().c_str());
332 ret = ge::GEFinalize();323 ret = ge::GEFinalize();
333 if (ret != SUCCESS) {324 if (ret != SUCCESS) {
334- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());325+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
335 return FAILED;326 return FAILED;
336 }327 }
337 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());328 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());
@@ -21,17 +21,14 @@ namespace optiling {
21 21 
22const std::set<ge::DataType> SUPPORT_DTYPE = {ge::DT_BOOL};22const std::set<ge::DataType> SUPPORT_DTYPE = {ge::DT_BOOL};
23 23 
24-bool StatelessRandomChoiceWithMaskSimtTiling::IsCapable()24+bool StatelessRandomChoiceWithMaskSimtTiling::IsCapable() { return true; }
25-{
26- return true;
27-}
28 25 
29ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::GetPlatformInfo()26ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::GetPlatformInfo()
30{27{
31 auto platformPtr = context_->GetPlatformInfo();28 auto platformPtr = context_->GetPlatformInfo();
32 if (platformPtr == nullptr) {29 if (platformPtr == nullptr) {
33- auto compileInfoPtr =30+ auto compileInfoPtr = reinterpret_cast<const StatelessRandomChoiceWithMaskCompileInfo*>(
34- reinterpret_cast<const StatelessRandomChoiceWithMaskCompileInfo*>(context_->GetCompileInfo());31+ context_->GetCompileInfo());
35 OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_, "compile info is null"), return ge::GRAPH_FAILED);32 OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_, "compile info is null"), return ge::GRAPH_FAILED);
36 coreNum_ = compileInfoPtr->coreNum;33 coreNum_ = compileInfoPtr->coreNum;
37 ubSize_ = compileInfoPtr->ubSize;34 ubSize_ = compileInfoPtr->ubSize;
@@ -44,12 +41,10 @@ ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::GetPlatformInfo()
44 ubSize_ = static_cast<uint64_t>(ubSizePlatform);41 ubSize_ = static_cast<uint64_t>(ubSizePlatform);
45 }42 }
46 ubSize_ = ubSize_ - DCACHE_SIZE;43 ubSize_ = ubSize_ - DCACHE_SIZE;
47- OP_CHECK_IF(44+ OP_CHECK_IF((coreNum_ <= 0 || ubSize_ <= 0),
48- (coreNum_ <= 0 || ubSize_ <= 0),45+ OP_LOGE(context_, "coreNum and ubSize should be greater than 0, but got coreNum [%ld] and ubSize [%ld]",
49- OP_LOGE(46+ coreNum_, ubSize_),
50- context_, "coreNum and ubSize should not be samller than 0, but got coreNum [%ld] and ubSize [%ld]",47+ return ge::GRAPH_FAILED);
51- coreNum_, ubSize_),
52- return ge::GRAPH_FAILED);
53 return ge::GRAPH_SUCCESS;48 return ge::GRAPH_SUCCESS;
54}49}
55 50 
@@ -61,9 +56,9 @@ ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::GetShapeAttrsInfo()
61 inputDim_ = xShape_.GetDimNum();56 inputDim_ = xShape_.GetDimNum();
62 if (inputDim_ > INPUT_X_MAX_DIM_NUM || inputDim_ < INPUT_X_MIN_DIM_NUM) {57 if (inputDim_ > INPUT_X_MAX_DIM_NUM || inputDim_ < INPUT_X_MIN_DIM_NUM) {
63 std::string valueStr = std::to_string(xShape_.GetDimNum());58 std::string valueStr = std::to_string(xShape_.GetDimNum());
64- std::string reasonMsg = "xDimNum should be greater than 1 and smaller than 5";59+ std::string reasonMsg = "xDimNum should be in range [1, 5]";
65- OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(60+ OP_LOGE_FOR_INVALID_SHAPEDIM_WITH_REASON(context_->GetNodeName(), "input tensor x", valueStr.c_str(),
66- context_->GetNodeName(), "input tensor x", valueStr.c_str(), reasonMsg.c_str());61+ reasonMsg.c_str());
67 return ge::GRAPH_FAILED;62 return ge::GRAPH_FAILED;
68 }63 }
69 inputSize_ = xShape_.GetShapeSize();64 inputSize_ = xShape_.GetShapeSize();
@@ -73,9 +68,9 @@ ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::GetShapeAttrsInfo()
73 OP_CHECK_NULL_WITH_CONTEXT(context_, seedValue);68 OP_CHECK_NULL_WITH_CONTEXT(context_, seedValue);
74 if (seed->GetShapeSize() <= 0) {69 if (seed->GetShapeSize() <= 0) {
75 std::string valueStr = std::to_string(seed->GetShapeSize());70 std::string valueStr = std::to_string(seed->GetShapeSize());
76- std::string reasonMsg = "inputSeed shapeSize need be greater than 0";71+ std::string reasonMsg = "inputSeed shapeSize must be greater than 0";
77- OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(72+ OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(context_->GetNodeName(), "input seed", valueStr.c_str(),
78- context_->GetNodeName(), "input seed", valueStr.c_str(), reasonMsg.c_str());73+ reasonMsg.c_str());
79 return ge::GRAPH_FAILED;74 return ge::GRAPH_FAILED;
80 }75 }
81 seed_ = seed->GetData<int64_t>()[0];76 seed_ = seed->GetData<int64_t>()[0];
@@ -85,9 +80,9 @@ ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::GetShapeAttrsInfo()
85 OP_CHECK_NULL_WITH_CONTEXT(context_, offsetValue);80 OP_CHECK_NULL_WITH_CONTEXT(context_, offsetValue);
86 if (offset->GetShapeSize() <= 0) {81 if (offset->GetShapeSize() <= 0) {
87 std::string valueStr = std::to_string(offset->GetShapeSize());82 std::string valueStr = std::to_string(offset->GetShapeSize());
88- std::string reasonMsg = "inputOffset shapeSize need be greater than 0";83+ std::string reasonMsg = "inputOffset shapeSize must be greater than 0";
89- OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(84+ OP_LOGE_FOR_INVALID_SHAPESIZE_WITH_REASON(context_->GetNodeName(), "input offset", valueStr.c_str(),
90- context_->GetNodeName(), "input offset", valueStr.c_str(), reasonMsg.c_str());85+ reasonMsg.c_str());
91 return ge::GRAPH_FAILED;86 return ge::GRAPH_FAILED;
92 }87 }
93 OP_CHECK_NULL_WITH_CONTEXT(context_, offset);88 OP_CHECK_NULL_WITH_CONTEXT(context_, offset);
@@ -96,8 +91,9 @@ ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::GetShapeAttrsInfo()
96 ge::DataType xDtype = xDesc->GetDataType();91 ge::DataType xDtype = xDesc->GetDataType();
97 if (SUPPORT_DTYPE.count(xDtype) == 0) {92 if (SUPPORT_DTYPE.count(xDtype) == 0) {
98 std::string valueStr = Ops::Base::ToString(xDtype);93 std::string valueStr = Ops::Base::ToString(xDtype);
99- std::string reasonMsg = "input x dtype only support BOOL currently";94+ std::string reasonMsg = "input x dtype only supports BOOL currently";
100- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor x", valueStr.c_str(), reasonMsg.c_str());95+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context_->GetNodeName(), "input tensor x", valueStr.c_str(),
96+ reasonMsg.c_str());
101 return ge::GRAPH_FAILED;97 return ge::GRAPH_FAILED;
102 }98 }
103 auto y = context_->GetOutputShape(OUTPUT_Y_IDX);99 auto y = context_->GetOutputShape(OUTPUT_Y_IDX);
@@ -120,8 +116,8 @@ ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::ComputeCoreNum()
120 116 
121ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::SetTilingData()117ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::SetTilingData()
122{118{
123- StatelessRandomChoiceWithMaskSimtTilingData* tilingData =119+ StatelessRandomChoiceWithMaskSimtTilingData*
124- context_->GetTilingData<StatelessRandomChoiceWithMaskSimtTilingData>();120+ tilingData = context_->GetTilingData<StatelessRandomChoiceWithMaskSimtTilingData>();
125 for (int64_t i = 0; i < static_cast<int64_t>(inputDim_); i++) {121 for (int64_t i = 0; i < static_cast<int64_t>(inputDim_); i++) {
126 tilingData->inputShape[i] = xShape_.GetDim(i);122 tilingData->inputShape[i] = xShape_.GetDim(i);
127 }123 }
@@ -155,10 +151,7 @@ ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::DoOpTiling()
155 return ge::GRAPH_SUCCESS;151 return ge::GRAPH_SUCCESS;
156}152}
157 153 
158-ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::DoLibApiTiling()154+ge::graphStatus StatelessRandomChoiceWithMaskSimtTiling::DoLibApiTiling() { return ge::GRAPH_SUCCESS; }
159-{
160- return ge::GRAPH_SUCCESS;
161-}
162 155 
163uint64_t StatelessRandomChoiceWithMaskSimtTiling::GetTilingKey() const156uint64_t StatelessRandomChoiceWithMaskSimtTiling::GetTilingKey() const
164{157{
@@ -228,4 +221,4 @@ IMPL_OP_OPTILING(StatelessRandomChoiceWithMask)
228 .TilingInputsDataDependency({INPUT_SEED_IDX, INPUT_OFFSET_IDX})221 .TilingInputsDataDependency({INPUT_SEED_IDX, INPUT_OFFSET_IDX})
229 .Tiling(Tiling4StatelessRandomChoiceWithMask)222 .Tiling(Tiling4StatelessRandomChoiceWithMask)
230 .TilingParse<StatelessRandomChoiceWithMaskCompileInfo>(TilingPrepare4StatelessRandomChoiceWithMask);223 .TilingParse<StatelessRandomChoiceWithMaskCompileInfo>(TilingPrepare4StatelessRandomChoiceWithMask);
231-} // namespace optiling224+} // namespace optiling
@@ -49,7 +49,7 @@ static Status ParseParamsRandomNormal(const Message* op_src, ge::Operator& op_de
49 };49 };
50 }50 }
51 if (shape_list.empty()) {51 if (shape_list.empty()) {
52- OP_LOGE(GetOpName(op_dest).c_str(), "Attr of shape must be not null.");52+ OP_LOGE(GetOpName(op_dest).c_str(), "Attr of shape must not be null.");
53 return FAILED;53 return FAILED;
54 }54 }
55 55 
@@ -83,14 +83,14 @@ static Status ParseOpToGraphRandomNormal(const ge::Operator& op, ge::Graph& grap
83 op.GetAttr("shape", shape);83 op.GetAttr("shape", shape);
84 84 
85 if (shape.empty()) {85 if (shape.empty()) {
86- OP_LOGE(GetOpName(op).c_str(), "Attr of shape must be not null.");86+ OP_LOGE(GetOpName(op).c_str(), "Attr of shape must not be null.");
87 return FAILED;87 return FAILED;
88 }88 }
89 89 
90 // cast from onnx dtype to tbe dtype90 // cast from onnx dtype to tbe dtype
91 std::map<int, ge::DataType> kvlist = {{1, ge::DT_FLOAT}, {10, ge::DT_FLOAT16}, {11, ge::DT_DOUBLE}};91 std::map<int, ge::DataType> kvlist = {{1, ge::DT_FLOAT}, {10, ge::DT_FLOAT16}, {11, ge::DT_DOUBLE}};
92 if (kvlist.find(dtype) == kvlist.end()) {92 if (kvlist.find(dtype) == kvlist.end()) {
93- OP_LOGE(GetOpName(op).c_str(), "only support float32/half/double, but got %d", dtype);93+ OP_LOGE(GetOpName(op).c_str(), "only float32/half/double are supported, but got %d", dtype);
94 return FAILED;94 return FAILED;
95 }95 }
96 96 
@@ -93,7 +93,8 @@ static bool CheckFloatAndTensorShapeOfStd(const aclTensor* std, const aclTensor*
93 }93 }
94 if (!(out->IsEmpty() && std->Size() == 1)) {94 if (!(out->IsEmpty() && std->Size() == 1)) {
95 if (out->GetViewShape() != std->GetViewShape()) {95 if (out->GetViewShape() != std->GetViewShape()) {
96- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Shape of std should be match with out.");96+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Shape of std %s should match with out %s.",
97+ op::ToString(std->GetViewShape()).GetString(), op::ToString(out->GetViewShape()).GetString());
97 return false;98 return false;
98 }99 }
99 }100 }
@@ -43,13 +43,14 @@ ge::graphStatus StatelessRandomNormalV2Tiling::GetPlatformInfo()
43 } else {43 } else {
44 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);44 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
45 auto aivNum = ascendcPlatform.GetCoreNumAiv();45 auto aivNum = ascendcPlatform.GetCoreNumAiv();
46- OP_CHECK_IF((aivNum <= 0), OP_LOGE(opName, "StatelessRandomNormalV2Tiling fail to get coreNum."),46+ OP_CHECK_IF((aivNum <= 0), OP_LOGE(opName, "StatelessRandomNormalV2Tiling fails to get coreNum."),
47 return ge::GRAPH_FAILED);47 return ge::GRAPH_FAILED);
48 coreNum_ = aivNum;48 coreNum_ = aivNum;
49 uint64_t ubSizePlatForm;49 uint64_t ubSizePlatForm;
50 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);50 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSizePlatForm);
51 OP_CHECK_IF((ubSizePlatForm <= REGBASE_CCEC_CACHE_SIZE),51 OP_CHECK_IF((ubSizePlatForm <= REGBASE_CCEC_CACHE_SIZE),
52- OP_LOGE(opName, "ub size less than REGBASE_CCEC_CACHE_SIZE Size. please check"),52+ OP_LOGE(opName, "ubSize %lu is less than REGBASE_CCEC_CACHE_SIZE %u, please check", ubSizePlatForm,
53+ REGBASE_CCEC_CACHE_SIZE),
53 return ge::GRAPH_FAILED);54 return ge::GRAPH_FAILED);
54 ubSize_ = ubSizePlatForm - REGBASE_CCEC_CACHE_SIZE;55 ubSize_ = ubSizePlatForm - REGBASE_CCEC_CACHE_SIZE;
55 }56 }
@@ -95,7 +96,7 @@ ge::graphStatus StatelessRandomNormalV2Tiling::GetInputInfo()
95 }96 }
96 if (alg_ != Algorithm::RNG_ALG_PHILOX) {97 if (alg_ != Algorithm::RNG_ALG_PHILOX) {
97 std::string valueStr = std::to_string(static_cast<int32_t>(alg_));98 std::string valueStr = std::to_string(static_cast<int32_t>(alg_));
98- std::string reasonMsg = "alg only support RNG_ALG_PHILOX";99+ std::string reasonMsg = "alg only supports RNG_ALG_PHILOX";
99 OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName, "input alg", valueStr.c_str(), reasonMsg.c_str());100 OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName, "input alg", valueStr.c_str(), reasonMsg.c_str());
100 return ge::GRAPH_FAILED;101 return ge::GRAPH_FAILED;
101 }102 }
@@ -140,8 +141,8 @@ void StatelessRandomNormalV2Tiling::BlockTiling()
140 blockNum_ = CeilDiv(outputSize_, blockTilingSize_);141 blockNum_ = CeilDiv(outputSize_, blockTilingSize_);
141 tailBlockTilingSize_ = outputSize_ - blockTilingSize_ * (blockNum_ - 1);142 tailBlockTilingSize_ = outputSize_ - blockTilingSize_ * (blockNum_ - 1);
142 OP_LOGD(opName,143 OP_LOGD(opName,
143- "outputSize = %lld, blockFactor = %lld, blockAlignFactor = %lld,"144+ "outputSize = %llu, blockFactor = %llu, blockAlignFactor = %llu, "
144- "blockTilingSize = %d, tailBlockTilingSize = %d",145+ "blockTilingSize = %u, tailBlockTilingSize = %u",
145 outputSize_, blockFactor, blockAlignFactor, blockTilingSize_, tailBlockTilingSize_);146 outputSize_, blockFactor, blockAlignFactor, blockTilingSize_, tailBlockTilingSize_);
146 return;147 return;
147}148}
@@ -26,9 +26,10 @@ ge::graphStatus StatelessRandomUniformV2Tiling::GetPlatformInfo()
26{26{
27 auto compileInfoPtr = reinterpret_cast<const StatelessRandomUniformV2CompileInfo*>(context_->GetCompileInfo());27 auto compileInfoPtr = reinterpret_cast<const StatelessRandomUniformV2CompileInfo*>(context_->GetCompileInfo());
28 OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_, "compile info is null"), return ge::GRAPH_FAILED);28 OP_CHECK_IF(compileInfoPtr == nullptr, OP_LOGE(context_, "compile info is null"), return ge::GRAPH_FAILED);
29- OP_CHECK_IF((compileInfoPtr->aivNum <= 0), OP_LOGE(opName, "StatelessRandomUniformV2Tiling fail to get coreNum."),29+ OP_CHECK_IF((compileInfoPtr->aivNum <= 0), OP_LOGE(opName, "StatelessRandomUniformV2Tiling fails to get coreNum."),
30 return ge::GRAPH_FAILED);30 return ge::GRAPH_FAILED);
31- OP_CHECK_IF((compileInfoPtr->ubSize <= 0), OP_LOGE(opName, "ub size less than 0 Size. please check"),31+ OP_CHECK_IF((compileInfoPtr->ubSize <= 0),
32+ OP_LOGE(opName, "ubSize %lu is invalid, must be greater than 0.", compileInfoPtr->ubSize),
32 return ge::GRAPH_FAILED);33 return ge::GRAPH_FAILED);
33 coreNum_ = compileInfoPtr->aivNum;34 coreNum_ = compileInfoPtr->aivNum;
34 ubSize_ = compileInfoPtr->ubSize;35 ubSize_ = compileInfoPtr->ubSize;
@@ -80,7 +81,7 @@ ge::graphStatus StatelessRandomUniformV2Tiling::GetInputInfo()
80 }81 }
81 if (alg_ != Algorithm::RNG_ALG_PHILOX) {82 if (alg_ != Algorithm::RNG_ALG_PHILOX) {
82 std::string valueStr = std::to_string(static_cast<int32_t>(alg_));83 std::string valueStr = std::to_string(static_cast<int32_t>(alg_));
83- std::string reasonMsg = "alg only support RNG_ALG_PHILOX";84+ std::string reasonMsg = "alg only supports RNG_ALG_PHILOX";
84 OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName, "input alg", valueStr.c_str(), reasonMsg.c_str());85 OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(opName, "input alg", valueStr.c_str(), reasonMsg.c_str());
85 return ge::GRAPH_FAILED;86 return ge::GRAPH_FAILED;
86 }87 }
@@ -129,7 +130,7 @@ void StatelessRandomUniformV2Tiling::BlockTiling()
129 blockNum_ = CeilDiv(outputSize_, blockTilingSize_);130 blockNum_ = CeilDiv(outputSize_, blockTilingSize_);
130 tailBlockTilingSize_ = outputSize_ - blockTilingSize_ * (blockNum_ - 1);131 tailBlockTilingSize_ = outputSize_ - blockTilingSize_ * (blockNum_ - 1);
131 OP_LOGD(opName,132 OP_LOGD(opName,
132- "outputSize = %lld, blockFactor = %lld, blockAlignFactor = %lld,"133+ "outputSize = %u, blockFactor = %u, blockAlignFactor = %u, "
133 "blockTilingSize = %d, tailBlockTilingSize = %d",134 "blockTilingSize = %d, tailBlockTilingSize = %d",
134 outputSize_, blockFactor, blockAlignFactor, blockTilingSize_, tailBlockTilingSize_);135 outputSize_, blockFactor, blockAlignFactor, blockTilingSize_, tailBlockTilingSize_);
135 return;136 return;
@@ -63,7 +63,8 @@ ge::graphStatus StatelessRandpermTiling::GetPlatformInfo()
63 63 
64 totalCoreNum_ = static_cast<int64_t>(compileInfo->totalCoreNum);64 totalCoreNum_ = static_cast<int64_t>(compileInfo->totalCoreNum);
65 ubSize_ = compileInfo->ubSize;65 ubSize_ = compileInfo->ubSize;
66- OP_CHECK_IF(ubSize_ <= 0, OP_LOGE(opName_, "UB size is invalid."), return ge::GRAPH_FAILED);66+ OP_CHECK_IF(ubSize_ <= 0, OP_LOGE(opName_, "UB size %ld is invalid, must be greater than 0.", ubSize_),
67+ return ge::GRAPH_FAILED);
67 OP_CHECK_IF(ubSize_ <= SIMT_DCACHE_SIZE,68 OP_CHECK_IF(ubSize_ <= SIMT_DCACHE_SIZE,
68 OP_LOGE(opName_, "UB size %ld bytes must be greater than simt dcache size %ld bytes, please check.",69 OP_LOGE(opName_, "UB size %ld bytes must be greater than simt dcache size %ld bytes, please check.",
69 ubSize_, SIMT_DCACHE_SIZE),70 ubSize_, SIMT_DCACHE_SIZE),
@@ -93,7 +94,8 @@ ge::graphStatus StatelessRandpermTiling::GetAttrs()
93 }94 }
94 if (OUTPUT_DTYPE.find(attrOutDtype_) == OUTPUT_DTYPE.end()) {95 if (OUTPUT_DTYPE.find(attrOutDtype_) == OUTPUT_DTYPE.end()) {
95 std::string valueStr = ToString(attrOutDtype_);96 std::string valueStr = ToString(attrOutDtype_);
96- std::string reasonMsg = "[attr]dtype only support int64, int32, int16, int8, float32, float16, bfloat16";97+ std::string
98+ reasonMsg = "[attr]dtype only supports int64, int32, int16, uint8, int8, float32, float16, bfloat16";
97 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "attr dtype", valueStr.c_str(), reasonMsg.c_str());99 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "attr dtype", valueStr.c_str(), reasonMsg.c_str());
98 return ge::GRAPH_FAILED;100 return ge::GRAPH_FAILED;
99 }101 }
@@ -236,7 +238,8 @@ ge::graphStatus StatelessRandpermTiling::GetOutputY()
236 auto outDtype = outDesc->GetDataType();238 auto outDtype = outDesc->GetDataType();
237 if (OUTPUT_DTYPE.count(outDtype) == 0) {239 if (OUTPUT_DTYPE.count(outDtype) == 0) {
238 std::string valueStr = ToString(outDtype);240 std::string valueStr = ToString(outDtype);
239- std::string reasonMsg = "output y dtype should be in int64, int32, int16, int8, float32, float16, bfloat16";241+ std::string
242+ reasonMsg = "output y dtype should be in int64, int32, int16, uint8, int8, float32, float16, bfloat16";
240 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "output tensor y", valueStr.c_str(), reasonMsg.c_str());243 OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(opName_, "output tensor y", valueStr.c_str(), reasonMsg.c_str());
241 return ge::GRAPH_FAILED;244 return ge::GRAPH_FAILED;
242 }245 }
@@ -25,7 +25,7 @@
25#include "stateless_randperm_tiling_for_sort.h"25#include "stateless_randperm_tiling_for_sort.h"
26 26 
27namespace optiling {27namespace optiling {
28-namespace statelessRandpermTiling{28+namespace statelessRandpermTiling {
29constexpr size_t WORK_SPACE_SIZE = 16 * 1024 * 1024;29constexpr size_t WORK_SPACE_SIZE = 16 * 1024 * 1024;
30const uint32_t BIN_NUM = 256; // 直方图一次处理256B30const uint32_t BIN_NUM = 256; // 直方图一次处理256B
31const uint32_t SMALL_TILE_DATA_NUM = 1024; // 测试数据得出一次至少处理1024,sort性能比较好31const uint32_t SMALL_TILE_DATA_NUM = 1024; // 测试数据得出一次至少处理1024,sort性能比较好
@@ -72,13 +72,10 @@ struct SortTileInfo {
72 int64_t unSortDimNum = 1;72 int64_t unSortDimNum = 1;
73};73};
74static const std::map<ge::DataType, uint32_t> tilingDataTypeBitMap = {74static const std::map<ge::DataType, uint32_t> tilingDataTypeBitMap = {
75- { ge::DT_INT64, 8 }, { ge::DT_INT32, 4 }, { ge::DT_INT16, 2 }, { ge::DT_INT8, 1 },75+ {ge::DT_INT64, 8}, {ge::DT_INT32, 4}, {ge::DT_INT16, 2}, {ge::DT_INT8, 1},
76- { ge::DT_UINT64, 8 }, { ge::DT_UINT32, 4 }, { ge::DT_UINT16, 2 }, { ge::DT_UINT8, 1 },76+ {ge::DT_UINT64, 8}, {ge::DT_UINT32, 4}, {ge::DT_UINT16, 2}, {ge::DT_UINT8, 1},
77- { ge::DT_FLOAT, 4 }, { ge::DT_FLOAT16, 2 }, { ge::DT_BF16, 2 }77+ {ge::DT_FLOAT, 4}, {ge::DT_FLOAT16, 2}, {ge::DT_BF16, 2}};
78-};78+static const std::map<ge::DataType, uint32_t> mergeType = {{ge::DT_FLOAT, 4}, {ge::DT_FLOAT16, 2}, {ge::DT_BF16, 2}};
79-static const std::map<ge::DataType, uint32_t> mergeType = { { ge::DT_FLOAT, 4 },
80- { ge::DT_FLOAT16, 2 },
81- { ge::DT_BF16, 2 } };
82 79 
83uint32_t CeilDiv(int64_t a, int64_t b)80uint32_t CeilDiv(int64_t a, int64_t b)
84{81{
@@ -89,7 +86,7 @@ uint32_t CeilDiv(int64_t a, int64_t b)
89}86}
90 87 
91template <typename T>88template <typename T>
92-auto CeilDivMul(int64_t a, int64_t b) ->T const89+auto CeilDivMul(int64_t a, int64_t b) -> T const
93{90{
94 if (b == 0) {91 if (b == 0) {
95 return static_cast<T>(a);92 return static_cast<T>(a);
@@ -97,10 +94,10 @@ auto CeilDivMul(int64_t a, int64_t b) ->T const
97 return static_cast<T>(((a + b - 1) / b) * b);94 return static_cast<T>(((a + b - 1) / b) * b);
98}95}
99 96 
100-ge::graphStatus CheckInputAndOutput(gert::TilingContext *context, SortTileInfo &sortTileInfo)97+ge::graphStatus CheckInputAndOutput(gert::TilingContext* context, SortTileInfo& sortTileInfo)
101{98{
102 auto platformInfo = context->GetPlatformInfo();99 auto platformInfo = context->GetPlatformInfo();
103- OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo); 100+ OP_CHECK_NULL_WITH_CONTEXT(context, platformInfo);
104 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);101 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfo);
105 uint64_t ubSize = 0;102 uint64_t ubSize = 0;
106 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);103 ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
@@ -111,29 +108,31 @@ ge::graphStatus CheckInputAndOutput(gert::TilingContext *context, SortTileInfo &
111 return ge::GRAPH_FAILED;108 return ge::GRAPH_FAILED;
112 }109 }
113 sortTileInfo.blockUbSize = Ops::Base::GetUbBlockSize(context);110 sortTileInfo.blockUbSize = Ops::Base::GetUbBlockSize(context);
114- sortTileInfo.ubSize = ubSize - SIMT_UB; // 侵入修改:MODIFY111+ sortTileInfo.ubSize = ubSize - SIMT_UB; // 侵入修改:MODIFY
115- OP_LOGI(context->GetNodeName(), "ubSize is %ld, simtDcache is %u, blockUbSize %u", sortTileInfo.ubSize, SIMT_UB, sortTileInfo.blockUbSize); // 侵入修改:MODIFY112+ OP_LOGI(context->GetNodeName(), "ubSize is %u, simtDcache is %u, blockUbSize %u", sortTileInfo.ubSize, SIMT_UB,
113+ sortTileInfo.blockUbSize); // 侵入修改:MODIFY
116 auto inputShapePtr = context->GetInputShape(0);114 auto inputShapePtr = context->GetInputShape(0);
117 OP_CHECK_NULL_WITH_CONTEXT(context, inputShapePtr);115 OP_CHECK_NULL_WITH_CONTEXT(context, inputShapePtr);
118- const gert::Shape &inputShape = Ops::Base::EnsureNotScalar(inputShapePtr->GetStorageShape());116+ const gert::Shape& inputShape = Ops::Base::EnsureNotScalar(inputShapePtr->GetStorageShape());
119 auto yStorage = context->GetOutputShape(0);117 auto yStorage = context->GetOutputShape(0);
120 OP_CHECK_NULL_WITH_CONTEXT(context, yStorage);118 OP_CHECK_NULL_WITH_CONTEXT(context, yStorage);
121- const gert::Shape &outShape = Ops::Base::EnsureNotScalar(yStorage->GetStorageShape());119+ const gert::Shape& outShape = Ops::Base::EnsureNotScalar(yStorage->GetStorageShape());
122 auto yStorage1 = context->GetOutputShape(1);120 auto yStorage1 = context->GetOutputShape(1);
123 OP_CHECK_NULL_WITH_CONTEXT(context, yStorage1);121 OP_CHECK_NULL_WITH_CONTEXT(context, yStorage1);
124- const gert::Shape &outShape1 = Ops::Base::EnsureNotScalar(yStorage1->GetStorageShape());122+ const gert::Shape& outShape1 = Ops::Base::EnsureNotScalar(yStorage1->GetStorageShape());
125 if (inputShape.GetShapeSize() == 0 || outShape.GetShapeSize() == 0) {123 if (inputShape.GetShapeSize() == 0 || outShape.GetShapeSize() == 0) {
126- std::string valueStr = std::to_string(inputShape.GetShapeSize()) + " and " + std::to_string(outShape.GetShapeSize());124+ std::string valueStr = std::to_string(inputShape.GetShapeSize()) + " and " +
127- std::string reasonMsg = "not support empty input or output";125+ std::to_string(outShape.GetShapeSize());
128- OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(126+ std::string reasonMsg = "empty input or output is not supported";
129- context->GetNodeName(), "input and output", valueStr.c_str(), reasonMsg.c_str());127+ OP_LOGE_FOR_INVALID_SHAPESIZES_WITH_REASON(context->GetNodeName(), "input and output", valueStr.c_str(),
128+ reasonMsg.c_str());
130 return ge::GRAPH_FAILED;129 return ge::GRAPH_FAILED;
131 }130 }
132 if (outShape != outShape1 || outShape != inputShape) {131 if (outShape != outShape1 || outShape != inputShape) {
133 std::string valueStr = Ops::Base::ToString(inputShape) + " and " + Ops::Base::ToString(outShape);132 std::string valueStr = Ops::Base::ToString(inputShape) + " and " + Ops::Base::ToString(outShape);
134 std::string reasonMsg = "input and outputs shape must be the same";133 std::string reasonMsg = "input and outputs shape must be the same";
135- OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(134+ OP_LOGE_FOR_INVALID_SHAPES_WITH_REASON(context->GetNodeName(), "input and output", valueStr.c_str(),
136- context->GetNodeName(), "input and output", valueStr.c_str(), reasonMsg.c_str());135+ reasonMsg.c_str());
137 return ge::GRAPH_FAILED;136 return ge::GRAPH_FAILED;
138 }137 }
139 int32_t xDimNum = inputShape.GetDimNum();138 int32_t xDimNum = inputShape.GetDimNum();
@@ -149,17 +148,18 @@ ge::graphStatus CheckInputAndOutput(gert::TilingContext *context, SortTileInfo &
149 return ge::GRAPH_SUCCESS;148 return ge::GRAPH_SUCCESS;
150}149}
151 150 
152-ge::graphStatus SortCheckParams(gert::TilingContext *context, SortTileInfo &sortTileInfo)151+ge::graphStatus SortCheckParams(gert::TilingContext* context, SortTileInfo& sortTileInfo)
153{152{
154 OP_CHECK_IF(CheckInputAndOutput(context, sortTileInfo) != ge::GRAPH_SUCCESS,153 OP_CHECK_IF(CheckInputAndOutput(context, sortTileInfo) != ge::GRAPH_SUCCESS,
155- OP_LOGE(context->GetNodeName(), "CheckInputAndOutput failed"), return ge::GRAPH_FAILED);154+ OP_LOGE(context->GetNodeName(), "CheckInputAndOutput failed"), return ge::GRAPH_FAILED);
156 auto inputDescPtr = context->GetInputDesc(0);155 auto inputDescPtr = context->GetInputDesc(0);
157 OP_CHECK_NULL_WITH_CONTEXT(context, inputDescPtr);156 OP_CHECK_NULL_WITH_CONTEXT(context, inputDescPtr);
158 ge::DataType dataType = inputDescPtr->GetDataType();157 ge::DataType dataType = inputDescPtr->GetDataType();
159 if (tilingDataTypeBitMap.count(dataType) == 0) {158 if (tilingDataTypeBitMap.count(dataType) == 0) {
160 std::string valueStr = Ops::Base::ToString(dataType);159 std::string valueStr = Ops::Base::ToString(dataType);
161 std::string reasonMsg = "Not supported data type";160 std::string reasonMsg = "Not supported data type";
162- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "input tensor x", valueStr.c_str(), reasonMsg.c_str());161+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "input tensor x", valueStr.c_str(),
162+ reasonMsg.c_str());
163 return ge::GRAPH_FAILED;163 return ge::GRAPH_FAILED;
164 }164 }
165 sortTileInfo.dataType = dataType;165 sortTileInfo.dataType = dataType;
@@ -172,39 +172,41 @@ ge::graphStatus SortCheckParams(gert::TilingContext *context, SortTileInfo &sort
172 auto y1DType = outDescPtr0->GetDataType();172 auto y1DType = outDescPtr0->GetDataType();
173 if ((y2DType != ge::DT_INT64) && (y2DType != ge::DT_INT32)) {173 if ((y2DType != ge::DT_INT64) && (y2DType != ge::DT_INT32)) {
174 std::string valueStr = Ops::Base::ToString(y2DType);174 std::string valueStr = Ops::Base::ToString(y2DType);
175- std::string reasonMsg = "y2 dtype only support int64 or int32";175+ std::string reasonMsg = "y2 dtype only supports int64 or int32";
176- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "output tensor y2", valueStr.c_str(), reasonMsg.c_str());176+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(context->GetNodeName(), "output tensor y2", valueStr.c_str(),
177+ reasonMsg.c_str());
177 return ge::GRAPH_FAILED;178 return ge::GRAPH_FAILED;
178 }179 }
179 if (y1DType != dataType) {180 if (y1DType != dataType) {
180 std::string valueStr = Ops::Base::ToString(dataType) + " and " + Ops::Base::ToString(y1DType);181 std::string valueStr = Ops::Base::ToString(dataType) + " and " + Ops::Base::ToString(y1DType);
181- std::string reasonMsg = "input0 dtype must be same as output0 dtype";182+ std::string reasonMsg = "input0 dtype must be the same as output0 dtype";
182- OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(183+ OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(context->GetNodeName(), "input and output", valueStr.c_str(),
183- context->GetNodeName(), "input and output", valueStr.c_str(), reasonMsg.c_str());184+ reasonMsg.c_str());
184 return ge::GRAPH_FAILED;185 return ge::GRAPH_FAILED;
185 }186 }
186 sortTileInfo.y2DtypeSize = tilingDataTypeBitMap.find(y2DType)->second;187 sortTileInfo.y2DtypeSize = tilingDataTypeBitMap.find(y2DType)->second;
187 auto const attrs = context->GetAttrs();188 auto const attrs = context->GetAttrs();
188 OP_CHECK_NULL_WITH_CONTEXT(context, attrs);189 OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
189- const bool *isDescending = attrs->GetAttrPointer<bool>(1);190+ const bool* isDescending = attrs->GetAttrPointer<bool>(1);
190- const int64_t *sortAxisPtr = attrs->GetAttrPointer<int64_t>(0);191+ const int64_t* sortAxisPtr = attrs->GetAttrPointer<int64_t>(0);
191 OP_CHECK_NULL_WITH_CONTEXT(context, isDescending);192 OP_CHECK_NULL_WITH_CONTEXT(context, isDescending);
192 OP_CHECK_NULL_WITH_CONTEXT(context, sortAxisPtr);193 OP_CHECK_NULL_WITH_CONTEXT(context, sortAxisPtr);
193 int32_t sortAxis = static_cast<int32_t>(*sortAxisPtr);194 int32_t sortAxis = static_cast<int32_t>(*sortAxisPtr);
194 sortAxis = sortAxis < 0 ? (sortAxis + sortTileInfo.xDimNum) : sortAxis;195 sortAxis = sortAxis < 0 ? (sortAxis + sortTileInfo.xDimNum) : sortAxis;
195 if (sortAxis != (sortTileInfo.xDimNum - 1)) {196 if (sortAxis != (sortTileInfo.xDimNum - 1)) {
196 std::string valueStr = std::to_string(sortAxis);197 std::string valueStr = std::to_string(sortAxis);
197- std::string reasonMsg = "only support last dim sort";198+ std::string reasonMsg = "only last dim sort is supported";
198- OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "attr sort_axis", valueStr.c_str(), reasonMsg.c_str());199+ OP_LOGE_FOR_INVALID_VALUE_WITH_REASON(context->GetNodeName(), "attr sort_axis", valueStr.c_str(),
200+ reasonMsg.c_str());
199 return ge::GRAPH_FAILED;201 return ge::GRAPH_FAILED;
200 }202 }
201 return ge::GRAPH_SUCCESS;203 return ge::GRAPH_SUCCESS;
202}204}
203 205 
204-void SetSortTmpSize(ge::DataType dataType, uint32_t tileData, bool isDescend, SortTileInfo &sortTileInfo)206+void SetSortTmpSize(ge::DataType dataType, uint32_t tileData, bool isDescend, SortTileInfo& sortTileInfo)
205{207{
206 int64_t realLen = std::min(sortTileInfo.sortAxisNum, static_cast<int64_t>(tileData));208 int64_t realLen = std::min(sortTileInfo.sortAxisNum, static_cast<int64_t>(tileData));
207- std::vector<int64_t> shapeVec = { realLen };209+ std::vector<int64_t> shapeVec = {realLen};
208 ge::Shape srcShape(shapeVec);210 ge::Shape srcShape(shapeVec);
209 AscendC::SortConfig config;211 AscendC::SortConfig config;
210 config.type = AscendC::SortType::RADIX_SORT;212 config.type = AscendC::SortType::RADIX_SORT;
@@ -218,22 +220,23 @@ void SetSortTmpSize(ge::DataType dataType, uint32_t tileData, bool isDescend, So
218 return;220 return;
219}221}
220 222 
221-bool IsMergeSort(SortTileInfo &sortTileInfo)223+bool IsMergeSort(SortTileInfo& sortTileInfo)
222{224{
223- bool support =225+ bool support = (sortTileInfo.sortAxisNum <= SMALL_SORT_MAX_DATA_SIZE) &&
224- (sortTileInfo.sortAxisNum <= SMALL_SORT_MAX_DATA_SIZE) && (mergeType.count(sortTileInfo.dataType) != 0);226+ (mergeType.count(sortTileInfo.dataType) != 0);
225 return support;227 return support;
226}228}
227 229 
228-bool IsMergeSortMultiCore(SortTileInfo &sortTileInfo)230+bool IsMergeSortMultiCore(SortTileInfo& sortTileInfo)
229{231{
230- bool isMuiltiCoreMergeSort =232+ bool isMuiltiCoreMergeSort = ((sortTileInfo.unSortDimNum == 1) &&
231- ((sortTileInfo.unSortDimNum == 1) && (sortTileInfo.sortAxisNum <= MULTI_CORE_MERGE_SORT_MAX_SIZE) &&233+ (sortTileInfo.sortAxisNum <= MULTI_CORE_MERGE_SORT_MAX_SIZE) &&
232- (sortTileInfo.sortAxisNum > SMALL_SORT_MAX_DATA_SIZE) && (sortTileInfo.dataType == ge::DT_FLOAT));234+ (sortTileInfo.sortAxisNum > SMALL_SORT_MAX_DATA_SIZE) &&
235+ (sortTileInfo.dataType == ge::DT_FLOAT));
233 return isMuiltiCoreMergeSort;236 return isMuiltiCoreMergeSort;
234}237}
235 238 
236-bool IsRadixSortOneCore(SortTileInfo &sortTileInfo)239+bool IsRadixSortOneCore(SortTileInfo& sortTileInfo)
237{240{
238 if (sortTileInfo.isInt32 == static_cast<uint32_t>(0)) {241 if (sortTileInfo.isInt32 == static_cast<uint32_t>(0)) {
239 return false;242 return false;
@@ -260,13 +263,13 @@ bool IsRadixSortOneCore(SortTileInfo &sortTileInfo)
260 return tmpUb <= remainUb;263 return tmpUb <= remainUb;
261}264}
262 265 
263-uint32_t ComputeRemainUb(SortTileInfo &sortTileInfo, uint32_t tileData, uint32_t ubExtra, uint32_t tileFactor)266+uint32_t ComputeRemainUb(SortTileInfo& sortTileInfo, uint32_t tileData, uint32_t ubExtra, uint32_t tileFactor)
264{267{
265 uint32_t tmpUb = sortTileInfo.ubSize - (ubExtra + tileFactor * tileData);268 uint32_t tmpUb = sortTileInfo.ubSize - (ubExtra + tileFactor * tileData);
266 return tmpUb;269 return tmpUb;
267}270}
268 271 
269-void AdjTmpUb(SortTileInfo &sortTileInfo, uint32_t tileData, uint32_t ubExtra, uint32_t tileFactor)272+void AdjTmpUb(SortTileInfo& sortTileInfo, uint32_t tileData, uint32_t ubExtra, uint32_t tileFactor)
270{273{
271 uint32_t remainUbNew = ComputeRemainUb(sortTileInfo, tileData, ubExtra, tileFactor) - sortTileInfo.tmpUbSize;274 uint32_t remainUbNew = ComputeRemainUb(sortTileInfo, tileData, ubExtra, tileFactor) - sortTileInfo.tmpUbSize;
272 remainUbNew = remainUbNew > sortTileInfo.blockUbSize ? (remainUbNew - sortTileInfo.blockUbSize) : uint32_t(0);275 remainUbNew = remainUbNew > sortTileInfo.blockUbSize ? (remainUbNew - sortTileInfo.blockUbSize) : uint32_t(0);
@@ -275,7 +278,7 @@ void AdjTmpUb(SortTileInfo &sortTileInfo, uint32_t tileData, uint32_t ubExtra, u
275 sortTileInfo.tmpUbSize = sortTileInfo.tmpUbSize + alignUbSize; // 剩余的ub都给tmpUbsize278 sortTileInfo.tmpUbSize = sortTileInfo.tmpUbSize + alignUbSize; // 剩余的ub都给tmpUbsize
276}279}
277 280 
278-void ComputeTileDataOne(SortTileInfo &sortTileInfo, uint32_t lastDimTileNum, uint32_t ubExtra, uint32_t &tileData,281+void ComputeTileDataOne(SortTileInfo& sortTileInfo, uint32_t lastDimTileNum, uint32_t ubExtra, uint32_t& tileData,
279 uint32_t tileFactor)282 uint32_t tileFactor)
280{283{
281 uint32_t allCore = CeilDivMul<uint32_t>(int64_t(lastDimTileNum), int64_t(sortTileInfo.maxCoreNum));284 uint32_t allCore = CeilDivMul<uint32_t>(int64_t(lastDimTileNum), int64_t(sortTileInfo.maxCoreNum));
@@ -287,11 +290,11 @@ void ComputeTileDataOne(SortTileInfo &sortTileInfo, uint32_t lastDimTileNum, ui
287 return;290 return;
288}291}
289 292 
290-bool NeedAdjTileData(SortTileInfo &sortTileInfo, uint32_t &tileData, uint32_t lastDimTileNum, uint32_t ubExtra,293+bool NeedAdjTileData(SortTileInfo& sortTileInfo, uint32_t& tileData, uint32_t lastDimTileNum, uint32_t ubExtra,
291 uint32_t tileFactor)294 uint32_t tileFactor)
292{295{
293 if (sortTileInfo.unSortDimNum == int64_t(1) && lastDimTileNum == uint32_t(1)) {296 if (sortTileInfo.unSortDimNum == int64_t(1) && lastDimTileNum == uint32_t(1)) {
294- OP_LOGI("RadixSortTiling", "unSortDimNum and lastDimTileNum is 1");297+ OP_LOGI("RadixSortTiling", "unSortDimNum and lastDimTileNum are both 1");
295 uint32_t newTileData = CeilDiv(sortTileInfo.sortAxisNum, int64_t(sortTileInfo.maxCoreNum));298 uint32_t newTileData = CeilDiv(sortTileInfo.sortAxisNum, int64_t(sortTileInfo.maxCoreNum));
296 newTileData = CeilDivMul<uint32_t>(int64_t(newTileData), int64_t(BIN_NUM));299 newTileData = CeilDivMul<uint32_t>(int64_t(newTileData), int64_t(BIN_NUM));
297 tileData = std::max(newTileData, SMALL_TILE_DATA_NUM);300 tileData = std::max(newTileData, SMALL_TILE_DATA_NUM);
@@ -301,13 +304,13 @@ bool NeedAdjTileData(SortTileInfo &sortTileInfo, uint32_t &tileData, uint32_t la
301 }304 }
302 if (sortTileInfo.unSortDimNum == int64_t(1) || (lastDimTileNum >= sortTileInfo.maxCoreNum)) {305 if (sortTileInfo.unSortDimNum == int64_t(1) || (lastDimTileNum >= sortTileInfo.maxCoreNum)) {
303 // b为1时,尽量均匀分核,同时保证处理的最小的tile_data为1024306 // b为1时,尽量均匀分核,同时保证处理的最小的tile_data为1024
304- OP_LOGI("RadixSortTiling", "unSortDimNum is 1 and lastDimTileNum greater than allCore");307+ OP_LOGI("RadixSortTiling", "unSortDimNum is 1 and lastDimTileNum is greater than allCore");
305 ComputeTileDataOne(sortTileInfo, lastDimTileNum, ubExtra, tileData, tileFactor);308 ComputeTileDataOne(sortTileInfo, lastDimTileNum, ubExtra, tileData, tileFactor);
306 return true;309 return true;
307 }310 }
308 if (sortTileInfo.unSortDimNum > int64_t(1) && sortTileInfo.unSortDimNum < int64_t(sortTileInfo.maxCoreNum) &&311 if (sortTileInfo.unSortDimNum > int64_t(1) && sortTileInfo.unSortDimNum < int64_t(sortTileInfo.maxCoreNum) &&
309 lastDimTileNum == uint32_t(1)) {312 lastDimTileNum == uint32_t(1)) {
310- OP_LOGI("RadixSortTiling", "unSortDimNum greater than 1,and unSortDimNum small and lastDimTileNum is one");313+ OP_LOGI("RadixSortTiling", "unSortDimNum is greater than 1, unSortDimNum is small and lastDimTileNum is one");
311 uint32_t hCore = sortTileInfo.maxCoreNum / static_cast<uint32_t>(sortTileInfo.unSortDimNum);314 uint32_t hCore = sortTileInfo.maxCoreNum / static_cast<uint32_t>(sortTileInfo.unSortDimNum);
312 uint32_t hTileData = static_cast<uint32_t>(sortTileInfo.sortAxisNum) / hCore;315 uint32_t hTileData = static_cast<uint32_t>(sortTileInfo.sortAxisNum) / hCore;
313 tileData = CeilDivMul<uint32_t>(int64_t(hTileData), int64_t(BIN_NUM));316 tileData = CeilDivMul<uint32_t>(int64_t(hTileData), int64_t(BIN_NUM));
@@ -317,7 +320,7 @@ bool NeedAdjTileData(SortTileInfo &sortTileInfo, uint32_t &tileData, uint32_t la
317 }320 }
318 if (sortTileInfo.unSortDimNum > int64_t(1) && lastDimTileNum > uint32_t(1)) {321 if (sortTileInfo.unSortDimNum > int64_t(1) && lastDimTileNum > uint32_t(1)) {
319 // b大于1且h轴循环次数小于总核数,也就是b轴核数大于1322 // b大于1且h轴循环次数小于总核数,也就是b轴核数大于1
320- OP_LOGI("RadixSortTiling", "unSortDimNum is one, lastDimTileNum greater than one");323+ OP_LOGI("RadixSortTiling", "unSortDimNum is one, lastDimTileNum is greater than one");
321 int64_t newTileData = sortTileInfo.sortAxisNum / int64_t(lastDimTileNum);324 int64_t newTileData = sortTileInfo.sortAxisNum / int64_t(lastDimTileNum);
322 tileData = CeilDivMul<uint32_t>(newTileData, int64_t(BIN_NUM));325 tileData = CeilDivMul<uint32_t>(newTileData, int64_t(BIN_NUM));
323 lastDimTileNum = CeilDiv(sortTileInfo.sortAxisNum, int64_t(tileData));326 lastDimTileNum = CeilDiv(sortTileInfo.sortAxisNum, int64_t(tileData));
@@ -340,7 +343,7 @@ bool NeedAdjTileData(SortTileInfo &sortTileInfo, uint32_t &tileData, uint32_t la
340 return false;343 return false;
341}344}
342 345 
343-uint32_t ComputeTileData(SortTileInfo &sortTileInfo)346+uint32_t ComputeTileData(SortTileInfo& sortTileInfo)
344{347{
345 uint32_t ubExtra;348 uint32_t ubExtra;
346 uint32_t tileFactor;349 uint32_t tileFactor;
@@ -367,8 +370,8 @@ uint32_t ComputeTileData(SortTileInfo &sortTileInfo)
367 }370 }
368 uint32_t lastDimTileNum = CeilDiv(sortTileInfo.sortAxisNum, int64_t(tileData));371 uint32_t lastDimTileNum = CeilDiv(sortTileInfo.sortAxisNum, int64_t(tileData));
369 OP_LOGI("RadixSortTiling", "tileData %u, lastDimTileNum %u, tmpUbSize %u", tileData, lastDimTileNum, tmpUbSize);372 OP_LOGI("RadixSortTiling", "tileData %u, lastDimTileNum %u, tmpUbSize %u", tileData, lastDimTileNum, tmpUbSize);
370- bool smallTile =373+ bool smallTile = (sortTileInfo.sortAxisNum <= static_cast<int64_t>(SMALL_TILE_DATA_NUM)) &&
371- (sortTileInfo.sortAxisNum <= static_cast<int64_t>(SMALL_TILE_DATA_NUM)) && lastDimTileNum == uint32_t(1);374+ lastDimTileNum == uint32_t(1);
372 if ((lastDimTileNum % sortTileInfo.maxCoreNum == static_cast<uint32_t>(0)) || smallTile) {375 if ((lastDimTileNum % sortTileInfo.maxCoreNum == static_cast<uint32_t>(0)) || smallTile) {
373 OP_LOGI("RadixSortTiling", "lastDimTileNum align or smallTile");376 OP_LOGI("RadixSortTiling", "lastDimTileNum align or smallTile");
374 AdjTmpUb(sortTileInfo, tileData, ubExtra, tileFactor);377 AdjTmpUb(sortTileInfo, tileData, ubExtra, tileFactor);
@@ -381,37 +384,40 @@ uint32_t ComputeTileData(SortTileInfo &sortTileInfo)
381 return tileData;384 return tileData;
382}385}
383 386 
384-void GetMergeSortMultiCore(gert::TilingContext *context, SortTileInfo &sortTileInfo) {387+void GetMergeSortMultiCore(gert::TilingContext* context, SortTileInfo& sortTileInfo)
385- uint32_t coreNumNeed = static_cast<uint32_t>((sortTileInfo.sortAxisNum + ONE_CORE_DATA_SIZE - 1) / ONE_CORE_DATA_SIZE);388+{
389+ uint32_t coreNumNeed = static_cast<uint32_t>((sortTileInfo.sortAxisNum + ONE_CORE_DATA_SIZE - 1) /
390+ ONE_CORE_DATA_SIZE);
386 uint32_t tileNum = static_cast<uint32_t>(sortTileInfo.sortAxisNum) / coreNumNeed;391 uint32_t tileNum = static_cast<uint32_t>(sortTileInfo.sortAxisNum) / coreNumNeed;
387 sortTileInfo.lastDimTileNum = static_cast<uint32_t>(sortTileInfo.sortAxisNum);392 sortTileInfo.lastDimTileNum = static_cast<uint32_t>(sortTileInfo.sortAxisNum);
388 sortTileInfo.lastDimNeedCore = coreNumNeed;393 sortTileInfo.lastDimNeedCore = coreNumNeed;
389 sortTileInfo.numTileDataSize = tileNum;394 sortTileInfo.numTileDataSize = tileNum;
390- sortTileInfo.coreNumNeed = coreNumNeed;395+ sortTileInfo.coreNumNeed = coreNumNeed;
391 396 
392- uint32_t byteNum = MERGE_SORT_DEALING_LIST_NUM * MERGE_SORT_DATASIZE * 2;//4list 8byte 2input/output397+ uint32_t byteNum = MERGE_SORT_DEALING_LIST_NUM * MERGE_SORT_DATASIZE * 2; // 4list 8byte 2input/output
393- byteNum += MERGE_SORT_DEALING_LIST_NUM * static_cast<uint32_t>(sizeof(uint32_t));//extract index398+ byteNum += MERGE_SORT_DEALING_LIST_NUM * static_cast<uint32_t>(sizeof(uint32_t)); // extract index
394 if (sortTileInfo.y2DtypeSize == sizeof(int64_t)) {399 if (sortTileInfo.y2DtypeSize == sizeof(int64_t)) {
395 byteNum += MERGE_SORT_DEALING_LIST_NUM * static_cast<uint32_t>(sizeof(int64_t));400 byteNum += MERGE_SORT_DEALING_LIST_NUM * static_cast<uint32_t>(sizeof(int64_t));
396 }401 }
397 if (sortTileInfo.dataType == ge::DT_BF16) {402 if (sortTileInfo.dataType == ge::DT_BF16) {
398 byteNum += MERGE_SORT_DEALING_LIST_NUM * mergeType.find(ge::DT_BF16)->second;403 byteNum += MERGE_SORT_DEALING_LIST_NUM * mergeType.find(ge::DT_BF16)->second;
399- byteNum += MERGE_SORT_DEALING_LIST_NUM * static_cast<uint32_t>(sizeof(float));//extract value404+ byteNum += MERGE_SORT_DEALING_LIST_NUM * static_cast<uint32_t>(sizeof(float)); // extract value
400 } else {405 } else {
401- byteNum += MERGE_SORT_DEALING_LIST_NUM * tilingDataTypeBitMap.find(sortTileInfo.dataType)->second;//extract value406+ byteNum += MERGE_SORT_DEALING_LIST_NUM *
402- }407+ tilingDataTypeBitMap.find(sortTileInfo.dataType)->second; // extract value
408+ }
403 sortTileInfo.keyParams0 = sortTileInfo.ubSize / byteNum;409 sortTileInfo.keyParams0 = sortTileInfo.ubSize / byteNum;
404 410 
405 OP_LOGI("[mergeSort]", "maxDealingNum: %u", sortTileInfo.keyParams0);411 OP_LOGI("[mergeSort]", "maxDealingNum: %u", sortTileInfo.keyParams0);
406 size_t* userWorkSpaceSize = context->GetWorkspaceSizes(1);412 size_t* userWorkSpaceSize = context->GetWorkspaceSizes(1);
407- size_t usrSize = static_cast<size_t>(413+ size_t usrSize = static_cast<size_t>(MERGE_SORT_WORKSPACE_PARAM * sortTileInfo.sortAxisNum *
408- MERGE_SORT_WORKSPACE_PARAM * sortTileInfo.sortAxisNum * static_cast<uint32_t>(sizeof(int32_t)));414+ static_cast<uint32_t>(sizeof(int32_t)));
409 userWorkSpaceSize[0] = usrSize + WORK_SPACE_SIZE;415 userWorkSpaceSize[0] = usrSize + WORK_SPACE_SIZE;
410 context->SetScheduleMode(1);416 context->SetScheduleMode(1);
411 return;417 return;
412}418}
413 419 
414-void GetRadixSortOneCore(gert::TilingContext *context, SortTileInfo &sortTileInfo)420+void GetRadixSortOneCore(gert::TilingContext* context, SortTileInfo& sortTileInfo)
415{421{
416 sortTileInfo.lastDimNeedCore = static_cast<uint32_t>(1);422 sortTileInfo.lastDimNeedCore = static_cast<uint32_t>(1);
417 sortTileInfo.numTileDataSize = static_cast<uint32_t>(sortTileInfo.sortAxisNum);423 sortTileInfo.numTileDataSize = static_cast<uint32_t>(sortTileInfo.sortAxisNum);
@@ -424,12 +430,12 @@ void GetRadixSortOneCore(gert::TilingContext *context, SortTileInfo &sortTileInf
424 sortTileInfo.coreNumNeed = core == uint32_t(0) ? sortTileInfo.maxCoreNum : core;430 sortTileInfo.coreNumNeed = core == uint32_t(0) ? sortTileInfo.maxCoreNum : core;
425 }431 }
426 sortTileInfo.unsortedDimParallel = sortTileInfo.coreNumNeed;432 sortTileInfo.unsortedDimParallel = sortTileInfo.coreNumNeed;
427- size_t *userWorkSpaceSize = context->GetWorkspaceSizes(1);433+ size_t* userWorkSpaceSize = context->GetWorkspaceSizes(1);
428 userWorkSpaceSize[0] = WORK_SPACE_SIZE;434 userWorkSpaceSize[0] = WORK_SPACE_SIZE;
429 return;435 return;
430}436}
431 437 
432-void ComputeWorkSpace(gert::TilingContext *context, SortTileInfo &sortTileInfo)438+void ComputeWorkSpace(gert::TilingContext* context, SortTileInfo& sortTileInfo)
433{439{
434 uint32_t dtypeSizeWk = static_cast<uint32_t>(sizeof(int32_t));440 uint32_t dtypeSizeWk = static_cast<uint32_t>(sizeof(int32_t));
435 if (sortTileInfo.isInt32 == static_cast<uint32_t>(0)) {441 if (sortTileInfo.isInt32 == static_cast<uint32_t>(0)) {
@@ -438,36 +444,36 @@ void ComputeWorkSpace(gert::TilingContext *context, SortTileInfo &sortTileInfo)
438 size_t excusiveBinsGmWkSize = static_cast<size_t>(sortTileInfo.keyParams1) * sortTileInfo.keyParams4 * dtypeSizeWk;444 size_t excusiveBinsGmWkSize = static_cast<size_t>(sortTileInfo.keyParams1) * sortTileInfo.keyParams4 * dtypeSizeWk;
439 excusiveBinsGmWkSize = CeilDivMul<size_t>(int64_t(excusiveBinsGmWkSize), int64_t(sortTileInfo.blockUbSize));445 excusiveBinsGmWkSize = CeilDivMul<size_t>(int64_t(excusiveBinsGmWkSize), int64_t(sortTileInfo.blockUbSize));
440 446 
441- size_t globalHistGmWkSize =447+ size_t globalHistGmWkSize = static_cast<size_t>(sortTileInfo.keyParams3) * sortTileInfo.keyParams2 *
442- static_cast<size_t>(sortTileInfo.keyParams3) * sortTileInfo.keyParams2 * sortTileInfo.keyParams0 * dtypeSizeWk;448+ sortTileInfo.keyParams0 * dtypeSizeWk;
443 globalHistGmWkSize = CeilDivMul<size_t>(int64_t(globalHistGmWkSize), int64_t(sortTileInfo.blockUbSize));449 globalHistGmWkSize = CeilDivMul<size_t>(int64_t(globalHistGmWkSize), int64_t(sortTileInfo.blockUbSize));
444 450 
445 size_t outIdxDbWK = static_cast<size_t>(sortTileInfo.sortAxisNum) * sortTileInfo.unsortedDimParallel * dtypeSizeWk;451 size_t outIdxDbWK = static_cast<size_t>(sortTileInfo.sortAxisNum) * sortTileInfo.unsortedDimParallel * dtypeSizeWk;
446 outIdxDbWK = CeilDivMul<size_t>(int64_t(outIdxDbWK), int64_t(sortTileInfo.blockUbSize));452 outIdxDbWK = CeilDivMul<size_t>(int64_t(outIdxDbWK), int64_t(sortTileInfo.blockUbSize));
447 453 
448- size_t histTileGmWk = static_cast<size_t>(sortTileInfo.lastDimTileNum) * BIN_NUM * sortTileInfo.unsortedDimParallel *454+ size_t histTileGmWk = static_cast<size_t>(sortTileInfo.lastDimTileNum) * BIN_NUM *
449- sizeof(int16_t) * CONST_2;455+ sortTileInfo.unsortedDimParallel * sizeof(int16_t) * CONST_2;
450 456 
451 size_t xB8GmWkSize = static_cast<size_t>(sortTileInfo.lastDimTileNum) * sortTileInfo.numTileDataSize *457 size_t xB8GmWkSize = static_cast<size_t>(sortTileInfo.lastDimTileNum) * sortTileInfo.numTileDataSize *
452- sortTileInfo.unsortedDimParallel;458+ sortTileInfo.unsortedDimParallel;
453 xB8GmWkSize = CeilDivMul<size_t>(int64_t(xB8GmWkSize), int64_t(sortTileInfo.blockUbSize));459 xB8GmWkSize = CeilDivMul<size_t>(int64_t(xB8GmWkSize), int64_t(sortTileInfo.blockUbSize));
454 460 
455- size_t outValueDbWKSize =461+ size_t outValueDbWKSize = static_cast<size_t>(sortTileInfo.sortAxisNum) * sortTileInfo.unsortedDimParallel *
456- static_cast<size_t>(sortTileInfo.sortAxisNum) * sortTileInfo.unsortedDimParallel * sortTileInfo.dtypeSize;462+ sortTileInfo.dtypeSize;
457 outValueDbWKSize = CeilDivMul<size_t>(int64_t(outValueDbWKSize), int64_t(sortTileInfo.blockUbSize));463 outValueDbWKSize = CeilDivMul<size_t>(int64_t(outValueDbWKSize), int64_t(sortTileInfo.blockUbSize));
458 464 
459 OP_LOGI("RadixSortTiling",465 OP_LOGI("RadixSortTiling",
460- "excusiveBinsGmWkSize %lu, globalHistGmWkSize %lu, outIdxDbWK %lu, histTileGmWk %lu,"466+ "excusiveBinsGmWkSize %lu, globalHistGmWkSize %lu, outIdxDbWK %lu, histTileGmWk %lu,"
461- " xB8GmWkSize %lu, outValueDbWKSize %lu ",467+ " xB8GmWkSize %lu, outValueDbWKSize %lu ",
462- excusiveBinsGmWkSize, globalHistGmWkSize, outIdxDbWK, histTileGmWk, xB8GmWkSize, outValueDbWKSize);468+ excusiveBinsGmWkSize, globalHistGmWkSize, outIdxDbWK, histTileGmWk, xB8GmWkSize, outValueDbWKSize);
463- size_t *userWorkSpaceSize = context->GetWorkspaceSizes(1);469+ size_t* userWorkSpaceSize = context->GetWorkspaceSizes(1);
464- size_t usrSize =470+ size_t usrSize = excusiveBinsGmWkSize + globalHistGmWkSize + outIdxDbWK + histTileGmWk + xB8GmWkSize +
465- excusiveBinsGmWkSize + globalHistGmWkSize + outIdxDbWK + histTileGmWk + xB8GmWkSize + outValueDbWKSize;471+ outValueDbWKSize;
466 userWorkSpaceSize[0] = usrSize + WORK_SPACE_SIZE;472 userWorkSpaceSize[0] = usrSize + WORK_SPACE_SIZE;
467 return;473 return;
468}474}
469 475 
470-void GetRadixSortMoreCore(gert::TilingContext *context, SortTileInfo &sortTileInfo)476+void GetRadixSortMoreCore(gert::TilingContext* context, SortTileInfo& sortTileInfo)
471{477{
472 // 侵入修改:DELETE。在stateless_randperm里已经减去SIMT的空间,这里不需要再减478 // 侵入修改:DELETE。在stateless_randperm里已经减去SIMT的空间,这里不需要再减
473 uint32_t tileData = ComputeTileData(sortTileInfo);479 uint32_t tileData = ComputeTileData(sortTileInfo);
@@ -475,7 +481,8 @@ void GetRadixSortMoreCore(gert::TilingContext *context, SortTileInfo &sortTileIn
475 if (sortTileInfo.maxCoreNum <= lastDimTileNum) {481 if (sortTileInfo.maxCoreNum <= lastDimTileNum) {
476 sortTileInfo.unsortedDimParallel = static_cast<uint32_t>(1);482 sortTileInfo.unsortedDimParallel = static_cast<uint32_t>(1);
477 } else {483 } else {
478- sortTileInfo.unsortedDimParallel = lastDimTileNum == 0 ? sortTileInfo.maxCoreNum : sortTileInfo.maxCoreNum / lastDimTileNum;484+ sortTileInfo.unsortedDimParallel = lastDimTileNum == 0 ? sortTileInfo.maxCoreNum :
485+ sortTileInfo.maxCoreNum / lastDimTileNum;
479 if (sortTileInfo.unSortDimNum < static_cast<int64_t>(sortTileInfo.unsortedDimParallel)) {486 if (sortTileInfo.unSortDimNum < static_cast<int64_t>(sortTileInfo.unsortedDimParallel)) {
480 sortTileInfo.unsortedDimParallel = static_cast<uint32_t>(sortTileInfo.unSortDimNum);487 sortTileInfo.unsortedDimParallel = static_cast<uint32_t>(sortTileInfo.unSortDimNum);
481 }488 }
@@ -493,15 +500,15 @@ void GetRadixSortMoreCore(gert::TilingContext *context, SortTileInfo &sortTileIn
493 uint32_t allNumGloblHist = BIN_NUM * lastDimTileNum * sortTileInfo.dtypeSize * sortTileInfo.unsortedDimParallel;500 uint32_t allNumGloblHist = BIN_NUM * lastDimTileNum * sortTileInfo.dtypeSize * sortTileInfo.unsortedDimParallel;
494 uint32_t allNumExcusiveBin = BIN_NUM * sortTileInfo.dtypeSize * sortTileInfo.unsortedDimParallel;501 uint32_t allNumExcusiveBin = BIN_NUM * sortTileInfo.dtypeSize * sortTileInfo.unsortedDimParallel;
495 uint32_t oneCoreSize = CeilDiv(int64_t(allNumGloblHist), int64_t(sortTileInfo.coreNumNeed));502 uint32_t oneCoreSize = CeilDiv(int64_t(allNumGloblHist), int64_t(sortTileInfo.coreNumNeed));
496- sortTileInfo.keyParams5 =503+ sortTileInfo.keyParams5 = std::max(static_cast<int64_t>(oneCoreSize),
497- std::max(static_cast<int64_t>(oneCoreSize), static_cast<int64_t>(sortTileInfo.blockUbSize));504+ static_cast<int64_t>(sortTileInfo.blockUbSize));
498 sortTileInfo.keyParams0 = CeilDiv(int64_t(allNumGloblHist), int64_t(sortTileInfo.keyParams5));505 sortTileInfo.keyParams0 = CeilDiv(int64_t(allNumGloblHist), int64_t(sortTileInfo.keyParams5));
499 sortTileInfo.keyParams3 = CeilDiv(int64_t(sortTileInfo.keyParams5), int64_t(ubSizeNum));506 sortTileInfo.keyParams3 = CeilDiv(int64_t(sortTileInfo.keyParams5), int64_t(ubSizeNum));
500 sortTileInfo.keyParams2 = sortTileInfo.keyParams5 > ubSizeNum ? ubSizeNum : sortTileInfo.keyParams5;507 sortTileInfo.keyParams2 = sortTileInfo.keyParams5 > ubSizeNum ? ubSizeNum : sortTileInfo.keyParams5;
501 508 
502 uint32_t oneCoreSize1 = CeilDiv(int64_t(allNumExcusiveBin), int64_t(sortTileInfo.coreNumNeed));509 uint32_t oneCoreSize1 = CeilDiv(int64_t(allNumExcusiveBin), int64_t(sortTileInfo.coreNumNeed));
503- sortTileInfo.keyParams4 =510+ sortTileInfo.keyParams4 = std::max(static_cast<int64_t>(oneCoreSize1),
504- std::max(static_cast<int64_t>(oneCoreSize1), static_cast<int64_t>(sortTileInfo.blockUbSize));511+ static_cast<int64_t>(sortTileInfo.blockUbSize));
505 512 
506 sortTileInfo.keyParams1 = CeilDiv(int64_t(allNumExcusiveBin), int64_t(sortTileInfo.keyParams4));513 sortTileInfo.keyParams1 = CeilDiv(int64_t(allNumExcusiveBin), int64_t(sortTileInfo.keyParams4));
507 ComputeWorkSpace(context, sortTileInfo);514 ComputeWorkSpace(context, sortTileInfo);
@@ -509,7 +516,7 @@ void GetRadixSortMoreCore(gert::TilingContext *context, SortTileInfo &sortTileIn
509 return;516 return;
510}517}
511 518 
512-void FillTilingDataSort(SortTileInfo &sortTileInfo, SortRegBaseTilingData *sortTilingData)519+void FillTilingDataSort(SortTileInfo& sortTileInfo, SortRegBaseTilingData* sortTilingData)
513{520{
514 sortTilingData->numTileDataSize = sortTileInfo.numTileDataSize;521 sortTilingData->numTileDataSize = sortTileInfo.numTileDataSize;
515 sortTilingData->unsortedDimParallel = sortTileInfo.unsortedDimParallel;522 sortTilingData->unsortedDimParallel = sortTileInfo.unsortedDimParallel;
@@ -528,21 +535,22 @@ void FillTilingDataSort(SortTileInfo &sortTileInfo, SortRegBaseTilingData *sortT
528 return;535 return;
529}536}
530 537 
531-void PrintTilindDataSort(gert::TilingContext *context, SortTileInfo &sortTileInfo)538+void PrintTilindDataSort(gert::TilingContext* context, SortTileInfo& sortTileInfo)
532{539{
533 OP_LOGI(context->GetNodeName(),540 OP_LOGI(context->GetNodeName(),
534- "realCoreNum %u, numTileDataSize %u, unsortedDimParallel %u, "541+ "realCoreNum %u, numTileDataSize %u, unsortedDimParallel %u, "
535- "lastDimTileNum %u, sortLoopTimes %u, lastDimNeedCore %u, keyParams0 %u, keyParams1 %u "542+ "lastDimTileNum %u, sortLoopTimes %u, lastDimNeedCore %u, keyParams0 %u, keyParams1 %u "
536- "keyParams2 %u, keyParams3 %u, keyParams4 %u, keyParams5 %u, tmpUbSize %u, "543+ "keyParams2 %u, keyParams3 %u, keyParams4 %u, keyParams5 %u, tmpUbSize %u, "
537- "lastAxisNum %ld, unsortedDimNum %ld ",544+ "lastAxisNum %ld, unsortedDimNum %ld ",
538- sortTileInfo.coreNumNeed, sortTileInfo.numTileDataSize, sortTileInfo.unsortedDimParallel,545+ sortTileInfo.coreNumNeed, sortTileInfo.numTileDataSize, sortTileInfo.unsortedDimParallel,
539- sortTileInfo.lastDimTileNum, sortTileInfo.sortLoopTimes, sortTileInfo.lastDimNeedCore, sortTileInfo.keyParams0,546+ sortTileInfo.lastDimTileNum, sortTileInfo.sortLoopTimes, sortTileInfo.lastDimNeedCore,
540- sortTileInfo.keyParams1, sortTileInfo.keyParams2, sortTileInfo.keyParams3, sortTileInfo.keyParams4,547+ sortTileInfo.keyParams0, sortTileInfo.keyParams1, sortTileInfo.keyParams2, sortTileInfo.keyParams3,
541- sortTileInfo.keyParams5, sortTileInfo.tmpUbSize, sortTileInfo.sortAxisNum, sortTileInfo.unSortDimNum);548+ sortTileInfo.keyParams4, sortTileInfo.keyParams5, sortTileInfo.tmpUbSize, sortTileInfo.sortAxisNum,
549+ sortTileInfo.unSortDimNum);
542 return;550 return;
543}551}
544 552 
545-void GetMergeSort(gert::TilingContext *context, SortTileInfo &sortTileInfo)553+void GetMergeSort(gert::TilingContext* context, SortTileInfo& sortTileInfo)
546{554{
547 uint32_t alignNum = CeilDivMul<uint32_t>(int64_t(sortTileInfo.sortAxisNum), int64_t(sortTileInfo.blockUbSize));555 uint32_t alignNum = CeilDivMul<uint32_t>(int64_t(sortTileInfo.sortAxisNum), int64_t(sortTileInfo.blockUbSize));
548 if (alignNum == 0) {556 if (alignNum == 0) {
@@ -564,8 +572,8 @@ void GetMergeSort(gert::TilingContext *context, SortTileInfo &sortTileInfo)
564 } else {572 } else {
565 coreNumNeed = sortTileInfo.maxCoreNum;573 coreNumNeed = sortTileInfo.maxCoreNum;
566 }574 }
567- uint32_t maxTypeSize =575+ uint32_t maxTypeSize = (sortTileInfo.dataType == ge::DT_BF16) ? mergeType.find(ge::DT_FLOAT)->second :
568- (sortTileInfo.dataType == ge::DT_BF16) ? mergeType.find(ge::DT_FLOAT)->second : sortTileInfo.dtypeSize;576+ sortTileInfo.dtypeSize;
569 auto platform_info = context->GetPlatformInfo();577 auto platform_info = context->GetPlatformInfo();
570 auto plat = platform_ascendc::PlatformAscendC(platform_info);578 auto plat = platform_ascendc::PlatformAscendC(platform_info);
571 uint32_t dataSizeNeed = AscendC::GetConcatTmpSize(plat, alignNum, maxTypeSize);579 uint32_t dataSizeNeed = AscendC::GetConcatTmpSize(plat, alignNum, maxTypeSize);
@@ -581,25 +589,25 @@ void GetMergeSort(gert::TilingContext *context, SortTileInfo &sortTileInfo)
581 sortTileInfo.keyParams1 = alignNum * oneCoreRowNum * sortTileInfo.dtypeSize;589 sortTileInfo.keyParams1 = alignNum * oneCoreRowNum * sortTileInfo.dtypeSize;
582 sortTileInfo.keyParams2 = alignNum * oneCoreRowNum * sortTileInfo.y2DtypeSize;590 sortTileInfo.keyParams2 = alignNum * oneCoreRowNum * sortTileInfo.y2DtypeSize;
583 sortTileInfo.keyParams3 = alignNum;591 sortTileInfo.keyParams3 = alignNum;
584- size_t *userWorkSpaceSize = context->GetWorkspaceSizes(1);592+ size_t* userWorkSpaceSize = context->GetWorkspaceSizes(1);
585 userWorkSpaceSize[0] = WORK_SPACE_SIZE;593 userWorkSpaceSize[0] = WORK_SPACE_SIZE;
586}594}
587 595 
588-ge::graphStatus RadixSortTiling(gert::TilingContext *context, int32_t maxCoreNum)596+ge::graphStatus RadixSortTiling(gert::TilingContext* context, int32_t maxCoreNum)
589{597{
590- SortRegBaseTilingData *sortTilingData{ nullptr };598+ SortRegBaseTilingData* sortTilingData{nullptr};
591 sortTilingData = context->GetTilingData<SortRegBaseTilingData>();599 sortTilingData = context->GetTilingData<SortRegBaseTilingData>();
592- OP_CHECK_IF(sortTilingData == nullptr,600+ OP_CHECK_IF(sortTilingData == nullptr, OP_LOGE(context->GetNodeName(), "get tilingdata ptr failed"),
593- OP_LOGE(context->GetNodeName(), "get tilingdata ptr failed"), return ge::GRAPH_FAILED);601+ return ge::GRAPH_FAILED);
594 OP_CHECK_IF((memset_s(sortTilingData, sizeof(SortRegBaseTilingData), 0, sizeof(SortRegBaseTilingData)) != EOK),602 OP_CHECK_IF((memset_s(sortTilingData, sizeof(SortRegBaseTilingData), 0, sizeof(SortRegBaseTilingData)) != EOK),
595- OP_LOGE(context->GetNodeName(), "memset tilingdata failed"), return ge::GRAPH_FAILED);603+ OP_LOGE(context->GetNodeName(), "memset tilingdata failed"), return ge::GRAPH_FAILED);
596 SortTileInfo sortTileInfo;604 SortTileInfo sortTileInfo;
597 OP_CHECK_IF(SortCheckParams(context, sortTileInfo) != ge::GRAPH_SUCCESS,605 OP_CHECK_IF(SortCheckParams(context, sortTileInfo) != ge::GRAPH_SUCCESS,
598- OP_LOGE(context->GetNodeName(), "check params failed"), return ge::GRAPH_FAILED);606+ OP_LOGE(context->GetNodeName(), "check params failed"), return ge::GRAPH_FAILED);
599 sortTileInfo.maxCoreNum = static_cast<uint32_t>(maxCoreNum);607 sortTileInfo.maxCoreNum = static_cast<uint32_t>(maxCoreNum);
600 int64_t int32Max = static_cast<int64_t>(std::numeric_limits<int32_t>::max());608 int64_t int32Max = static_cast<int64_t>(std::numeric_limits<int32_t>::max());
601 uint64_t isInt32 = static_cast<uint64_t>((sortTileInfo.sortAxisNum <= int32Max));609 uint64_t isInt32 = static_cast<uint64_t>((sortTileInfo.sortAxisNum <= int32Max));
602- const bool *isDescending = context->GetAttrs()->GetAttrPointer<bool>(1);610+ const bool* isDescending = context->GetAttrs()->GetAttrPointer<bool>(1);
603 uint64_t isDescend = *isDescending;611 uint64_t isDescend = *isDescending;
604 sortTileInfo.isDescend = static_cast<bool>(isDescend);612 sortTileInfo.isDescend = static_cast<bool>(isDescend);
605 sortTileInfo.isInt32 = static_cast<uint32_t>(isInt32);613 sortTileInfo.isInt32 = static_cast<uint32_t>(isInt32);
@@ -625,14 +633,14 @@ ge::graphStatus RadixSortTiling(gert::TilingContext *context, int32_t maxCoreNum
625 context->SetLocalMemorySize(sortTileInfo.ubSize);633 context->SetLocalMemorySize(sortTileInfo.ubSize);
626 FillTilingDataSort(sortTileInfo, sortTilingData);634 FillTilingDataSort(sortTileInfo, sortTilingData);
627 PrintTilindDataSort(context, sortTileInfo);635 PrintTilindDataSort(context, sortTileInfo);
628- OP_LOGI(context->GetNodeName(), "end RadixSortTIling ");636+ OP_LOGI(context->GetNodeName(), "end RadixSortTiling ");
629 return ge::GRAPH_SUCCESS;637 return ge::GRAPH_SUCCESS;
630}638}
631 639 
632-ge::graphStatus SortTilingSimt(gert::TilingContext *context, int32_t maxCoreNum)640+ge::graphStatus SortTilingSimt(gert::TilingContext* context, int32_t maxCoreNum)
633{641{
634 return RadixSortTiling(context, maxCoreNum);642 return RadixSortTiling(context, maxCoreNum);
635}643}
636 644 
637-}645+} // namespace statelessRandpermTiling
638-}646+} // namespace optiling
@@ -14,11 +14,7 @@ import numpy as np
14from ml_dtypes import bfloat1614from ml_dtypes import bfloat16
15 15 
16 16 
17-__golden__ = {17+__golden__ = {"kernel": {"stateless_randperm": "stateless_randperm_golden"}}
18- "kernel": {
19- "stateless_randperm": "stateless_randperm_golden"
20- }
21-}
22 18 
23 19 
24class RandpermGpuClient:20class RandpermGpuClient:
@@ -36,10 +32,11 @@ class RandpermGpuClient:
36 import struct32 import struct
37 import torch33 import torch
38 import numpy as np34 import numpy as np
35+ 
39 self._deps_loaded = True36 self._deps_loaded = True
40 37 
41 def _recv_all(self, sock, n):38 def _recv_all(self, sock, n):
42- data = b''39+ data = b""
43 while len(data) < n:40 while len(data) < n:
44 packet = sock.recv(n - len(data))41 packet = sock.recv(n - len(data))
45 if not packet:42 if not packet:
@@ -49,18 +46,18 @@ class RandpermGpuClient:
49 46 
50 def _send_msg(self, sock, msg):47 def _send_msg(self, sock, msg):
51 msg = pickle.dumps(msg)48 msg = pickle.dumps(msg)
52- msg = struct.pack('>Q', len(msg)) + msg49+ msg = struct.pack(">Q", len(msg)) + msg
53 sock.sendall(msg)50 sock.sendall(msg)
54 51 
55 def _recv_msg(self, sock):52 def _recv_msg(self, sock):
56 raw_msglen = self._recv_all(sock, 8)53 raw_msglen = self._recv_all(sock, 8)
57 if not raw_msglen:54 if not raw_msglen:
58 return None55 return None
59- msglen = struct.unpack('>Q', raw_msglen)[0]56+ msglen = struct.unpack(">Q", raw_msglen)[0]
60 return pickle.loads(self._recv_all(sock, msglen))57 return pickle.loads(self._recv_all(sock, msglen))
61 58 
62 def compute_on_gpu(self, seed=42, offset=10, n=100, dtype=9):59 def compute_on_gpu(self, seed=42, offset=10, n=100, dtype=9):
63- request = {'seed': seed, 'offset': offset, 'n': n, 'dtype': dtype}60+ request = {"seed": seed, "offset": offset, "n": n, "dtype": dtype}
64 try:61 try:
65 with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:62 with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
66 s.settimeout(30)63 s.settimeout(30)
@@ -69,29 +66,31 @@ class RandpermGpuClient:
69 result = self._recv_msg(s)66 result = self._recv_msg(s)
70 return result67 return result
71 except Exception as e:68 except Exception as e:
72- print(f"连接错误: {e}")69+ print(f"Connection error: {e}")
73 return None70 return None
74 71 
75 72 
76def compute_local(seed, offset, n):73def compute_local(seed, offset, n):
77 import torch74 import torch
78- generator = torch.Generator(device='cpu')75+ 
76+ generator = torch.Generator(device="cpu")
79 generator.manual_seed(seed)77 generator.manual_seed(seed)
80 generator.set_offset(offset)78 generator.set_offset(offset)
81- result = torch.randperm(n, generator=generator, device='cpu')79+ result = torch.randperm(n, generator=generator, device="cpu")
82 return result.numpy()80 return result.numpy()
83 81 
84 82 
85def stateless_randperm_golden(n, seed, offset, layout=0, dtype=9, **kwargs):83def stateless_randperm_golden(n, seed, offset, layout=0, dtype=9, **kwargs):
86- '''84+ """
87 Kernel golden for stateless_randperm.85 Kernel golden for stateless_randperm.
88 All the parameters follow @stateless_randperm_def.cpp without outputs.86 All the parameters follow @stateless_randperm_def.cpp without outputs.
89 All the input Tensors are numpy.ndarray.87 All the input Tensors are numpy.ndarray.
90 kwargs may contain: short_soc_version, input_ori_shapes, output_ori_shapes,88 kwargs may contain: short_soc_version, input_ori_shapes, output_ori_shapes,
91 input_formats, output_formats, input_ori_formats, output_ori_formats,89 input_formats, output_formats, input_ori_formats, output_ori_formats,
92 input_dtypes, output_dtypes.90 input_dtypes, output_dtypes.
93- '''91+ """
94 import logging92 import logging
93+ 
95 GPU_SERVER_IP = "x.x.x.x"94 GPU_SERVER_IP = "x.x.x.x"
96 GPU_SERVER_PORT = 3232395 GPU_SERVER_PORT = 32323
97 96 
@@ -100,13 +99,19 @@ def stateless_randperm_golden(n, seed, offset, layout=0, dtype=9, **kwargs):
100 offset_val = int(np.array(offset).flatten()[0])99 offset_val = int(np.array(offset).flatten()[0])
101 100 
102 client = RandpermGpuClient(GPU_SERVER_IP, GPU_SERVER_PORT)101 client = RandpermGpuClient(GPU_SERVER_IP, GPU_SERVER_PORT)
103- result = client.compute_on_gpu(seed=seed_val, offset=offset_val, n=n_val, dtype=dtype)102+ result = client.compute_on_gpu(
103+ seed=seed_val, offset=offset_val, n=n_val, dtype=dtype
104+ )
104 if dtype == 27:105 if dtype == 27:
105 result = result.astype(bfloat16)106 result = result.astype(bfloat16)
106- logging.info("remote gpu computation is done, the result type is {}.".format(type(result)))107+ logging.info(
108+ "remote gpu computation is done, the result type is {}.".format(type(result))
109+ )
107 110 
108 if result is None:111 if result is None:
109- logging.warning(f"remote gpu computation failed, switch to local cpu computation.")112+ logging.warning(
113+ "remote gpu computation failed, switch to local cpu computation."
114+ )
110 result = compute_local(seed=seed_val, offset=offset_val, n=n_val)115 result = compute_local(seed=seed_val, offset=offset_val, n=n_val)
111 116 
112 return result117 return result
@@ -92,7 +92,7 @@ bool CheckTilingDataDefinitions()
92 typesMatch = CompareStructMembers<SortRegBaseTilingData, SortRegBaseTilingDataForRandperm, count1 - 1>::value;92 typesMatch = CompareStructMembers<SortRegBaseTilingData, SortRegBaseTilingDataForRandperm, count1 - 1>::value;
93 93 
94 if (!typesMatch) {94 if (!typesMatch) {
95- std::cout << "Error: Struct member type is not same." << std::endl;95+ std::cout << "Error: Struct member type is not the same." << std::endl;
96 return false;96 return false;
97 }97 }
98 98 
@@ -65,7 +65,7 @@ static Status ParseOpToGraphMultinomial(const ge::Operator& op, Graph& graph)
65 new_op.SetAttr("seed", seed);65 new_op.SetAttr("seed", seed);
66 ge::DataType dtype_om = GetOmDtypeFromOnnxDtype(data_type);66 ge::DataType dtype_om = GetOmDtypeFromOnnxDtype(data_type);
67 if (dtype_om == ge::DT_UNDEFINED) {67 if (dtype_om == ge::DT_UNDEFINED) {
68- OP_LOGE(GetOpName(op).c_str(), "dtype[%d] is wrong,please select right dtype", data_type);68+ OP_LOGE(GetOpName(op).c_str(), "dtype[%d] is wrong, please select a valid ONNX dtype", data_type);
69 return FAILED;69 return FAILED;
70 }70 }
71 int int_seed = static_cast<int>(seed);71 int int_seed = static_cast<int>(seed);
@@ -153,27 +153,29 @@ static bool CheckDtypeValidTensor(const aclTensor* self, const aclTensor* seedTe
153static bool CheckShape(const aclTensor* self, int64_t numsamples, const aclTensor* out)153static bool CheckShape(const aclTensor* self, int64_t numsamples, const aclTensor* out)
154{154{
155 if (self->GetViewShape().GetDimNum() != DIM_NUM_ONE && self->GetViewShape().GetDimNum() != DIM_NUM_TWO) {155 if (self->GetViewShape().GetDimNum() != DIM_NUM_ONE && self->GetViewShape().GetDimNum() != DIM_NUM_TWO) {
156- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Dim of self only can be 1 or 2.");156+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Dim of self can only be 1 or 2, but got %zu.",
157+ self->GetViewShape().GetDimNum());
157 return false;158 return false;
158 }159 }
159 if (self->GetViewShape().GetDimNum() != out->GetViewShape().GetDimNum()) {160 if (self->GetViewShape().GetDimNum() != out->GetViewShape().GetDimNum()) {
160- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Dim of self should be equal to dim of out.");161+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Dim of self %zu should be equal to dim of out %zu.",
162+ self->GetViewShape().GetDimNum(), out->GetViewShape().GetDimNum());
161 return false;163 return false;
162 }164 }
163 auto dimNum = out->GetViewShape().GetDimNum();165 auto dimNum = out->GetViewShape().GetDimNum();
164 auto nCategories = out->GetViewShape().GetDim(dimNum - 1);166 auto nCategories = out->GetViewShape().GetDim(dimNum - 1);
165 if (nCategories != numsamples) {167 if (nCategories != numsamples) {
166 OP_LOGE(ACLNN_ERR_PARAM_INVALID,168 OP_LOGE(ACLNN_ERR_PARAM_INVALID,
167- "excepted the size of out at last dim must be equal with numsamples %ld, but got %ld.", numsamples,169+ "expected the size of out at last dim must be equal to numsamples %ld, but got %ld.", numsamples,
168 nCategories);170 nCategories);
169 return false;171 return false;
170 }172 }
171 if (self->GetViewShape().GetDimNum() != DIM_NUM_ONE &&173 if (self->GetViewShape().GetDimNum() != DIM_NUM_ONE &&
172 self->GetViewShape().GetDim(DIM_ZERO) != out->GetViewShape().GetDim(DIM_ZERO)) {174 self->GetViewShape().GetDim(DIM_ZERO) != out->GetViewShape().GetDim(DIM_ZERO)) {
173 OP_LOGE(ACLNN_ERR_PARAM_INVALID,175 OP_LOGE(ACLNN_ERR_PARAM_INVALID,
174- "excepted the size of out at first dim %ld must be equal with the size of self at first dim %ld when "176+ "expected the size of out at first dim %ld must be equal to the size of self at first dim %ld when "
175 "dimNum > 1",177 "dimNum > 1",
176- self->GetViewShape().GetDim(DIM_ZERO), out->GetViewShape().GetDim(DIM_ZERO));178+ out->GetViewShape().GetDim(DIM_ZERO), self->GetViewShape().GetDim(DIM_ZERO));
177 return false;179 return false;
178 }180 }
179 return true;181 return true;
@@ -182,17 +184,19 @@ static bool CheckShape(const aclTensor* self, int64_t numsamples, const aclTenso
182static bool CheckValueRange(const aclTensor* self, int64_t numsamples, bool replacement)184static bool CheckValueRange(const aclTensor* self, int64_t numsamples, bool replacement)
183{185{
184 if (numsamples <= 0) {186 if (numsamples <= 0) {
185- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Numsamples must > 0.");187+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Numsamples must be greater than 0, but got %ld.", numsamples);
186 return false;188 return false;
187 }189 }
188 auto dimNum = self->GetViewShape().GetDimNum();190 auto dimNum = self->GetViewShape().GetDimNum();
189 auto nCategories = self->GetViewShape().GetDim(dimNum - 1);191 auto nCategories = self->GetViewShape().GetDim(dimNum - 1);
190 if (!replacement && (numsamples > nCategories)) {192 if (!replacement && (numsamples > nCategories)) {
191- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "numsamples must <= shape.GetDim(dimNum - 1) without replacement");193+ OP_LOGE(ACLNN_ERR_PARAM_INVALID,
194+ "numsamples %ld must be no more than the size of last dim %ld without replacement", numsamples,
195+ nCategories);
192 return false;196 return false;
193 }197 }
194 if (nCategories > FLOAT32_MAX_CONSECUTIVE_INT) {198 if (nCategories > FLOAT32_MAX_CONSECUTIVE_INT) {
195- OP_LOGE(ACLNN_ERR_PARAM_INVALID, "number of categories cannot exceed 2^24");199+ OP_LOGE(ACLNN_ERR_PARAM_INVALID, "number of categories %ld cannot exceed 2^24", nCategories);
196 return false;200 return false;
197 }201 }
198 return true;202 return true;
@@ -80,9 +80,12 @@ ge::graphStatus StatelessSampleMultinomialTiling::CheckXRankAndNormProbsShape()
80 auto normProbsShapePtr = context_->GetOptionalInputShape(INPUT_IDX_NORM_PROBS);80 auto normProbsShapePtr = context_->GetOptionalInputShape(INPUT_IDX_NORM_PROBS);
81 if (normProbsShapePtr != nullptr) {81 if (normProbsShapePtr != nullptr) {
82 const auto& normProbsShape = normProbsShapePtr->GetStorageShape();82 const auto& normProbsShape = normProbsShapePtr->GetStorageShape();
83- OP_CHECK_IF(normProbsShape != xShape,83+ OP_CHECK_IF(
84- OP_LOGE(context_->GetNodeName(), "the shapes of x and norm_probs must be the same"),84+ normProbsShape != xShape,
85- return ge::GRAPH_FAILED);85+ OP_LOGE(context_->GetNodeName(),
86+ "the shapes of x and norm_probs must be the same, x shape size: %ld, norm_probs shape size: %ld.",
87+ xShape.GetShapeSize(), normProbsShape.GetShapeSize()),
88+ return ge::GRAPH_FAILED);
86 }89 }
87 return ge::GRAPH_SUCCESS;90 return ge::GRAPH_SUCCESS;
88}91}
@@ -37,10 +37,10 @@ using std::map;
37using std::string;37using std::string;
38using std::vector;38using std::vector;
39 39 
40-#define LOG_PRINT(message, ...) \40+#define LOG_PRINT(message, ...) \
41- do { \41+ do { \
42- printf(message, ##__VA_ARGS__); \42+ printf(message, ##__VA_ARGS__); \
43- } while (0)43+ } while (0)
44 44 
45string GetTime()45string GetTime()
46{46{
@@ -53,26 +53,33 @@ string GetTime()
53 53 
54uint32_t GetDataTypeSize(DataType dt)54uint32_t GetDataTypeSize(DataType dt)
55{55{
56- if (dt == ge::DT_FLOAT) return 4;56+ if (dt == ge::DT_FLOAT)
57- if (dt == ge::DT_FLOAT16) return 2;57+ return 4;
58- if (dt == ge::DT_BF16) return 2;58+ if (dt == ge::DT_FLOAT16)
59- if (dt == ge::DT_INT32) return 4;59+ return 2;
60- if (dt == ge::DT_UINT32) return 4;60+ if (dt == ge::DT_BF16)
61- if (dt == ge::DT_INT64) return 8;61+ return 2;
62- if (dt == ge::DT_UINT64) return 8;62+ if (dt == ge::DT_INT32)
63+ return 4;
64+ if (dt == ge::DT_UINT32)
65+ return 4;
66+ if (dt == ge::DT_INT64)
67+ return 8;
68+ if (dt == ge::DT_UINT64)
69+ return 8;
63 return 1;70 return 1;
64}71}
65 72 
66-int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t *inputData)73+int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
67{74{
68- FILE *fp = fopen(bin_file.c_str(), "w");75+ FILE* fp = fopen(bin_file.c_str(), "w");
69 fwrite(inputData, sizeof(uint8_t), data_size, fp);76 fwrite(inputData, sizeof(uint8_t), data_size, fp);
70 fclose(fp);77 fclose(fp);
71 return SUCCESS;78 return SUCCESS;
72}79}
73 80 
74-int CreateOppInGraph(std::vector<ge::Tensor> &input, std::vector<Operator> &inputs,81+int CreateOppInGraph(std::vector<ge::Tensor>& input, std::vector<Operator>& inputs, std::vector<Operator>& outputs,
75- std::vector<Operator> &outputs, Graph &graph)82+ Graph& graph)
76{83{
77 // StatelessTruncatedNormalV2 op84 // StatelessTruncatedNormalV2 op
78 auto op1 = op::StatelessTruncatedNormalV2("stateless_truncated_normal_v2");85 auto op1 = op::StatelessTruncatedNormalV2("stateless_truncated_normal_v2");
@@ -82,7 +89,7 @@ int CreateOppInGraph(std::vector<ge::Tensor> &input, std::vector<Operator> &inpu
82 auto shapeNode = op::Const("shape_const");89 auto shapeNode = op::Const("shape_const");
83 TensorDesc shapeDesc(ge::Shape(shapeShape), FORMAT_ND, DT_INT32);90 TensorDesc shapeDesc(ge::Shape(shapeShape), FORMAT_ND, DT_INT32);
84 shapeDesc.SetPlacement(ge::kPlacementHost);91 shapeDesc.SetPlacement(ge::kPlacementHost);
85- int32_t shapeData[] = {4, 8}; // output shape: [4, 8]92+ int32_t shapeData[] = {4, 8}; // output shape: [4, 8]
86 Tensor shapeTensor(shapeDesc, reinterpret_cast<uint8_t*>(shapeData), sizeof(shapeData));93 Tensor shapeTensor(shapeDesc, reinterpret_cast<uint8_t*>(shapeData), sizeof(shapeData));
87 shapeNode.SetAttr("value", shapeTensor);94 shapeNode.SetAttr("value", shapeTensor);
88 shapeNode.update_output_desc_y(shapeDesc);95 shapeNode.update_output_desc_y(shapeDesc);
@@ -122,7 +129,7 @@ int CreateOppInGraph(std::vector<ge::Tensor> &input, std::vector<Operator> &inpu
122 auto algNode = op::Data("alg_data").set_attr_index(3);129 auto algNode = op::Data("alg_data").set_attr_index(3);
123 TensorDesc algDesc(ge::Shape(algShape), FORMAT_ND, DT_INT32);130 TensorDesc algDesc(ge::Shape(algShape), FORMAT_ND, DT_INT32);
124 algDesc.SetPlacement(ge::kPlacementHost);131 algDesc.SetPlacement(ge::kPlacementHost);
125- int32_t algData[] = {1}; // 1 = Philox132+ int32_t algData[] = {1}; // 1 = Philox
126 Tensor algTensor(algDesc, reinterpret_cast<uint8_t*>(algData), sizeof(algData));133 Tensor algTensor(algDesc, reinterpret_cast<uint8_t*>(algData), sizeof(algData));
127 algNode.update_input_desc_x(algDesc);134 algNode.update_input_desc_x(algDesc);
128 input.push_back(algTensor);135 input.push_back(algTensor);
@@ -131,7 +138,7 @@ int CreateOppInGraph(std::vector<ge::Tensor> &input, std::vector<Operator> &inpu
131 inputs.push_back(algNode);138 inputs.push_back(algNode);
132 139 
133 // Attr: dtype140 // Attr: dtype
134- op1.set_attr_dtype(0); // 0 = float32141+ op1.set_attr_dtype(0); // 0 = float32
135 142 
136 // Output: y143 // Output: y
137 std::vector<int64_t> outShape = {4, 8};144 std::vector<int64_t> outShape = {4, 8};
@@ -142,9 +149,9 @@ int CreateOppInGraph(std::vector<ge::Tensor> &input, std::vector<Operator> &inpu
142 return SUCCESS;149 return SUCCESS;
143}150}
144 151 
145-int main(int argc, char *argv[])152+int main(int argc, char* argv[])
146{153{
147- const char *graph_name = "tc_ge_irrun_test";154+ const char* graph_name = "tc_ge_irrun_test";
148 Graph graph(graph_name);155 Graph graph(graph_name);
149 std::vector<ge::Tensor> input;156 std::vector<ge::Tensor> input;
150 157 
@@ -152,7 +159,7 @@ int main(int argc, char *argv[])
152 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};159 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
153 Status ret = ge::GEInitialize(global_options);160 Status ret = ge::GEInitialize(global_options);
154 if (ret != SUCCESS) {161 if (ret != SUCCESS) {
155- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());162+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
156 return FAILED;163 return FAILED;
157 }164 }
158 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());165 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -172,7 +179,7 @@ int main(int argc, char *argv[])
172 179 
173 std::map<AscendString, AscendString> build_options = {};180 std::map<AscendString, AscendString> build_options = {};
174 printf("%s - INFO - [XIR]: Start to create ir session\n", GetTime().c_str());181 printf("%s - INFO - [XIR]: Start to create ir session\n", GetTime().c_str());
175- ge::Session *session = new Session(build_options);182+ ge::Session* session = new Session(build_options);
176 if (session == nullptr) {183 if (session == nullptr) {
177 printf("%s - ERROR - [XIR]: Create ir session failed\n", GetTime().c_str());184 printf("%s - ERROR - [XIR]: Create ir session failed\n", GetTime().c_str());
178 return FAILED;185 return FAILED;
@@ -191,7 +198,7 @@ int main(int argc, char *argv[])
191 std::vector<ge::Tensor> output;198 std::vector<ge::Tensor> output;
192 ret = session->RunGraph(graph_id, input, output);199 ret = session->RunGraph(graph_id, input, output);
193 if (ret != SUCCESS) {200 if (ret != SUCCESS) {
194- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());201+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
195 delete session;202 delete session;
196 GEFinalize();203 GEFinalize();
197 return FAILED;204 return FAILED;
@@ -202,12 +209,12 @@ int main(int argc, char *argv[])
202 for (int i = 0; i < output_num; i++) {209 for (int i = 0; i < output_num; i++) {
203 std::cout << "output " << i << " dtype: " << output[i].GetTensorDesc().GetDataType() << std::endl;210 std::cout << "output " << i << " dtype: " << output[i].GetTensorDesc().GetDataType() << std::endl;
204 string output_file = "./stateless_truncated_normal_v2_output_" + std::to_string(i) + ".bin";211 string output_file = "./stateless_truncated_normal_v2_output_" + std::to_string(i) + ".bin";
205- uint8_t *output_data_i = output[i].GetData();212+ uint8_t* output_data_i = output[i].GetData();
206 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();213 int64_t output_shape = output[i].GetTensorDesc().GetShape().GetShapeSize();
207 std::cout << "output " << i << " shape size = " << output_shape << std::endl;214 std::cout << "output " << i << " shape size = " << output_shape << std::endl;
208 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());215 uint32_t data_size = output_shape * GetDataTypeSize(output[i].GetTensorDesc().GetDataType());
209 WriteDataToFile(output_file, data_size, output_data_i);216 WriteDataToFile(output_file, data_size, output_data_i);
210- float *resultData = (float*)output_data_i;217+ float* resultData = (float*)output_data_i;
211 for (int64_t j = 0; j < output_shape && j < 32; j++) {218 for (int64_t j = 0; j < output_shape && j < 32; j++) {
212 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);219 LOG_PRINT("result[%ld] is: %f\n", j, resultData[j]);
213 }220 }
@@ -216,7 +223,7 @@ int main(int argc, char *argv[])
216 printf("%s - INFO - [XIR]: Start to finalize\n", GetTime().c_str());223 printf("%s - INFO - [XIR]: Start to finalize\n", GetTime().c_str());
217 ret = ge::GEFinalize();224 ret = ge::GEFinalize();
218 if (ret != SUCCESS) {225 if (ret != SUCCESS) {
219- printf("%s - INFO - [XIR]: Finalize failed\n", GetTime().c_str());226+ printf("%s - ERROR - [XIR]: Finalize failed\n", GetTime().c_str());
220 return FAILED;227 return FAILED;
221 }228 }
222 printf("%s - INFO - [XIR]: Finalize success\n", GetTime().c_str());229 printf("%s - INFO - [XIR]: Finalize success\n", GetTime().c_str());
@@ -100,7 +100,7 @@ OpTilingConfig StatelessTruncatedNormalV2Tiling::BuildOpConfig()
100 100 
101ge::graphStatus StatelessTruncatedNormalV2Tiling::DoSimtBlockTiling()101ge::graphStatus StatelessTruncatedNormalV2Tiling::DoSimtBlockTiling()
102{102{
103- OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is less than or equal to 0. please check."),103+ OP_CHECK_IF((totalCoreNum_ <= 0), OP_LOGE(opName_, "totalCoreNum is %ld, must be greater than 0.", totalCoreNum_),
104 return ge::GRAPH_FAILED);104 return ge::GRAPH_FAILED);
105 int64_t threadNum = Ops::Base::CeilAlign(simtTilingData_.outputSize, THREAD_DISPOSAL_NUM);105 int64_t threadNum = Ops::Base::CeilAlign(simtTilingData_.outputSize, THREAD_DISPOSAL_NUM);
106 int64_t coreNum = Ops::Base::CeilAlign(threadNum, MAX_THREAD_NUM);106 int64_t coreNum = Ops::Base::CeilAlign(threadNum, MAX_THREAD_NUM);
@@ -38,107 +38,97 @@ using namespace ge;
38using std::map;38using std::map;
39using std::string;39using std::string;
40using std::vector;40using std::vector;
41-#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \41+#define ADD_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
42- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \42+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
43- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \43+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
44- TensorDesc placeholder##intputIndex##_desc = \44+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
45- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \45+ intputDtype); \
46- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \46+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
47- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \47+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
48- Tensor tensor_placeholder##intputIndex; \48+ Tensor tensor_placeholder##intputIndex; \
49- ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, \49+ ret = GenOnesDataFloat32(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
50- tensor_placeholder##intputIndex, \50+ placeholder##intputIndex##_desc, value); \
51- placeholder##intputIndex##_desc, \51+ if (ret != SUCCESS) { \
52- value); \52+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
53- if (ret != SUCCESS) { \53+ return FAILED; \
54- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \54+ } \
55- return FAILED; \55+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
56- } \56+ input.push_back(tensor_placeholder##intputIndex); \
57- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \57+ graph.AddOp(placeholder##intputIndex); \
58- input.push_back(tensor_placeholder##intputIndex); \58+ add1.set_input_##intputName(placeholder##intputIndex); \
59- graph.AddOp(placeholder##intputIndex); \
60- add1.set_input_##intputName(placeholder##intputIndex); \
61 inputs.push_back(placeholder##intputIndex)59 inputs.push_back(placeholder##intputIndex)
62 60 
63-#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \61+#define ADD_INT_INPUT(intputIndex, intputName, intputDtype, inputShape, value) \
64- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \62+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
65- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \63+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
66- TensorDesc placeholder##intputIndex##_desc = \64+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
67- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \65+ intputDtype); \
68- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \66+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
69- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \67+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
70- Tensor tensor_placeholder##intputIndex; \68+ Tensor tensor_placeholder##intputIndex; \
71- ret = GenOnesDataInt64(placeholder##intputIndex##_shape, \69+ ret = GenOnesDataInt64(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
72- tensor_placeholder##intputIndex, \70+ placeholder##intputIndex##_desc, value); \
73- placeholder##intputIndex##_desc, \71+ if (ret != SUCCESS) { \
74- value); \72+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
75- if (ret != SUCCESS) { \73+ return FAILED; \
76- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \74+ } \
77- return FAILED; \75+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
78- } \76+ input.push_back(tensor_placeholder##intputIndex); \
79- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \77+ graph.AddOp(placeholder##intputIndex); \
80- input.push_back(tensor_placeholder##intputIndex); \78+ add1.set_input_##intputName(placeholder##intputIndex); \
81- graph.AddOp(placeholder##intputIndex); \
82- add1.set_input_##intputName(placeholder##intputIndex); \
83 inputs.push_back(placeholder##intputIndex)79 inputs.push_back(placeholder##intputIndex)
84 80 
85-#define ADD_DOUBLE_INPUT(intputIndex, intputName, inputShape, value) \81+#define ADD_DOUBLE_INPUT(intputIndex, intputName, inputShape, value) \
86- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \82+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
87- auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \83+ auto placeholder##intputIndex = op::Data("placeholder" + intputIndex).set_attr_index(0); \
88- TensorDesc placeholder##intputIndex##_desc = \84+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
89- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, ge::DT_DOUBLE); \85+ ge::DT_DOUBLE); \
90- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \86+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
91- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \87+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
92- Tensor tensor_placeholder##intputIndex; \88+ Tensor tensor_placeholder##intputIndex; \
93- ret = GenOnesDataDouble(placeholder##intputIndex##_shape, \89+ ret = GenOnesDataDouble(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
94- tensor_placeholder##intputIndex, \90+ placeholder##intputIndex##_desc, value); \
95- placeholder##intputIndex##_desc, \91+ if (ret != SUCCESS) { \
96- value); \92+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
97- if (ret != SUCCESS) { \93+ return FAILED; \
98- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \94+ } \
99- return FAILED; \95+ placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \
100- } \96+ input.push_back(tensor_placeholder##intputIndex); \
101- placeholder##intputIndex.update_input_desc_x(placeholder##intputIndex##_desc); \97+ graph.AddOp(placeholder##intputIndex); \
102- input.push_back(tensor_placeholder##intputIndex); \98+ add1.set_input_##intputName(placeholder##intputIndex); \
103- graph.AddOp(placeholder##intputIndex); \
104- add1.set_input_##intputName(placeholder##intputIndex); \
105 inputs.push_back(placeholder##intputIndex)99 inputs.push_back(placeholder##intputIndex)
106 100 
107-#define ADD_INPUT_ATTR(attrName, attrValue) \101+#define ADD_INPUT_ATTR(attrName, attrValue) add1.set_attr_##attrName(attrValue)
108- add1.set_attr_##attrName(attrValue)
109 102 
110-#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \103+#define ADD_OUTPUT(outputIndex, outputName, outputDtype, outputShape) \
111- TensorDesc outputName##outputIndex##_desc = \104+ TensorDesc outputName##outputIndex##_desc = TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
112- TensorDesc(ge::Shape(outputShape), FORMAT_ND, outputDtype); \
113 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)105 add1.update_output_desc_##outputName(outputName##outputIndex##_desc)
114 106 
115-#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape, constValues) \107+#define ADD_CONST_INPUT(intputIndex, intputName, intputDtype, inputShape, constValues) \
116- vector<int64_t> placeholder##intputIndex##_shape = inputShape; \108+ vector<int64_t> placeholder##intputIndex##_shape = inputShape; \
117- auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \109+ auto placeholder##intputIndex = op::Const("placeholder" + intputIndex); \
118- TensorDesc placeholder##intputIndex##_desc = \110+ TensorDesc placeholder##intputIndex##_desc = TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, \
119- TensorDesc(ge::Shape(placeholder##intputIndex##_shape), FORMAT_ND, intputDtype); \111+ intputDtype); \
120- placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \112+ placeholder##intputIndex##_desc.SetPlacement(ge::kPlacementHost); \
121- placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \113+ placeholder##intputIndex##_desc.SetFormat(FORMAT_ND); \
122- Tensor tensor_placeholder##intputIndex; \114+ Tensor tensor_placeholder##intputIndex; \
123- ret = GenConstDataInt64(placeholder##intputIndex##_shape, \115+ ret = GenConstDataInt64(placeholder##intputIndex##_shape, tensor_placeholder##intputIndex, \
124- tensor_placeholder##intputIndex, \116+ placeholder##intputIndex##_desc, constValues); \
125- placeholder##intputIndex##_desc, \117+ if (ret != SUCCESS) { \
126- constValues); \118+ printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \
127- if (ret != SUCCESS) { \119+ return FAILED; \
128- printf("%s - ERROR - [XIR]: Generate input data failed\n", GetTime().c_str()); \120+ } \
129- return FAILED; \121+ placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \
130- } \122+ placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \
131- placeholder##intputIndex.SetAttr("value", tensor_placeholder##intputIndex); \123+ graph.AddOp(placeholder##intputIndex); \
132- placeholder##intputIndex.update_output_desc_y(placeholder##intputIndex##_desc); \124+ add1.set_input_##intputName(placeholder##intputIndex); \
133- graph.AddOp(placeholder##intputIndex); \125+ add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
134- add1.set_input_##intputName(placeholder##intputIndex); \
135- add1.update_input_desc_##intputName(placeholder##intputIndex##_desc); \
136 inputs.push_back(placeholder##intputIndex)126 inputs.push_back(placeholder##intputIndex)
137 127 
138-#define LOG_PRINT(message, ...) \128+#define LOG_PRINT(message, ...) \
139- do { \129+ do { \
140- printf(message, ##__VA_ARGS__); \130+ printf(message, ##__VA_ARGS__); \
141- } while (0)131+ } while (0)
142 132 
143string GetTime()133string GetTime()
144{134{
@@ -176,7 +166,7 @@ uint32_t GetDataTypeSize(DataType dt)
176 return oneByte;166 return oneByte;
177}167}
178 168 
179-int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, float value)169+int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, float value)
180{170{
181 input_tensor_desc.SetRealDimCnt(shapes.size());171 input_tensor_desc.SetRealDimCnt(shapes.size());
182 size_t size = 1;172 size_t size = 1;
@@ -184,16 +174,16 @@ int32_t GenOnesDataFloat32(vector<int64_t> shapes, Tensor &input_tensor, TensorD
184 size *= shapes[i];174 size *= shapes[i];
185 }175 }
186 uint32_t data_len = size * sizeof(float);176 uint32_t data_len = size * sizeof(float);
187- float *pData = new (std::nothrow) float[size];177+ float* pData = new (std::nothrow) float[size];
188 for (size_t i = 0; i < size; ++i) {178 for (size_t i = 0; i < size; ++i) {
189 *(pData + i) = value;179 *(pData + i) = value;
190 }180 }
191- input_tensor = Tensor(input_tensor_desc, (uint8_t *)pData, data_len);181+ input_tensor = Tensor(input_tensor_desc, (uint8_t*)pData, data_len);
192 delete[] pData;182 delete[] pData;
193 return SUCCESS;183 return SUCCESS;
194}184}
195 185 
196-int32_t GenOnesDataInt64(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, int64_t value)186+int32_t GenOnesDataInt64(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, int64_t value)
197{187{
198 input_tensor_desc.SetRealDimCnt(shapes.size());188 input_tensor_desc.SetRealDimCnt(shapes.size());
199 size_t size = 1;189 size_t size = 1;
@@ -201,16 +191,16 @@ int32_t GenOnesDataInt64(vector<int64_t> shapes, Tensor &input_tensor, TensorDes
201 size *= shapes[i];191 size *= shapes[i];
202 }192 }
203 uint32_t data_len = size * sizeof(int64_t);193 uint32_t data_len = size * sizeof(int64_t);
204- int64_t *pData = new (std::nothrow) int64_t[size];194+ int64_t* pData = new (std::nothrow) int64_t[size];
205 for (size_t i = 0; i < size; ++i) {195 for (size_t i = 0; i < size; ++i) {
206 *(pData + i) = value;196 *(pData + i) = value;
207 }197 }
208- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);198+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
209 delete[] pData;199 delete[] pData;
210 return SUCCESS;200 return SUCCESS;
211}201}
212 202 
213-int32_t GenOnesDataDouble(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc, double value)203+int32_t GenOnesDataDouble(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc, double value)
214{204{
215 input_tensor_desc.SetRealDimCnt(shapes.size());205 input_tensor_desc.SetRealDimCnt(shapes.size());
216 size_t size = 1;206 size_t size = 1;
@@ -218,17 +208,17 @@ int32_t GenOnesDataDouble(vector<int64_t> shapes, Tensor &input_tensor, TensorDe
218 size *= shapes[i];208 size *= shapes[i];
219 }209 }
220 uint32_t data_len = size * sizeof(double);210 uint32_t data_len = size * sizeof(double);
221- double *pData = new (std::nothrow) double[size];211+ double* pData = new (std::nothrow) double[size];
222 for (size_t i = 0; i < size; ++i) {212 for (size_t i = 0; i < size; ++i) {
223 *(pData + i) = value;213 *(pData + i) = value;
224 }214 }
225- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);215+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
226 delete[] pData;216 delete[] pData;
227 return SUCCESS;217 return SUCCESS;
228}218}
229 219 
230-int32_t GenConstDataInt64(vector<int64_t> shapes, Tensor &input_tensor, TensorDesc &input_tensor_desc,220+int32_t GenConstDataInt64(vector<int64_t> shapes, Tensor& input_tensor, TensorDesc& input_tensor_desc,
231- const vector<int64_t> &values)221+ const vector<int64_t>& values)
232{222{
233 input_tensor_desc.SetRealDimCnt(shapes.size());223 input_tensor_desc.SetRealDimCnt(shapes.size());
234 size_t size = 1;224 size_t size = 1;
@@ -236,25 +226,25 @@ int32_t GenConstDataInt64(vector<int64_t> shapes, Tensor &input_tensor, TensorDe
236 size *= shapes[i];226 size *= shapes[i];
237 }227 }
238 uint32_t data_len = size * sizeof(int64_t);228 uint32_t data_len = size * sizeof(int64_t);
239- int64_t *pData = new (std::nothrow) int64_t[size];229+ int64_t* pData = new (std::nothrow) int64_t[size];
240 for (size_t i = 0; i < size; ++i) {230 for (size_t i = 0; i < size; ++i) {
241 *(pData + i) = (i < values.size()) ? values[i] : 0;231 *(pData + i) = (i < values.size()) ? values[i] : 0;
242 }232 }
243- input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t *>(pData), data_len);233+ input_tensor = Tensor(input_tensor_desc, reinterpret_cast<uint8_t*>(pData), data_len);
244 delete[] pData;234 delete[] pData;
245 return SUCCESS;235 return SUCCESS;
246}236}
247 237 
248-int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t *inputData)238+int32_t WriteDataToFile(string bin_file, uint64_t data_size, uint8_t* inputData)
249{239{
250- FILE *fp = fopen(bin_file.c_str(), "w");240+ FILE* fp = fopen(bin_file.c_str(), "w");
251 fwrite(inputData, sizeof(uint8_t), data_size, fp);241 fwrite(inputData, sizeof(uint8_t), data_size, fp);
252 fclose(fp);242 fclose(fp);
253 return SUCCESS;243 return SUCCESS;
254}244}
255 245 
256-int CreateOppInGraph(std::vector<ge::Tensor> &input, std::vector<Operator> &inputs,246+int CreateOppInGraph(std::vector<ge::Tensor>& input, std::vector<Operator>& inputs, std::vector<Operator>& outputs,
257- std::vector<Operator> &outputs, Graph &graph)247+ Graph& graph)
258{248{
259 Status ret = SUCCESS;249 Status ret = SUCCESS;
260 // StatelessUniform 算子定义250 // StatelessUniform 算子定义
@@ -295,9 +285,9 @@ int CreateOppInGraph(std::vector<ge::Tensor> &input, std::vector<Operator> &inpu
295 return SUCCESS;285 return SUCCESS;
296}286}
297 287 
298-int main(int argc, char *argv[])288+int main(int argc, char* argv[])
299{289{
300- const char *graph_name = "tc_ge_irrun_test";290+ const char* graph_name = "tc_ge_irrun_test";
301 Graph graph(graph_name);291 Graph graph(graph_name);
302 std::vector<ge::Tensor> input;292 std::vector<ge::Tensor> input;
303 293 
@@ -305,7 +295,7 @@ int main(int argc, char *argv[])
305 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};295 std::map<AscendString, AscendString> global_options = {{"ge.exec.deviceId", "0"}, {"ge.graphRunMode", "1"}};
306 Status ret = ge::GEInitialize(global_options);296 Status ret = ge::GEInitialize(global_options);
307 if (ret != SUCCESS) {297 if (ret != SUCCESS) {
308- printf("%s - INFO - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());298+ printf("%s - ERROR - [XIR]: Initialize ge using ge global options failed\n", GetTime().c_str());
309 return FAILED;299 return FAILED;
310 }300 }
311 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());301 printf("%s - INFO - [XIR]: Initialize ge using ge global options success\n", GetTime().c_str());
@@ -329,7 +319,7 @@ int main(int argc, char *argv[])
329 319 
330 std::map<AscendString, AscendString> build_options = {};320 std::map<AscendString, AscendString> build_options = {};
331 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());321 printf("%s - INFO - [XIR]: Start to create ir session using build options\n", GetTime().c_str());
332- ge::Session *session = new Session(build_options);322+ ge::Session* session = new Session(build_options);
333 323 
334 if (session == nullptr) {324 if (session == nullptr) {
335 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());325 printf("%s - ERROR - [XIR]: Create ir session using build options failed\n", GetTime().c_str());
@@ -350,7 +340,7 @@ int main(int argc, char *argv[])
350 std::vector<ge::Tensor> output;340 std::vector<ge::Tensor> output;
351 ret = session->RunGraph(graph_id, input, output);341 ret = session->RunGraph(graph_id, input, output);
352 if (ret != SUCCESS) {342 if (ret != SUCCESS) {
353- printf("%s - INFO - [XIR]: Run graph failed\n", GetTime().c_str());343+ printf("%s - ERROR - [XIR]: Run graph failed\n", GetTime().c_str());
354 delete session;344 delete session;
355 GEFinalize();345 GEFinalize();
356 return FAILED;346 return FAILED;
@@ -361,12 +351,11 @@ int main(int argc, char *argv[])
361 printf("\n========== INPUT SUMMARY ==========\n");351 printf("\n========== INPUT SUMMARY ==========\n");
362 printf("Total inputs: %zu\n", input.size());352 printf("Total inputs: %zu\n", input.size());
363 int input_num = input.size();353 int input_num = input.size();
364- const char* dtypeNames[] = {354+ const char* dtypeNames[] = {"FLOAT(0)", "FLOAT16(1)", "INT8(2)", "INT32(3)", "UINT8(4)", "",
365- "FLOAT(0)", "FLOAT16(1)", "INT8(2)", "INT32(3)", "UINT8(4)", "",355+ "INT16(6)", "UINT16(7)", "UINT32(8)", "INT64(9)", "UINT64(10)", "DOUBLE(11)",
366- "INT16(6)", "UINT16(7)", "UINT32(8)", "INT64(9)", "UINT64(10)",356+ "BOOL(12)", "", "UINT1(14)", "", "", "",
367- "DOUBLE(11)", "BOOL(12)", "", "UINT1(14)", "", "", "", "", "", "",357+ "", "", "", "", "", "",
368- "", "", "", "", "", "", "BF16(27)"358+ "", "", "", "BF16(27)"};
369- };
370 for (int i = 0; i < input_num; i++) {359 for (int i = 0; i < input_num; i++) {
371 printf("---------- Input %d ----------\n", i);360 printf("---------- Input %d ----------\n", i);
372 DataType dt = input[i].GetTensorDesc().GetDataType();361 DataType dt = input[i].GetTensorDesc().GetDataType();
@@ -387,7 +376,8 @@ int main(int argc, char *argv[])
387 printf(" shape : [");376 printf(" shape : [");
388 for (size_t d = 0; d < inDimNum; d++) {377 for (size_t d = 0; d < inDimNum; d++) {
389 printf("%ld", inShape.GetDim(d));378 printf("%ld", inShape.GetDim(d));
390- if (d + 1 < inDimNum) printf(", ");379+ if (d + 1 < inDimNum)
380+ printf(", ");
391 }381 }
392 printf("] (dims=%zu, elements=%ld)\n", inDimNum, inShapeSize);382 printf("] (dims=%zu, elements=%ld)\n", inDimNum, inShapeSize);
393 383 
@@ -398,41 +388,46 @@ int main(int argc, char *argv[])
398 printf(" data size : %u bytes (%ld elems * %u bytes/elem)\n", dataBytes, inShapeSize, elemSize);388 printf(" data size : %u bytes (%ld elems * %u bytes/elem)\n", dataBytes, inShapeSize, elemSize);
399 389 
400 // print actual values390 // print actual values
401- uint8_t *inData = input[i].GetData();391+ uint8_t* inData = input[i].GetData();
402 if (inData != nullptr && inShapeSize > 0) {392 if (inData != nullptr && inShapeSize > 0) {
403 printf(" values : ");393 printf(" values : ");
404 if (dt == ge::DT_INT64) {394 if (dt == ge::DT_INT64) {
405- int64_t *vals = (int64_t*)inData;395+ int64_t* vals = (int64_t*)inData;
406 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {396 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {
407 printf("%ld", vals[j]);397 printf("%ld", vals[j]);
408- if (j + 1 < inShapeSize && j + 1 < 16) printf(", ");398+ if (j + 1 < inShapeSize && j + 1 < 16)
399+ printf(", ");
409 }400 }
410 } else if (dt == ge::DT_DOUBLE) {401 } else if (dt == ge::DT_DOUBLE) {
411- double *vals = (double*)inData;402+ double* vals = (double*)inData;
412 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {403 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {
413 printf("%.6f", vals[j]);404 printf("%.6f", vals[j]);
414- if (j + 1 < inShapeSize && j + 1 < 16) printf(", ");405+ if (j + 1 < inShapeSize && j + 1 < 16)
406+ printf(", ");
415 }407 }
416 } else if (dt == ge::DT_FLOAT) {408 } else if (dt == ge::DT_FLOAT) {
417- float *vals = (float*)inData;409+ float* vals = (float*)inData;
418 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {410 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {
419 printf("%.6f", vals[j]);411 printf("%.6f", vals[j]);
420- if (j + 1 < inShapeSize && j + 1 < 16) printf(", ");412+ if (j + 1 < inShapeSize && j + 1 < 16)
413+ printf(", ");
421 }414 }
422 } else if (dt == ge::DT_INT32) {415 } else if (dt == ge::DT_INT32) {
423- int32_t *vals = (int32_t*)inData;416+ int32_t* vals = (int32_t*)inData;
424 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {417 for (int64_t j = 0; j < inShapeSize && j < 16; j++) {
425 printf("%d", vals[j]);418 printf("%d", vals[j]);
426- if (j + 1 < inShapeSize && j + 1 < 16) printf(", ");419+ if (j + 1 < inShapeSize && j + 1 < 16)
420+ printf(", ");
427 }421 }
428 }422 }
429- if (inShapeSize > 16) printf(" ... (%ld more)", inShapeSize - 16);423+ if (inShapeSize > 16)
424+ printf(" ... (%ld more)", inShapeSize - 16);
430 printf("\n");425 printf("\n");
431 }426 }
432 427 
433 // write to file428 // write to file
434 string input_file = "./tc_ge_irrun_test_npu_input_" + std::to_string(i) + ".bin";429 string input_file = "./tc_ge_irrun_test_npu_input_" + std::to_string(i) + ".bin";
435- WriteDataToFile((const char *)input_file.c_str(), dataBytes, inData);430+ WriteDataToFile((const char*)input_file.c_str(), dataBytes, inData);
436 printf(" saved to : %s\n", input_file.c_str());431 printf(" saved to : %s\n", input_file.c_str());
437 }432 }
438 433 
@@ -460,7 +455,8 @@ int main(int argc, char *argv[])
460 printf(" inferred shape : [");455 printf(" inferred shape : [");
461 for (size_t d = 0; d < outDimNum; d++) {456 for (size_t d = 0; d < outDimNum; d++) {
462 printf("%ld", outShape.GetDim(d));457 printf("%ld", outShape.GetDim(d));
463- if (d + 1 < outDimNum) printf(", ");458+ if (d + 1 < outDimNum)
459+ printf(", ");
464 }460 }
465 printf("] (dims=%zu, elements=%ld)\n", outDimNum, outShapeSize);461 printf("] (dims=%zu, elements=%ld)\n", outDimNum, outShapeSize);
466 462 
@@ -472,22 +468,30 @@ int main(int argc, char *argv[])
472 468 
473 // write to file469 // write to file
474 string output_file = "./tc_ge_irrun_test_npu_output_" + std::to_string(i) + ".bin";470 string output_file = "./tc_ge_irrun_test_npu_output_" + std::to_string(i) + ".bin";
475- uint8_t *output_data_i = output[i].GetData();471+ uint8_t* output_data_i = output[i].GetData();
476- WriteDataToFile((const char *)output_file.c_str(), dataBytes, output_data_i);472+ WriteDataToFile((const char*)output_file.c_str(), dataBytes, output_data_i);
477 printf(" saved to : %s\n", output_file.c_str());473 printf(" saved to : %s\n", output_file.c_str());
478 474 
479 // print values with statistics475 // print values with statistics
480 if (dt == ge::DT_FLOAT && output_data_i != nullptr && outShapeSize > 0) {476 if (dt == ge::DT_FLOAT && output_data_i != nullptr && outShapeSize > 0) {
481- float *resultData = (float*)output_data_i;477+ float* resultData = (float*)output_data_i;
482 float minVal = resultData[0], maxVal = resultData[0];478 float minVal = resultData[0], maxVal = resultData[0];
483 double sum = 0.0;479 double sum = 0.0;
484 int nanCount = 0, infCount = 0;480 int nanCount = 0, infCount = 0;
485 for (int64_t j = 0; j < outShapeSize; j++) {481 for (int64_t j = 0; j < outShapeSize; j++) {
486 float v = resultData[j];482 float v = resultData[j];
487- if (std::isnan(v)) { nanCount++; continue; }483+ if (std::isnan(v)) {
488- if (std::isinf(v)) { infCount++; continue; }484+ nanCount++;
489- if (v < minVal) minVal = v;485+ continue;
490- if (v > maxVal) maxVal = v;486+ }
487+ if (std::isinf(v)) {
488+ infCount++;
489+ continue;
490+ }
491+ if (v < minVal)
492+ minVal = v;
493+ if (v > maxVal)
494+ maxVal = v;
491 sum += v;495 sum += v;
492 }496 }
493 printf(" --- Statistics ---\n");497 printf(" --- Statistics ---\n");
@@ -496,8 +500,10 @@ int main(int argc, char *argv[])
496 printf(" mean : %.6f\n", sum / outShapeSize);500 printf(" mean : %.6f\n", sum / outShapeSize);
497 printf(" range check : all in [0, 1) ? %s\n",501 printf(" range check : all in [0, 1) ? %s\n",
498 (minVal >= 0.0f && maxVal < 1.0f && nanCount == 0) ? "YES" : "NO");502 (minVal >= 0.0f && maxVal < 1.0f && nanCount == 0) ? "YES" : "NO");
499- if (nanCount > 0) printf(" NaN count : %d\n", nanCount);503+ if (nanCount > 0)
500- if (infCount > 0) printf(" Inf count : %d\n", infCount);504+ printf(" NaN count : %d\n", nanCount);
505+ if (infCount > 0)
506+ printf(" Inf count : %d\n", infCount);
501 507 
502 printf(" --- Values ---\n");508 printf(" --- Values ---\n");
503 for (int64_t j = 0; j < outShapeSize && j < 64; j++) {509 for (int64_t j = 0; j < outShapeSize && j < 64; j++) {
@@ -508,7 +514,7 @@ int main(int argc, char *argv[])
508 }514 }
509 } else if (dt == ge::DT_FLOAT16 && output_data_i != nullptr && outShapeSize > 0) {515 } else if (dt == ge::DT_FLOAT16 && output_data_i != nullptr && outShapeSize > 0) {
510 printf(" --- Values (fp16 raw hex) ---\n");516 printf(" --- Values (fp16 raw hex) ---\n");
511- uint16_t *fp16Data = (uint16_t*)output_data_i;517+ uint16_t* fp16Data = (uint16_t*)output_data_i;
512 for (int64_t j = 0; j < outShapeSize && j < 32; j++) {518 for (int64_t j = 0; j < outShapeSize && j < 32; j++) {
513 printf(" result[%ld] = 0x%04X\n", j, fp16Data[j]);519 printf(" result[%ld] = 0x%04X\n", j, fp16Data[j]);
514 }520 }
@@ -527,7 +533,7 @@ int main(int argc, char *argv[])
527 delete session;533 delete session;
528 ret = ge::GEFinalize();534 ret = ge::GEFinalize();
529 if (ret != SUCCESS) {535 if (ret != SUCCESS) {
530- printf("%s - INFO - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());536+ printf("%s - ERROR - [XIR]: Finalize ir graph session failed\n", GetTime().c_str());
531 return FAILED;537 return FAILED;
532 }538 }
533 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());539 printf("%s - INFO - [XIR]: Finalize ir graph session success\n", GetTime().c_str());