已合并
批量算子使用elewise16B接口 #4390
ligen75创建于 8月3日
批量算子使用elewise16B接口 #4390
已合并
ligen75创建于 8月3日
已删除 :master合入到cann/ops-mathmaster
33 个文件变更+487-521
@@ -24,7 +24,6 @@
24#include "../../op_kernel/arch35/abs_complex_dag.h"24#include "../../op_kernel/arch35/abs_complex_dag.h"
25 25 
26using namespace ge;26using namespace ge;
27-using namespace AbsNs;
28 27 
29namespace optiling {28namespace optiling {
30constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_BF16 = 101;29constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_BF16 = 101;
@@ -32,19 +31,19 @@ constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_OTHER = 102;
32constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_COMPLEX = 103;31constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_COMPLEX = 103;
33constexpr uint64_t ABS_WORKSPACE_RESERVE_BYTE = 16777216;32constexpr uint64_t ABS_WORKSPACE_RESERVE_BYTE = 16777216;
34 33 
35-ge::graphStatus AbsTiling::SetTilingData()34+ge::graphStatus AbsTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
36{35{
37 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);36 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);
38 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);37 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
39 currentWorkspace[0] = ABS_WORKSPACE_RESERVE_BYTE;38 currentWorkspace[0] = ABS_WORKSPACE_RESERVE_BYTE;
40 if (this->outputDtype == ge::DT_BF16) {39 if (this->outputDtype == ge::DT_BF16) {
41 tilingContext->SetTilingKey(ABS_TILING_KEY_ELEMENTWISE_BF16);40 tilingContext->SetTilingKey(ABS_TILING_KEY_ELEMENTWISE_BF16);
42- } else if (this->inputDtype == ge::DT_COMPLEX64 || this->inputDtype == ge::DT_COMPLEX32) { // 新增complex分支41+ } else if (this->inputDtype == ge::DT_COMPLEX64 || this->inputDtype == ge::DT_COMPLEX32) {
43 tilingContext->SetTilingKey(ABS_TILING_KEY_ELEMENTWISE_COMPLEX);42 tilingContext->SetTilingKey(ABS_TILING_KEY_ELEMENTWISE_COMPLEX);
44 } else {43 } else {
45 tilingContext->SetTilingKey(ABS_TILING_KEY_ELEMENTWISE_OTHER);44 tilingContext->SetTilingKey(ABS_TILING_KEY_ELEMENTWISE_OTHER);
46 }45 }
47- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);46+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
48 return ge::GRAPH_SUCCESS;47 return ge::GRAPH_SUCCESS;
49}48}
50 49 
@@ -60,22 +59,26 @@ ge::graphStatus AbsTiling::CalcOutputDtype()
60 59 
61 if (this->inputDtype != ge::DT_COMPLEX64 && this->inputDtype != ge::DT_COMPLEX32) {60 if (this->inputDtype != ge::DT_COMPLEX64 && this->inputDtype != ge::DT_COMPLEX32) {
62 OP_CHECK_IF(this->inputDtype != this->outputDtype,61 OP_CHECK_IF(this->inputDtype != this->outputDtype,
63- OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(tilingContext->GetNodeName(), "x, y",62+ OP_LOGE_FOR_INVALID_DTYPES_WITH_REASON(
64- std::string(ge::TypeUtils::DataTypeToSerialString(this->inputDtype)) + ", " + std::string(ge::TypeUtils::DataTypeToSerialString(this->outputDtype)),63+ tilingContext->GetNodeName(), "x, y",
65- "The dtypes of x and y must be the same when the dtype of x is not COMPLEX64 or COMPLEX32"),64+ std::string(ge::TypeUtils::DataTypeToSerialString(this->inputDtype)) + ", " +
65+ std::string(ge::TypeUtils::DataTypeToSerialString(this->outputDtype)),
66+ "The dtypes of x and y must be the same when the dtype of x is not COMPLEX64 or COMPLEX32"),
66 return ge::GRAPH_FAILED);67 return ge::GRAPH_FAILED);
67 } else if (inputDtype == ge::DT_COMPLEX64) {68 } else if (inputDtype == ge::DT_COMPLEX64) {
68- OP_CHECK_IF(this->outputDtype != ge::DT_FLOAT, 69+ OP_CHECK_IF(
69- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(tilingContext->GetNodeName(), "outputDtype",70+ this->outputDtype != ge::DT_FLOAT,
70- ge::TypeUtils::DataTypeToSerialString(this->outputDtype),71+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(
71- "The dtype of outputDtype must be FLOAT when the dtype of inputDtype is COMPLEX64"),72+ tilingContext->GetNodeName(), "outputDtype", ge::TypeUtils::DataTypeToSerialString(this->outputDtype),
72- return ge::GRAPH_FAILED); 73+ "The dtype of outputDtype must be FLOAT when the dtype of inputDtype is COMPLEX64"),
74+ return ge::GRAPH_FAILED);
73 } else if (inputDtype == ge::DT_COMPLEX32) {75 } else if (inputDtype == ge::DT_COMPLEX32) {
74- OP_CHECK_IF(this->outputDtype != ge::DT_FLOAT16, 76+ OP_CHECK_IF(
75- OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(tilingContext->GetNodeName(), "outputDtype",77+ this->outputDtype != ge::DT_FLOAT16,
76- ge::TypeUtils::DataTypeToSerialString(this->outputDtype),78+ OP_LOGE_FOR_INVALID_DTYPE_WITH_REASON(
77- "The dtype of outputDtype must be FLOAT16 when the dtype of inputDtype is COMPLEX32"),79+ tilingContext->GetNodeName(), "outputDtype", ge::TypeUtils::DataTypeToSerialString(this->outputDtype),
78- return ge::GRAPH_FAILED); 80+ "The dtype of outputDtype must be FLOAT16 when the dtype of inputDtype is COMPLEX32"),
81+ return ge::GRAPH_FAILED);
79 }82 }
80 return ge::GRAPH_SUCCESS;83 return ge::GRAPH_SUCCESS;
81}84}
@@ -83,51 +86,46 @@ ge::graphStatus AbsTiling::CalcOutputDtype()
83ge::graphStatus AbsTiling::RunTiling()86ge::graphStatus AbsTiling::RunTiling()
84{87{
85 ElewiseBaseTiling elewiseBaseTiling(tilingContext);88 ElewiseBaseTiling elewiseBaseTiling(tilingContext);
86- OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED,89+ OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext, "get output dtype failed"),
87- OP_LOGE(tilingContext, "get output dtype failed"),90+ return ge::GRAPH_FAILED);
88- return ge::GRAPH_FAILED);
89 91 
92+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
atomgit-bot
atomgit-botatomgit-bot8月3日

🟡 Medium Priority

RunTiling() 中,tilingContext->GetTilingData<EleBaseTilingData16B>() 的返回值 tiling 在没有任何空指针检查的情况下,立即在下一行通过 *tiling 解引用传给 DoTiling(第 95 行 res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<half, half>::OpDag>(*tiling);)。

对比同批次改造的其他两个算子:

如果 GetTilingData 返回 nullptr(例如框架内部分配失败),abs 算子会直接在 *tiling 处触发空指针解引用崩溃(段错误 / AIC Error),而另外两个算子会以 GRAPH_FAILED 正常终止。这是一个可靠性缺陷。

改动建议
92
- auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
92
+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
93
+ OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling);
应用建议
likedislike
90 ge::graphStatus res = ge::GRAPH_FAILED;93 ge::graphStatus res = ge::GRAPH_FAILED;
91- tiling = tilingContext->GetTilingData<AbsTilingData>();
92 if (this->inputDtype == ge::DT_FLOAT16) {94 if (this->inputDtype == ge::DT_FLOAT16) {
93- res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<half, half>::OpDag>(tiling->baseTiling);95+ res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<half, half>::OpDag>(*tiling);
94 } else if (this->inputDtype == ge::DT_COMPLEX64) {96 } else if (this->inputDtype == ge::DT_COMPLEX64) {
95- res = elewiseBaseTiling.DoTiling<AbsOp::AbscomplexDag<int64_t, float>::OpDag>(tiling->baseTiling);97+ res = elewiseBaseTiling.DoTiling<AbsOp::AbscomplexDag<int64_t, float>::OpDag>(*tiling);
96 } else if (this->inputDtype == ge::DT_COMPLEX32) {98 } else if (this->inputDtype == ge::DT_COMPLEX32) {
97- res = elewiseBaseTiling.DoTiling<AbsOp::AbscomplexDag<int32_t, half>::OpDag>(tiling->baseTiling); 99+ res = elewiseBaseTiling.DoTiling<AbsOp::AbscomplexDag<int32_t, half>::OpDag>(*tiling);
98 } else if (this->inputDtype == ge::DT_FLOAT) {100 } else if (this->inputDtype == ge::DT_FLOAT) {
99- res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<float, float>::OpDag>(tiling->baseTiling);101+ res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<float, float>::OpDag>(*tiling);
100 } else if (this->inputDtype == ge::DT_BF16) {102 } else if (this->inputDtype == ge::DT_BF16) {
101- res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<bfloat16_t, float>::OpDag>(tiling->baseTiling);103+ res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<bfloat16_t, float>::OpDag>(*tiling);
102 } else if (this->inputDtype == ge::DT_INT8) {104 } else if (this->inputDtype == ge::DT_INT8) {
103- res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int8_t, int8_t>::OpDag>(tiling->baseTiling);105+ res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int8_t, int8_t>::OpDag>(*tiling);
104 } else if (this->inputDtype == ge::DT_INT16) {106 } else if (this->inputDtype == ge::DT_INT16) {
105- res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int16_t, int16_t>::OpDag>(tiling->baseTiling);107+ res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int16_t, int16_t>::OpDag>(*tiling);
106 } else if (this->inputDtype == ge::DT_INT32) {108 } else if (this->inputDtype == ge::DT_INT32) {
107- res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int32_t, int32_t>::OpDag>(tiling->baseTiling);109+ res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int32_t, int32_t>::OpDag>(*tiling);
108 } else if (this->inputDtype == ge::DT_INT64) {110 } else if (this->inputDtype == ge::DT_INT64) {
109- res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int64_t, int64_t>::OpDag>(tiling->baseTiling);111+ res = elewiseBaseTiling.DoTiling<AbsOp::AbsDag<int64_t, int64_t>::OpDag>(*tiling);
110 } else {112 } else {
111 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "inputDtype",113 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "inputDtype",
112- ge::TypeUtils::DataTypeToSerialString(this->inputDtype),114+ ge::TypeUtils::DataTypeToSerialString(this->inputDtype),
113- "FLOAT16, FLOAT, BF16, INT8, INT16, INT32, INT64, COMPLEX64, COMPLEX32");115+ "FLOAT16, FLOAT, BF16, INT8, INT16, INT32, INT64, COMPLEX64, COMPLEX32");
114 return ge::GRAPH_FAILED;116 return ge::GRAPH_FAILED;
115 }117 }
116 118 
117- OP_CHECK_IF(res == ge::GRAPH_FAILED,119+ OP_CHECK_IF(res == ge::GRAPH_FAILED, OP_LOGE(tilingContext, "DoTiling failed"), return ge::GRAPH_FAILED);
118- OP_LOGE(tilingContext, "DoTiling failed"),
119- return ge::GRAPH_FAILED);
120 120 
121- ge::graphStatus result = SetTilingData();121+ ge::graphStatus result = SetTilingData(elewiseBaseTiling);
122 return result;122 return result;
123}123}
124 124 
125-static ge::graphStatus TilingForAbs(gert::TilingContext *context)125+static ge::graphStatus TilingForAbs(gert::TilingContext* context)
126{126{
127 OP_LOGD("AbsTiling", "Enter TilingForAbs");127 OP_LOGD("AbsTiling", "Enter TilingForAbs");
128- OP_CHECK_IF(context == nullptr,128+ OP_CHECK_IF(context == nullptr, OP_LOGE(context, "Tiling context is null"), return ge::GRAPH_FAILED);
129- OP_LOGE(context, "Tiling context is null"),
130- return ge::GRAPH_FAILED);
131 129 
132 // 走新的模板tiling130 // 走新的模板tiling
133 OP_LOGD("AbsTiling", "Enter new AbsTiling");131 OP_LOGD("AbsTiling", "Enter new AbsTiling");
@@ -147,6 +145,5 @@ ge::graphStatus TilingPrepareForAbs(gert::TilingParseContext* context)
147 return ge::GRAPH_SUCCESS;145 return ge::GRAPH_SUCCESS;
148}146}
149 147 
150-IMPL_OP_OPTILING(Abs).Tiling(TilingForAbs)148+IMPL_OP_OPTILING(Abs).Tiling(TilingForAbs).TilingParse<AbsCompileInfo>(TilingPrepareForAbs);
151- .TilingParse<AbsCompileInfo>(TilingPrepareForAbs);149+} // namespace optiling
152-} // namespace optiling
@@ -17,10 +17,8 @@
17 17 
18#include "register/tilingdata_base.h"18#include "register/tilingdata_base.h"
19#include "atvoss/elewise/elewise_tiling.h"19#include "atvoss/elewise/elewise_tiling.h"
20-#include "../../op_kernel/abs_struct.h"
21 20 
22namespace optiling {21namespace optiling {
23-using namespace AbsNs;
24using namespace Ops::Base;22using namespace Ops::Base;
25 23 
26struct AbsCompileInfo {24struct AbsCompileInfo {
@@ -35,13 +33,12 @@ public:
35 33 
36protected:34protected:
37 ge::graphStatus CalcOutputDtype();35 ge::graphStatus CalcOutputDtype();
38- ge::graphStatus SetTilingData();36+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
39 37 
40private:38private:
41 gert::TilingContext* tilingContext = nullptr;39 gert::TilingContext* tilingContext = nullptr;
42 ge::DataType outputDtype = ge::DT_UNDEFINED;40 ge::DataType outputDtype = ge::DT_UNDEFINED;
43 ge::DataType inputDtype = ge::DT_UNDEFINED;41 ge::DataType inputDtype = ge::DT_UNDEFINED;
44- AbsTilingData* tiling = nullptr;
45};42};
46 43 
47} // namespace optiling44} // namespace optiling
@@ -16,60 +16,57 @@
16#include "kernel_tiling/kernel_tiling.h"16#include "kernel_tiling/kernel_tiling.h"
17#include "arch35/abs_dag.h"17#include "arch35/abs_dag.h"
18#include "arch35/abs_complex_dag.h"18#include "arch35/abs_complex_dag.h"
19-#include "atvoss/elewise/elewise_sch.h"19+#include "atvoss/elewise/elewise_sch_16b.h"
20-#include "abs_struct.h"
21 20 
22using namespace AscendC;21using namespace AscendC;
23-using namespace AbsNs;
24using namespace AbsOp;22using namespace AbsOp;
25 23 
26extern "C" __global__ __aicore__ void abs(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)24extern "C" __global__ __aicore__ void abs(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
27{25{
28 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);26 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
29- REGISTER_TILING_DEFAULT(AbsTilingData);27+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
30- GET_TILING_DATA_WITH_STRUCT(AbsTilingData, tilingData, tiling);28+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
31 29 
32- TPipe pipe;
33 if (TILING_KEY_IS(101UL)) {30 if (TILING_KEY_IS(101UL)) {
34- ElementwiseSch<0UL, AbsDag<bfloat16_t, float>::OpDag> sch(&(tilingData.baseTiling), &pipe);31+ ElementwiseSch16B<0UL, AbsDag<bfloat16_t, float>::OpDag> sch(tilingData);
35 sch.Init(x, y);32 sch.Init(x, y);
36 sch.Process();33 sch.Process();
37 } else if (TILING_KEY_IS(102UL)) {34 } else if (TILING_KEY_IS(102UL)) {
38 if constexpr (std::is_same<DTYPE_X, half>::value) {35 if constexpr (std::is_same<DTYPE_X, half>::value) {
39- ElementwiseSch<0UL, AbsDag<half, half>::OpDag> sch(&(tilingData.baseTiling), &pipe);36+ ElementwiseSch16B<0UL, AbsDag<half, half>::OpDag> sch(tilingData);
40 sch.Init(x, y);37 sch.Init(x, y);
41 sch.Process();38 sch.Process();
42 } else if constexpr (std::is_same<DTYPE_X, float>::value) {39 } else if constexpr (std::is_same<DTYPE_X, float>::value) {
43- ElementwiseSch<0UL, AbsDag<float, float>::OpDag> sch(&(tilingData.baseTiling), &pipe);40+ ElementwiseSch16B<0UL, AbsDag<float, float>::OpDag> sch(tilingData);
44 sch.Init(x, y);41 sch.Init(x, y);
45 sch.Process();42 sch.Process();
46 } else if constexpr (std::is_same<DTYPE_X, int8_t>::value) {43 } else if constexpr (std::is_same<DTYPE_X, int8_t>::value) {
47- ElementwiseSch<0UL, AbsDag<int8_t, int8_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);44+ ElementwiseSch16B<0UL, AbsDag<int8_t, int8_t>::OpDag> sch(tilingData);
48 sch.Init(x, y);45 sch.Init(x, y);
49 sch.Process();46 sch.Process();
50 } else if constexpr (std::is_same<DTYPE_X, int16_t>::value) {47 } else if constexpr (std::is_same<DTYPE_X, int16_t>::value) {
51- ElementwiseSch<0UL, AbsDag<int16_t, int16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);48+ ElementwiseSch16B<0UL, AbsDag<int16_t, int16_t>::OpDag> sch(tilingData);
52 sch.Init(x, y);49 sch.Init(x, y);
53 sch.Process();50 sch.Process();
54 } else if constexpr (std::is_same<DTYPE_X, int32_t>::value) {51 } else if constexpr (std::is_same<DTYPE_X, int32_t>::value) {
55- ElementwiseSch<0UL, AbsDag<int32_t, int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);52+ ElementwiseSch16B<0UL, AbsDag<int32_t, int32_t>::OpDag> sch(tilingData);
56 sch.Init(x, y);53 sch.Init(x, y);
57 sch.Process();54 sch.Process();
58 } else if constexpr (std::is_same<DTYPE_X, int64_t>::value) {55 } else if constexpr (std::is_same<DTYPE_X, int64_t>::value) {
59- ElementwiseSch<0UL, AbsDag<int64_t, int64_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);56+ ElementwiseSch16B<0UL, AbsDag<int64_t, int64_t>::OpDag> sch(tilingData);
60 sch.Init(x, y);57 sch.Init(x, y);
61 sch.Process();58 sch.Process();
62 }59 }
63 } else if (TILING_KEY_IS(103UL)) {60 } else if (TILING_KEY_IS(103UL)) {
64 if constexpr (std::is_same<DTYPE_X, complex64>::value) {61 if constexpr (std::is_same<DTYPE_X, complex64>::value) {
65- ElementwiseSch<0UL, AbscomplexDag<complex64, float>::OpDag> sch(&(tilingData.baseTiling), &pipe);62+ ElementwiseSch16B<0UL, AbscomplexDag<complex64, float>::OpDag> sch(tilingData);
66 sch.Init(x, y);63 sch.Init(x, y);
67 sch.Process();64 sch.Process();
68 } else if constexpr (std::is_same<DTYPE_X, complex32>::value) {65 } else if constexpr (std::is_same<DTYPE_X, complex32>::value) {
69- ElementwiseSch<0UL, AbscomplexDag<complex32, half>::OpDag> sch(&(tilingData.baseTiling), &pipe);66+ ElementwiseSch16B<0UL, AbscomplexDag<complex32, half>::OpDag> sch(tilingData);
70 sch.Init(x, y);67 sch.Init(x, y);
71 sch.Process();68 sch.Process();
72 }69 }
73 }70 }
74 return;71 return;
75-}72+}
@@ -22,16 +22,10 @@
22using namespace std;22using namespace std;
23 23 
24class AbsTilingTest : public testing::Test {24class AbsTilingTest : public testing::Test {
25- protected:25+protected:
26- static void SetUpTestCase()26+ static void SetUpTestCase() { std::cout << "AbsTilingTest SetUp" << std::endl; }
27- {
28- std::cout << "AbsTilingTest SetUp" << std::endl;
29- }
30 27 
31- static void TearDownTestCase()28+ static void TearDownTestCase() { std::cout << "AbsTilingTest TearDown" << std::endl; }
32- {
33- std::cout << "AbsTilingTest TearDown" << std::endl;
34- }
35};29};
36 30 
37TEST_F(AbsTilingTest, test_tiling_fp16_001)31TEST_F(AbsTilingTest, test_tiling_fp16_001)
@@ -39,14 +33,14 @@ TEST_F(AbsTilingTest, test_tiling_fp16_001)
39 optiling::AbsCompileInfo compileInfo = {64, 262144};33 optiling::AbsCompileInfo compileInfo = {64, 262144};
40 gert::TilingContextPara tilingContextPara("Abs",34 gert::TilingContextPara tilingContextPara("Abs",
41 {35 {
42- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},36+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},
43 },37 },
44 {38 {
45- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},39+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},
46 },40 },
47 &compileInfo);41 &compileInfo);
48 uint64_t expectTilingKey = 102;42 uint64_t expectTilingKey = 102;
49- string expectTilingData = "8192 140737488355332 2048 4 1 1 2048 2048 32768 1 ";43+ string expectTilingData = "8192 140737488355332 ";
50 std::vector<size_t> expectWorkspaces = {16777216};44 std::vector<size_t> expectWorkspaces = {16777216};
51 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);45 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
52}46}
@@ -56,14 +50,14 @@ TEST_F(AbsTilingTest, test_tiling_fp32_002)
56 optiling::AbsCompileInfo compileInfo = {64, 262144};50 optiling::AbsCompileInfo compileInfo = {64, 262144};
57 gert::TilingContextPara tilingContextPara("Abs",51 gert::TilingContextPara tilingContextPara("Abs",
58 {52 {
59- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},53+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},
60 },54 },
61 {55 {
62- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},56+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},
63 },57 },
64 &compileInfo);58 &compileInfo);
65 uint64_t expectTilingKey = 102;59 uint64_t expectTilingKey = 102;
66- string expectTilingData = "8192 70368744177672 1024 8 1 1 1024 1024 16384 1 ";60+ string expectTilingData = "8192 70368744177672 ";
67 std::vector<size_t> expectWorkspaces = {16777216};61 std::vector<size_t> expectWorkspaces = {16777216};
68 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);62 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
69}63}
@@ -73,14 +67,14 @@ TEST_F(AbsTilingTest, test_tiling_int8_003)
73 optiling::AbsCompileInfo compileInfo = {64, 262144};67 optiling::AbsCompileInfo compileInfo = {64, 262144};
74 gert::TilingContextPara tilingContextPara("Abs",68 gert::TilingContextPara tilingContextPara("Abs",
75 {69 {
76- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT8, ge::FORMAT_ND},70+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT8, ge::FORMAT_ND},
77 },71 },
78 {72 {
79- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT8, ge::FORMAT_ND},73+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT8, ge::FORMAT_ND},
80 },74 },
81 &compileInfo);75 &compileInfo);
82 uint64_t expectTilingKey = 102;76 uint64_t expectTilingKey = 102;
83- string expectTilingData = "8192 281474976710658 4096 2 1 1 4096 4096 65536 1 ";77+ string expectTilingData = "8192 281474976710658 ";
84 std::vector<size_t> expectWorkspaces = {16777216};78 std::vector<size_t> expectWorkspaces = {16777216};
85 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);79 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
86}80}
@@ -90,14 +84,14 @@ TEST_F(AbsTilingTest, test_tiling_int16_004)
90 optiling::AbsCompileInfo compileInfo = {64, 262144};84 optiling::AbsCompileInfo compileInfo = {64, 262144};
91 gert::TilingContextPara tilingContextPara("Abs",85 gert::TilingContextPara tilingContextPara("Abs",
92 {86 {
93- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT16, ge::FORMAT_ND},87+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT16, ge::FORMAT_ND},
94 },88 },
95 {89 {
96- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT16, ge::FORMAT_ND},90+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT16, ge::FORMAT_ND},
97 },91 },
98 &compileInfo);92 &compileInfo);
99 uint64_t expectTilingKey = 102;93 uint64_t expectTilingKey = 102;
100- string expectTilingData = "8192 140737488355332 2048 4 1 1 2048 2048 32768 1 ";94+ string expectTilingData = "8192 140737488355332 ";
101 std::vector<size_t> expectWorkspaces = {16777216};95 std::vector<size_t> expectWorkspaces = {16777216};
102 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);96 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
103}97}
@@ -107,14 +101,14 @@ TEST_F(AbsTilingTest, test_tiling_int32_005)
107 optiling::AbsCompileInfo compileInfo = {64, 262144};101 optiling::AbsCompileInfo compileInfo = {64, 262144};
108 gert::TilingContextPara tilingContextPara("Abs",102 gert::TilingContextPara tilingContextPara("Abs",
109 {103 {
110- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT32, ge::FORMAT_ND},104+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT32, ge::FORMAT_ND},
111 },105 },
112 {106 {
113- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT32, ge::FORMAT_ND},107+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT32, ge::FORMAT_ND},
114 },108 },
115 &compileInfo);109 &compileInfo);
116 uint64_t expectTilingKey = 102;110 uint64_t expectTilingKey = 102;
117- string expectTilingData = "8192 70368744177672 1024 8 1 1 1024 1024 16384 1 ";111+ string expectTilingData = "8192 70368744177672 ";
118 std::vector<size_t> expectWorkspaces = {16777216};112 std::vector<size_t> expectWorkspaces = {16777216};
119 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);113 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
120}114}
@@ -124,14 +118,14 @@ TEST_F(AbsTilingTest, test_tiling_int64_006)
124 optiling::AbsCompileInfo compileInfo = {64, 262144};118 optiling::AbsCompileInfo compileInfo = {64, 262144};
125 gert::TilingContextPara tilingContextPara("Abs",119 gert::TilingContextPara tilingContextPara("Abs",
126 {120 {
127- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT64, ge::FORMAT_ND},121+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT64, ge::FORMAT_ND},
128 },122 },
129 {123 {
130- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT64, ge::FORMAT_ND},124+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_INT64, ge::FORMAT_ND},
131 },125 },
132 &compileInfo);126 &compileInfo);
133 uint64_t expectTilingKey = 102;127 uint64_t expectTilingKey = 102;
134- string expectTilingData = "8192 35184372088848 512 16 1 1 512 512 8192 1 ";128+ string expectTilingData = "8192 35184372088848 ";
135 std::vector<size_t> expectWorkspaces = {16777216};129 std::vector<size_t> expectWorkspaces = {16777216};
136 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);130 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
137}131}
@@ -141,14 +135,14 @@ TEST_F(AbsTilingTest, test_tiling_bf16_007)
141 optiling::AbsCompileInfo compileInfo = {64, 262144};135 optiling::AbsCompileInfo compileInfo = {64, 262144};
142 gert::TilingContextPara tilingContextPara("Abs",136 gert::TilingContextPara tilingContextPara("Abs",
143 {137 {
144- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},138+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},
145 },139 },
146 {140 {
147- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},141+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},
148 },142 },
149 &compileInfo);143 &compileInfo);
150 uint64_t expectTilingKey = 101;144 uint64_t expectTilingKey = 101;
151- string expectTilingData = "8192 46729244180484 2048 4 1 1 2048 2048 10880 1 ";145+ string expectTilingData = "8192 46729244180484 ";
152 std::vector<size_t> expectWorkspaces = {16777216};146 std::vector<size_t> expectWorkspaces = {16777216};
153 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);147 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
154}148}
@@ -158,10 +152,10 @@ TEST_F(AbsTilingTest, test_tiling_failed_bf16_fp32_008)
158 optiling::AbsCompileInfo compileInfo = {64, 262144};152 optiling::AbsCompileInfo compileInfo = {64, 262144};
159 gert::TilingContextPara tilingContextPara("Abs",153 gert::TilingContextPara tilingContextPara("Abs",
160 {154 {
161- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},155+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},
162 },156 },
163 {157 {
164- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},158+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_BF16, ge::FORMAT_ND},
165 },159 },
166 &compileInfo);160 &compileInfo);
167 uint64_t expectTilingKey = 0;161 uint64_t expectTilingKey = 0;
@@ -172,13 +166,13 @@ TEST_F(AbsTilingTest, test_tiling_failed_bf16_fp32_008)
172 166 
173TEST_F(AbsTilingTest, test_tiling_failed_empty_tensor_009)167TEST_F(AbsTilingTest, test_tiling_failed_empty_tensor_009)
174{168{
175-optiling::AbsCompileInfo compileInfo = {64, 262144};169+ optiling::AbsCompileInfo compileInfo = {64, 262144};
176 gert::TilingContextPara tilingContextPara("Abs",170 gert::TilingContextPara tilingContextPara("Abs",
177 {171 {
178- {{{1, 0, 2, 64}, {1, 0, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},172+ {{{1, 0, 2, 64}, {1, 0, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},
179 },173 },
180 {174 {
181- {{{1, 0, 2, 64}, {1, 0, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},175+ {{{1, 0, 2, 64}, {1, 0, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},
182 },176 },
183 &compileInfo);177 &compileInfo);
184 uint64_t expectTilingKey = 0;178 uint64_t expectTilingKey = 0;
@@ -189,13 +183,13 @@ optiling::AbsCompileInfo compileInfo = {64, 262144};
189 183 
190TEST_F(AbsTilingTest, test_tiling_failed_empty_tensor_010)184TEST_F(AbsTilingTest, test_tiling_failed_empty_tensor_010)
191{185{
192-optiling::AbsCompileInfo compileInfo = {64, 262144};186+ optiling::AbsCompileInfo compileInfo = {64, 262144};
193 gert::TilingContextPara tilingContextPara("Abs",187 gert::TilingContextPara tilingContextPara("Abs",
194 {188 {
195- {{{1, 1, 2, 64}, {1, 1, 2, 64}}, ge::DT_DOUBLE, ge::FORMAT_ND},189+ {{{1, 1, 2, 64}, {1, 1, 2, 64}}, ge::DT_DOUBLE, ge::FORMAT_ND},
196 },190 },
197 {191 {
198- {{{1, 1, 2, 64}, {1, 1, 2, 64}}, ge::DT_DOUBLE, ge::FORMAT_ND},192+ {{{1, 1, 2, 64}, {1, 1, 2, 64}}, ge::DT_DOUBLE, ge::FORMAT_ND},
199 },193 },
200 &compileInfo);194 &compileInfo);
201 uint64_t expectTilingKey = 0;195 uint64_t expectTilingKey = 0;
@@ -209,14 +203,14 @@ TEST_F(AbsTilingTest, test_tiling_complex64_011)
209 optiling::AbsCompileInfo compileInfo = {64, 262144};203 optiling::AbsCompileInfo compileInfo = {64, 262144};
210 gert::TilingContextPara tilingContextPara("Abs",204 gert::TilingContextPara tilingContextPara("Abs",
211 {205 {
212- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_COMPLEX64, ge::FORMAT_ND},206+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_COMPLEX64, ge::FORMAT_ND},
213 },207 },
214 {208 {
215- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},209+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT, ge::FORMAT_ND},
216 },210 },
217 &compileInfo);211 &compileInfo);
218 uint64_t expectTilingKey = 103;212 uint64_t expectTilingKey = 103;
219- string expectTilingData = "8192 35184372088840 1024 8 1 1 1024 1024 8192 1 ";213+ string expectTilingData = "8192 35184372088840 ";
220 std::vector<size_t> expectWorkspaces = {16777216};214 std::vector<size_t> expectWorkspaces = {16777216};
221 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);215 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
222}216}
@@ -226,14 +220,14 @@ TEST_F(AbsTilingTest, test_tiling_complex32_012)
226 optiling::AbsCompileInfo compileInfo = {64, 262144};220 optiling::AbsCompileInfo compileInfo = {64, 262144};
227 gert::TilingContextPara tilingContextPara("Abs",221 gert::TilingContextPara tilingContextPara("Abs",
228 {222 {
229- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_COMPLEX32, ge::FORMAT_ND},223+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_COMPLEX32, ge::FORMAT_ND},
230 },224 },
231 {225 {
232- {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},226+ {{{1, 64, 2, 64}, {1, 64, 2, 64}}, ge::DT_FLOAT16, ge::FORMAT_ND},
233 },227 },
234 &compileInfo);228 &compileInfo);
235 uint64_t expectTilingKey = 103;229 uint64_t expectTilingKey = 103;
236- string expectTilingData = "8192 70368744177668 2048 4 1 1 2048 2048 16384 1 ";230+ string expectTilingData = "8192 70368744177668 ";
237 std::vector<size_t> expectWorkspaces = {16777216};231 std::vector<size_t> expectWorkspaces = {16777216};
238 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);232 ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
239-}233+}
@@ -27,19 +27,16 @@
27namespace optiling {27namespace optiling {
28const size_t ASCEND_WORKSPACE = 16777216; // 16M28const size_t ASCEND_WORKSPACE = 16777216; // 16M
29 29 
30-ge::graphStatus CeilTiling::SetTilingData()30+ge::graphStatus CeilTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
31{31{
32- auto rawTilingData = tilingContext->GetRawTilingData();
33- OP_CHECK_NULL_WITH_CONTEXT(tilingContext, rawTilingData);
34- 
35 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);32 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);
36 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);33 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
37 currentWorkspace[0] = ASCEND_WORKSPACE;34 currentWorkspace[0] = ASCEND_WORKSPACE;
38 35 
39- const uint64_t tilingKey = GET_TPL_TILING_KEY(static_cast<uint64_t>(tiling->baseTiling.scheMode), dType);36+ const uint64_t tilingKey = GET_TPL_TILING_KEY(1, dType);
40 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);37 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);
41 tilingContext->SetTilingKey(tilingKey);38 tilingContext->SetTilingKey(tilingKey);
42- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);39+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
43 return ge::GRAPH_SUCCESS;40 return ge::GRAPH_SUCCESS;
44}41}
45 42 
@@ -105,16 +102,16 @@ ge::graphStatus CeilTiling::RunTiling()
105 return ge::GRAPH_FAILED);102 return ge::GRAPH_FAILED);
106 103 
107 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;104 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;
108- tiling = tilingContext->GetTilingData<CeilTilingData>();105+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
109 if (this->outputDtype == ge::DT_FLOAT16) {106 if (this->outputDtype == ge::DT_FLOAT16) {
110 dType = TPL_FP16;107 dType = TPL_FP16;
111- baseTilingResult = elewiseBaseTiling.DoTiling<CeilOp::CeilDAG<half>::OpDag>(tiling->baseTiling);108+ baseTilingResult = elewiseBaseTiling.DoTiling<CeilOp::CeilDAG<half>::OpDag>(*tiling);
112 } else if (this->outputDtype == ge::DT_BF16) {109 } else if (this->outputDtype == ge::DT_BF16) {
113 dType = TPL_BF16;110 dType = TPL_BF16;
114- baseTilingResult = elewiseBaseTiling.DoTiling<CeilOp::CeilDAG<bfloat16_t>::OpDag>(tiling->baseTiling);111+ baseTilingResult = elewiseBaseTiling.DoTiling<CeilOp::CeilDAG<bfloat16_t>::OpDag>(*tiling);
115 } else if (this->outputDtype == ge::DT_FLOAT) {112 } else if (this->outputDtype == ge::DT_FLOAT) {
116 dType = TPL_FP32;113 dType = TPL_FP32;
117- baseTilingResult = elewiseBaseTiling.DoTiling<CeilOp::CeilDAG<float>::OpDag>(tiling->baseTiling);114+ baseTilingResult = elewiseBaseTiling.DoTiling<CeilOp::CeilDAG<float>::OpDag>(*tiling);
118 } else {115 } else {
119 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",116 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",
120 ge::TypeUtils::DataTypeToSerialString(this->outputDtype), "FLOAT16, BF16, FLOAT");117 ge::TypeUtils::DataTypeToSerialString(this->outputDtype), "FLOAT16, BF16, FLOAT");
@@ -123,7 +120,7 @@ ge::graphStatus CeilTiling::RunTiling()
123 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),120 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),
124 return ge::GRAPH_FAILED);121 return ge::GRAPH_FAILED);
125 122 
126- return SetTilingData();123+ return SetTilingData(elewiseBaseTiling);
127}124}
128 125 
129static ge::graphStatus Tiling4Ceil(gert::TilingContext* tilingContextGen)126static ge::graphStatus Tiling4Ceil(gert::TilingContext* tilingContextGen)
@@ -17,10 +17,8 @@
17 17 
18#include "register/tilingdata_base.h"18#include "register/tilingdata_base.h"
19#include "atvoss/elewise/elewise_tiling.h"19#include "atvoss/elewise/elewise_tiling.h"
20-#include "../../op_kernel/arch35/ceil_tiling_struct.h"
21 20 
22namespace optiling {21namespace optiling {
23-using namespace CeilNs;
24using namespace Ops::Base;22using namespace Ops::Base;
25 23 
26struct CeilCompileInfo {24struct CeilCompileInfo {
@@ -32,13 +30,12 @@ class CeilTiling {
32public:30public:
33 explicit CeilTiling(gert::TilingContext* context) : tilingContext(context) {};31 explicit CeilTiling(gert::TilingContext* context) : tilingContext(context) {};
34 ge::graphStatus RunTiling();32 ge::graphStatus RunTiling();
35- CeilTilingData* tiling = nullptr;
36 33 
37protected:34protected:
38 ge::graphStatus CalcOutputDtype();35 ge::graphStatus CalcOutputDtype();
39 ge::graphStatus CalcInputDtype();36 ge::graphStatus CalcInputDtype();
40 ge::graphStatus CheckShape() const;37 ge::graphStatus CheckShape() const;
41- ge::graphStatus SetTilingData();38+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
42 39 
43private:40private:
44 gert::TilingContext* tilingContext = nullptr;41 gert::TilingContext* tilingContext = nullptr;
@@ -16,32 +16,29 @@
16#include "kernel_tiling/kernel_tiling.h"16#include "kernel_tiling/kernel_tiling.h"
17#include "arch35/ceil_dag.h"17#include "arch35/ceil_dag.h"
18#include "arch35/ceil_struct.h"18#include "arch35/ceil_struct.h"
19-#include "atvoss/elewise/elewise_sch.h"19+#include "atvoss/elewise/elewise_sch_16b.h"
20-#include "arch35/ceil_tiling_struct.h"
21 20 
22using namespace AscendC;21using namespace AscendC;
23using namespace Ops::Base;22using namespace Ops::Base;
24-using namespace CeilNs;
25using namespace CeilOp;23using namespace CeilOp;
26template <uint64_t schMode, uint64_t dType>24template <uint64_t schMode, uint64_t dType>
27__global__ __aicore__ void ceil(GM_ADDR input_x, GM_ADDR output_y, GM_ADDR workspace, GM_ADDR tiling)25__global__ __aicore__ void ceil(GM_ADDR input_x, GM_ADDR output_y, GM_ADDR workspace, GM_ADDR tiling)
28{26{
29 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);27 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
30- REGISTER_TILING_DEFAULT(CeilTilingData);28+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
31- GET_TILING_DATA_WITH_STRUCT(CeilTilingData, tilingData, tiling);29+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
32- TPipe pipe;
33 if constexpr (dType == TPL_FP16) {30 if constexpr (dType == TPL_FP16) {
34- ElementwiseSch<schMode, CeilDAG<half>::OpDag> sch(&(tilingData.baseTiling), &pipe);31+ ElementwiseSch16B<schMode, CeilDAG<half>::OpDag> sch(tilingData);
35 sch.Init(input_x, output_y);32 sch.Init(input_x, output_y);
36 sch.Process();33 sch.Process();
37 } else if constexpr (dType == TPL_BF16) {34 } else if constexpr (dType == TPL_BF16) {
38- ElementwiseSch<schMode, CeilDAG<bfloat16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);35+ ElementwiseSch16B<schMode, CeilDAG<bfloat16_t>::OpDag> sch(tilingData);
39 sch.Init(input_x, output_y);36 sch.Init(input_x, output_y);
40 sch.Process();37 sch.Process();
41 } else if constexpr (dType == TPL_FP32) {38 } else if constexpr (dType == TPL_FP32) {
42- ElementwiseSch<schMode, CeilDAG<float>::OpDag> sch(&(tilingData.baseTiling), &pipe);39+ ElementwiseSch16B<schMode, CeilDAG<float>::OpDag> sch(tilingData);
43 sch.Init(input_x, output_y);40 sch.Init(input_x, output_y);
44 sch.Process();41 sch.Process();
45 }42 }
46 return;43 return;
47-}44+}
@@ -26,7 +26,7 @@ const int64_t ASCEND_WORKSPACE = 16777216; // 16M
26const int64_t ASCEND_API_BUFFER = 122880; // 120K26const int64_t ASCEND_API_BUFFER = 122880; // 120K
27const int64_t DCACHE_SIZE = 32768;27const int64_t DCACHE_SIZE = 32768;
28 28 
29-ge::graphStatus CosTiling::SetTilingData()29+ge::graphStatus CosTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
30{30{
31 OP_LOGD(tilingContext->GetNodeName(), "CosTiling SetTilingData enter.");31 OP_LOGD(tilingContext->GetNodeName(), "CosTiling SetTilingData enter.");
32 32 
@@ -34,10 +34,10 @@ ge::graphStatus CosTiling::SetTilingData()
34 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);34 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
35 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);35 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);
36 36 
37- const uint64_t tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, dType);37+ const uint64_t tilingKey = GET_TPL_TILING_KEY(1, dType);
38 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);38 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);
39 tilingContext->SetTilingKey(tilingKey);39 tilingContext->SetTilingKey(tilingKey);
40- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);40+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
41 41 
42 uint64_t ubSize = 0;42 uint64_t ubSize = 0;
43 auto platformInfo = tilingContext->GetPlatformInfo();43 auto platformInfo = tilingContext->GetPlatformInfo();
@@ -119,20 +119,19 @@ ge::graphStatus CosTiling::RunTiling()
119 OP_CHECK_IF(CheckShape() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),119 OP_CHECK_IF(CheckShape() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),
120 return ge::GRAPH_FAILED);120 return ge::GRAPH_FAILED);
121 121 
122- tiling = tilingContext->GetTilingData<CosTilingData>();122+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
atomgit-bot
atomgit-botatomgit-bot8月3日

🟠 High Priority

变更前,CosTiling::RunTiling() 在调用 GetTilingData 后通过 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling) 检查返回值。变更后该检查被完全移除,紧接着第 127 行直接对 tiling 解引用(*tiling)传给 DoTiling。若 GetTilingData 返回 nullptr,将触发空指针解引用崩溃。触发条件:GetTilingData 返回 nullptr;失败模式:空指针解引用导致进程崩溃。

改动建议
122
- auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
122
+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
123
+ OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling);
应用建议
likedislike
123- OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling);
124 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;123 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;
125 if (this->outputDtype == ge::DT_FLOAT16) {124 if (this->outputDtype == ge::DT_FLOAT16) {
126 dType = TPL_FP16;125 dType = TPL_FP16;
127 baseTilingResult = elewiseBaseTiling.DoTiling<CosOp::CosDAG<Ops::Base::half>::OpDag>(126 baseTilingResult = elewiseBaseTiling.DoTiling<CosOp::CosDAG<Ops::Base::half>::OpDag>(
128- tiling->baseTiling, ASCEND_API_BUFFER + DCACHE_SIZE);127+ *tiling, ASCEND_API_BUFFER + DCACHE_SIZE);
129 } else if (this->outputDtype == ge::DT_BF16) {128 } else if (this->outputDtype == ge::DT_BF16) {
130 dType = TPL_BF16;129 dType = TPL_BF16;
131 baseTilingResult = elewiseBaseTiling.DoTiling<CosOp::CosDAG<Ops::Base::bfloat16_t>::OpDag>(130 baseTilingResult = elewiseBaseTiling.DoTiling<CosOp::CosDAG<Ops::Base::bfloat16_t>::OpDag>(
132- tiling->baseTiling, ASCEND_API_BUFFER + DCACHE_SIZE);131+ *tiling, ASCEND_API_BUFFER + DCACHE_SIZE);
133 } else if (this->outputDtype == ge::DT_FLOAT) {132 } else if (this->outputDtype == ge::DT_FLOAT) {
134 dType = TPL_FP32;133 dType = TPL_FP32;
135- baseTilingResult = elewiseBaseTiling.DoTiling<CosOp::CosDAG<float>::OpDag>(tiling->baseTiling,134+ baseTilingResult = elewiseBaseTiling.DoTiling<CosOp::CosDAG<float>::OpDag>(*tiling,
136 ASCEND_API_BUFFER + DCACHE_SIZE);135 ASCEND_API_BUFFER + DCACHE_SIZE);
137 } else {136 } else {
138 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",137 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",
@@ -142,7 +141,7 @@ ge::graphStatus CosTiling::RunTiling()
142 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),141 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),
143 return ge::GRAPH_FAILED);142 return ge::GRAPH_FAILED);
144 143 
145- return SetTilingData();144+ return SetTilingData(elewiseBaseTiling);
146}145}
147 146 
148static ge::graphStatus TilingForCos(gert::TilingContext* tilingContextGen)147static ge::graphStatus TilingForCos(gert::TilingContext* tilingContextGen)
@@ -16,21 +16,20 @@
16#define OPS_MATH_COS_OP_HOST_COS_ARCH35_H_16#define OPS_MATH_COS_OP_HOST_COS_ARCH35_H_
17 17 
18#include "atvoss/elewise/elewise_tiling.h"18#include "atvoss/elewise/elewise_tiling.h"
19-#include "math/cos/op_kernel/arch35/cos_tilingdata.h"
20 19 
21namespace optiling {20namespace optiling {
21+using namespace Ops::Base;
22 22 
23class CosTiling {23class CosTiling {
24public:24public:
25 explicit CosTiling(gert::TilingContext* context) : tilingContext(context) {};25 explicit CosTiling(gert::TilingContext* context) : tilingContext(context) {};
26 ge::graphStatus RunTiling();26 ge::graphStatus RunTiling();
27- CosTilingData* tiling = nullptr;
28 27 
29protected:28protected:
30 ge::graphStatus CalcOutputDtype();29 ge::graphStatus CalcOutputDtype();
31 ge::graphStatus CalcInputDtype();30 ge::graphStatus CalcInputDtype();
32 ge::graphStatus CheckShape() const;31 ge::graphStatus CheckShape() const;
33- ge::graphStatus SetTilingData();32+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
34 33 
35private:34private:
36 gert::TilingContext* tilingContext = nullptr;35 gert::TilingContext* tilingContext = nullptr;
@@ -12,34 +12,34 @@
12 * \file cos.cpp12 * \file cos.cpp
13 * \brief z = cos(x)13 * \brief z = cos(x)
14 */14 */
15- 15+ 
16#include "kernel_operator.h"16#include "kernel_operator.h"
17#include "kernel_tiling/kernel_tiling.h"17#include "kernel_tiling/kernel_tiling.h"
18#include "arch35/cos_dag.h"18#include "arch35/cos_dag.h"
19#include "arch35/cos_struct.h"19#include "arch35/cos_struct.h"
20-#include "arch35/cos_tilingdata.h"20+#include "atvoss/elewise/elewise_sch_16b.h"
21-#include "atvoss/elewise/elewise_sch.h"
22#include "atvoss/util/dfx.h"21#include "atvoss/util/dfx.h"
23 22 
24using namespace AscendC;23using namespace AscendC;
24+using namespace Ops::Base;
25template <uint64_t schMode, uint64_t dType>25template <uint64_t schMode, uint64_t dType>
26-__global__ __aicore__ void cos(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) {26+__global__ __aicore__ void cos(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
27+{
27 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);28 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
28- REGISTER_TILING_DEFAULT(CosTilingData);29+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
29- GET_TILING_DATA_WITH_STRUCT(CosTilingData, tilingData, tiling);30+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
30- TPipe pipe;31+ if constexpr (dType == TPL_FP16) {
31- if constexpr(dType == TPL_FP16) {32+ Ops::Base::ElementwiseSch16B<schMode, CosOp::CosDAG<half>::OpDag> sch(tilingData);
32- Ops::Base::ElementwiseSch<schMode, CosOp::CosDAG<half>::OpDag> sch(&(tilingData.baseTiling), &pipe);
33 sch.Init(x, y);33 sch.Init(x, y);
34 sch.Process();34 sch.Process();
35- } else if constexpr(dType == TPL_BF16) {35+ } else if constexpr (dType == TPL_BF16) {
36- Ops::Base::ElementwiseSch<schMode, CosOp::CosDAG<bfloat16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);36+ Ops::Base::ElementwiseSch16B<schMode, CosOp::CosDAG<bfloat16_t>::OpDag> sch(tilingData);
37 sch.Init(x, y);37 sch.Init(x, y);
38 sch.Process();38 sch.Process();
39- } else if constexpr(dType == TPL_FP32) {39+ } else if constexpr (dType == TPL_FP32) {
40- Ops::Base::ElementwiseSch<schMode, CosOp::CosDAG<float>::OpDag> sch(&(tilingData.baseTiling), &pipe);40+ Ops::Base::ElementwiseSch16B<schMode, CosOp::CosDAG<float>::OpDag> sch(tilingData);
41 sch.Init(x, y);41 sch.Init(x, y);
42 sch.Process();42 sch.Process();
43 }43 }
44 return;44 return;
45-}45+}
@@ -13,8 +13,10 @@
13 * \brief13 * \brief
14 */14 */
15#include "log1p_tiling_arch35.h"15#include "log1p_tiling_arch35.h"
16+#include <graph/utils/type_utils.h>
16#include "log/log.h"17#include "log/log.h"
17#include "register/op_impl_registry.h"18#include "register/op_impl_registry.h"
19+#include "tiling/platform/platform_ascendc.h"
18#include "op_host/tiling_base_util.h"20#include "op_host/tiling_base_util.h"
19#include "atvoss/elewise/elewise_tiling.h"21#include "atvoss/elewise/elewise_tiling.h"
20#include "math/log1p/op_kernel/arch35/log1p_dag.h"22#include "math/log1p/op_kernel/arch35/log1p_dag.h"
@@ -24,7 +26,7 @@ using namespace Ops::Base;
24namespace optiling {26namespace optiling {
25const int64_t ASCEND_WORKSPACE = static_cast<int64_t>(16) * 1024 * 1024;27const int64_t ASCEND_WORKSPACE = static_cast<int64_t>(16) * 1024 * 1024;
26 28 
27-ge::graphStatus Log1pTiling::SetTilingData()29+ge::graphStatus Log1pTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
28{30{
29 OP_LOGD(tilingContext->GetNodeName(), "Log1pTiling SetTilingData enter.");31 OP_LOGD(tilingContext->GetNodeName(), "Log1pTiling SetTilingData enter.");
30 32 
@@ -32,10 +34,10 @@ ge::graphStatus Log1pTiling::SetTilingData()
32 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);34 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
33 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);35 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);
34 36 
35- const uint64_t tilingKey = GET_TPL_TILING_KEY(static_cast<uint64_t>(tiling->baseTiling.scheMode), dType);37+ const uint64_t tilingKey = GET_TPL_TILING_KEY(1, dType);
36 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);38 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);
37 tilingContext->SetTilingKey(tilingKey);39 tilingContext->SetTilingKey(tilingKey);
38- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);40+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
39 return ge::GRAPH_SUCCESS;41 return ge::GRAPH_SUCCESS;
40}42}
41 43 
@@ -99,19 +101,18 @@ ge::graphStatus Log1pTiling::RunTiling()
99 return ge::GRAPH_FAILED);101 return ge::GRAPH_FAILED);
100 OP_CHECK_IF(CheckShape() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),102 OP_CHECK_IF(CheckShape() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),
101 return ge::GRAPH_FAILED);103 return ge::GRAPH_FAILED);
102- tiling = tilingContext->GetTilingData<Log1pNs::Log1pTilingData>();104+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
atomgit-bot
atomgit-botatomgit-bot8月3日

🟠 High Priority

变更前,Log1pTiling::RunTiling() 在调用 GetTilingData 后通过 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling) 检查返回值。变更后该检查被完全移除,紧接着第 109 行直接对 tiling 解引用(*tiling)传给 DoTiling。若 GetTilingData 返回 nullptr,将触发空指针解引用崩溃。触发条件:GetTilingData 返回 nullptr;失败模式:空指针解引用导致进程崩溃。

改动建议
104
- auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
104
+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
105
+ OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling);
应用建议
likedislike
103- OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling);
104 105 
105 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;106 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;
106 if (this->outputDtype == ge::DT_FLOAT16) {107 if (this->outputDtype == ge::DT_FLOAT16) {
107 dType = TPL_FP16;108 dType = TPL_FP16;
108- baseTilingResult = elewiseBaseTiling.DoTiling<Log1pOp::Log1pDAG<half>::OpDag>(tiling->baseTiling);109+ baseTilingResult = elewiseBaseTiling.DoTiling<Log1pOp::Log1pDAG<half>::OpDag>(*tiling);
109 } else if (this->outputDtype == ge::DT_BF16) {110 } else if (this->outputDtype == ge::DT_BF16) {
110 dType = TPL_BF16;111 dType = TPL_BF16;
111- baseTilingResult = elewiseBaseTiling.DoTiling<Log1pOp::Log1pDAG<bfloat16_t>::OpDag>(tiling->baseTiling);112+ baseTilingResult = elewiseBaseTiling.DoTiling<Log1pOp::Log1pDAG<bfloat16_t>::OpDag>(*tiling);
112 } else if (this->outputDtype == ge::DT_FLOAT) {113 } else if (this->outputDtype == ge::DT_FLOAT) {
113 dType = TPL_FP32;114 dType = TPL_FP32;
114- baseTilingResult = elewiseBaseTiling.DoTiling<Log1pOp::Log1pDAG<float>::OpDag>(tiling->baseTiling);115+ baseTilingResult = elewiseBaseTiling.DoTiling<Log1pOp::Log1pDAG<float>::OpDag>(*tiling);
115 } else {116 } else {
116 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",117 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",
117 ge::TypeUtils::DataTypeToSerialString(this->outputDtype), "FLOAT16, BF16, FLOAT");118 ge::TypeUtils::DataTypeToSerialString(this->outputDtype), "FLOAT16, BF16, FLOAT");
@@ -120,7 +121,7 @@ ge::graphStatus Log1pTiling::RunTiling()
120 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),121 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),
121 return ge::GRAPH_FAILED);122 return ge::GRAPH_FAILED);
122 123 
123- return SetTilingData();124+ return SetTilingData(elewiseBaseTiling);
124}125}
125 126 
126static ge::graphStatus TilingPrepareForLog1p(gert::TilingParseContext* context)127static ge::graphStatus TilingPrepareForLog1p(gert::TilingParseContext* context)
@@ -15,10 +15,11 @@
15#ifndef OPS_OP_TILING_RUNTIME_LOG1P_TILING_H15#ifndef OPS_OP_TILING_RUNTIME_LOG1P_TILING_H
16#define OPS_OP_TILING_RUNTIME_LOG1P_TILING_H16#define OPS_OP_TILING_RUNTIME_LOG1P_TILING_H
17 17 
18-#include "op_host/tiling_base_class.h"18+#include "register/tilingdata_base.h"
19-#include "../../op_kernel/arch35/log1p_tiling_struct.h"19+#include "atvoss/elewise/elewise_tiling.h"
20 20 
21namespace optiling {21namespace optiling {
22+using namespace Ops::Base;
22 23 
23struct Log1pCompileInfo {24struct Log1pCompileInfo {
24 uint64_t coreNum = 0;25 uint64_t coreNum = 0;
@@ -29,13 +30,12 @@ class Log1pTiling {
29public:30public:
30 explicit Log1pTiling(gert::TilingContext* context) : tilingContext(context) {};31 explicit Log1pTiling(gert::TilingContext* context) : tilingContext(context) {};
31 ge::graphStatus RunTiling();32 ge::graphStatus RunTiling();
32- Log1pNs::Log1pTilingData* tiling = nullptr;
33 33 
34protected:34protected:
35 ge::graphStatus CalcOutputDtype();35 ge::graphStatus CalcOutputDtype();
36 ge::graphStatus CalcInputDtype();36 ge::graphStatus CalcInputDtype();
37 ge::graphStatus CheckShape() const;37 ge::graphStatus CheckShape() const;
38- ge::graphStatus SetTilingData();38+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
39 39 
40private:40private:
41 gert::TilingContext* tilingContext = nullptr;41 gert::TilingContext* tilingContext = nullptr;
@@ -16,31 +16,31 @@
16#include "kernel_tiling/kernel_tiling.h"16#include "kernel_tiling/kernel_tiling.h"
17#include "arch35/log1p_dag.h"17#include "arch35/log1p_dag.h"
18#include "arch35/log1p_struct.h"18#include "arch35/log1p_struct.h"
19-#include "atvoss/elewise/elewise_sch.h"19+#include "atvoss/elewise/elewise_sch_16b.h"
20#include "atvoss/util/dfx.h"20#include "atvoss/util/dfx.h"
21-#include "arch35/log1p_tiling_struct.h"
22 21 
23using namespace AscendC;22using namespace AscendC;
23+using namespace Ops::Base;
24using namespace Log1pOp;24using namespace Log1pOp;
25template <uint64_t schMode, uint64_t dType>25template <uint64_t schMode, uint64_t dType>
26-__global__ __aicore__ void log1p(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) {26+__global__ __aicore__ void log1p(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
27+{
27 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);28 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
28- REGISTER_TILING_DEFAULT(Log1pNs::Log1pTilingData);29+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
29- GET_TILING_DATA_WITH_STRUCT(Log1pNs::Log1pTilingData, tilingData, tiling);30+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
30- TPipe pipe;31+ if constexpr (dType == TPL_FP16) {
31- if constexpr(dType == TPL_FP16) {32+ ElementwiseSch16B<schMode, Log1pDAG<half>::OpDag> sch(tilingData);
32- ElementwiseSch<schMode, Log1pDAG<half>::OpDag> sch(&(tilingData.baseTiling), &pipe);
33 sch.Init(x, y);33 sch.Init(x, y);
34 sch.Process();34 sch.Process();
35 } else if constexpr (dType == TPL_BF16) {35 } else if constexpr (dType == TPL_BF16) {
36- ElementwiseSch<schMode, Log1pDAG<bfloat16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);36+ ElementwiseSch16B<schMode, Log1pDAG<bfloat16_t>::OpDag> sch(tilingData);
37 sch.Init(x, y);37 sch.Init(x, y);
38 sch.Process();38 sch.Process();
39 } else if constexpr (dType == TPL_FP32) {39 } else if constexpr (dType == TPL_FP32) {
40- ElementwiseSch<schMode, Log1pDAG<float>::OpDag> sch(&(tilingData.baseTiling), &pipe);40+ ElementwiseSch16B<schMode, Log1pDAG<float>::OpDag> sch(tilingData);
41 sch.Init(x, y);41 sch.Init(x, y);
42 sch.Process();42 sch.Process();
43 }43 }
44- 44+ 
45 return;45 return;
46-}46+}
@@ -27,27 +27,26 @@ const uint64_t LOGICAL_NOT_SYS_WORKSPACE = 16777216; // 16M
27 27 
28class LogicalNotTiling {28class LogicalNotTiling {
29public:29public:
30- explicit LogicalNotTiling(gert::TilingContext *context) : tilingContext(context){};30+ explicit LogicalNotTiling(gert::TilingContext* context) : tilingContext(context) {};
31 ge::graphStatus RunTiling();31 ge::graphStatus RunTiling();
32+ 
32private:33private:
33- gert::TilingContext *tilingContext;34+ gert::TilingContext* tilingContext;
34};35};
35 36 
36ge::graphStatus LogicalNotTiling::RunTiling()37ge::graphStatus LogicalNotTiling::RunTiling()
37{38{
38-OP_CHECK_IF(tilingContext == nullptr,39+ OP_CHECK_IF(tilingContext == nullptr, OP_LOGE("CheckContextValid", "tilingContext is nullptr"),
39- OP_LOGE("CheckContextValid", "tilingContext is nullptr"),40+ return ge::GRAPH_FAILED);
40- return ge::GRAPH_FAILED);
41 41 
42- auto tiling = tilingContext->GetTilingData<EleBaseTilingDataV2>();42+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
43-OP_CHECK_IF((tiling == nullptr),43+ OP_CHECK_IF((tiling == nullptr),
44- OP_LOGE_FOR_INVALID_VALUE(tilingContext->GetNodeName(), "tiling_data", "nullptr", "not nullptr"),44+ OP_LOGE_FOR_INVALID_VALUE(tilingContext->GetNodeName(), "tiling_data", "nullptr", "not nullptr"),
45- return ge::GRAPH_FAILED);45+ return ge::GRAPH_FAILED);
46 46 
47 ElewiseBaseTiling elewiseBaseTiling(tilingContext);47 ElewiseBaseTiling elewiseBaseTiling(tilingContext);
48-OP_CHECK_IF(elewiseBaseTiling.DoTiling<LogicalNotOp::LogicalNotDag<int8_t>::OpDag>(*tiling) != ge::GRAPH_SUCCESS,48+ OP_CHECK_IF(elewiseBaseTiling.DoTiling<LogicalNotOp::LogicalNotDag<int8_t>::OpDag>(*tiling) != ge::GRAPH_SUCCESS,
49- OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling DoTiling failed"),49+ OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling DoTiling failed"), return ge::GRAPH_FAILED);
50- return ge::GRAPH_FAILED);
51 50 
52 // set workspace/tilingkey/blockdim51 // set workspace/tilingkey/blockdim
53 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);52 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);
@@ -55,12 +54,12 @@ OP_CHECK_IF(elewiseBaseTiling.DoTiling<LogicalNotOp::LogicalNotDag<int8_t>::OpDa
55 currentWorkspace[0] = LOGICAL_NOT_SYS_WORKSPACE;54 currentWorkspace[0] = LOGICAL_NOT_SYS_WORKSPACE;
56 55 
57 tilingContext->SetTilingKey(101UL);56 tilingContext->SetTilingKey(101UL);
58- tilingContext->SetBlockDim(tiling->blockNum);57+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
59 58 
60 return ge::GRAPH_SUCCESS;59 return ge::GRAPH_SUCCESS;
61}60}
62 61 
63-static ge::graphStatus Tiling4LogicalNot(gert::TilingContext *context)62+static ge::graphStatus Tiling4LogicalNot(gert::TilingContext* context)
64{63{
65 auto compileInfo = context->GetCompileInfo<ElewiseCompileInfo>();64 auto compileInfo = context->GetCompileInfo<ElewiseCompileInfo>();
66 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);65 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfo);
@@ -69,13 +68,12 @@ static ge::graphStatus Tiling4LogicalNot(gert::TilingContext *context)
69 return logicalNotTiling.RunTiling();68 return logicalNotTiling.RunTiling();
70}69}
71 70 
72- 71+ge::graphStatus TilingPrepareForLogicalNot(gert::TilingParseContext* context)
73-ge::graphStatus TilingPrepareForLogicalNot(gert::TilingParseContext *context)
74{72{
75 OP_LOGD("ElewiseTiling", "Enter TilingPrepareForLogicalNot.");73 OP_LOGD("ElewiseTiling", "Enter TilingPrepareForLogicalNot.");
76 auto compileInfoPtr = context->GetCompiledInfo<ElewiseCompileInfo>();74 auto compileInfoPtr = context->GetCompiledInfo<ElewiseCompileInfo>();
77 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfoPtr);75 OP_CHECK_NULL_WITH_CONTEXT(context, compileInfoPtr);
78- fe::PlatFormInfos *platformInfoPtr = context->GetPlatformInfo();76+ fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo();
79 OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr);77 OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr);
80 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);78 auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
81 compileInfoPtr->coreNum = ascendcPlatform.GetCoreNumAiv();79 compileInfoPtr->coreNum = ascendcPlatform.GetCoreNumAiv();
@@ -83,6 +81,5 @@ ge::graphStatus TilingPrepareForLogicalNot(gert::TilingParseContext *context)
83 return ge::GRAPH_SUCCESS;81 return ge::GRAPH_SUCCESS;
84}82}
85 83 
86-IMPL_OP_OPTILING(LogicalNot).Tiling(Tiling4LogicalNot)84+IMPL_OP_OPTILING(LogicalNot).Tiling(Tiling4LogicalNot).TilingParse<ElewiseCompileInfo>(TilingPrepareForLogicalNot);
87- .TilingParse<ElewiseCompileInfo>(TilingPrepareForLogicalNot);85+} // namespace optiling
88-} // namespace optiling
@@ -15,7 +15,7 @@
15#include "kernel_operator.h"15#include "kernel_operator.h"
16#include "kernel_tiling/kernel_tiling.h"16#include "kernel_tiling/kernel_tiling.h"
17#include "arch35/logical_not_dag.h"17#include "arch35/logical_not_dag.h"
18-#include "atvoss/elewise/elewise_sch.h"18+#include "atvoss/elewise/elewise_sch_16b.h"
19#include "atvoss/util/dfx.h"19#include "atvoss/util/dfx.h"
20 20 
21using namespace AscendC;21using namespace AscendC;
@@ -23,8 +23,8 @@ using namespace Ops::Base;
23 23 
24#define LOGICAL_NOT_DEFAULT_BOOL_TILING_KEY 101UL24#define LOGICAL_NOT_DEFAULT_BOOL_TILING_KEY 101UL
25 25 
26-__global__ __aicore__ void logical_not(GM_ADDR x, GM_ADDR y, GM_ADDR workspace,26+__global__ __aicore__ void logical_not(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
27- GM_ADDR tiling) {27+{
28 if (workspace == nullptr) {28 if (workspace == nullptr) {
29 return;29 return;
30 }30 }
@@ -33,14 +33,13 @@ __global__ __aicore__ void logical_not(GM_ADDR x, GM_ADDR y, GM_ADDR workspace,
33 if (userWS == nullptr) {33 if (userWS == nullptr) {
34 return;34 return;
35 }35 }
36- REGISTER_TILING_DEFAULT(EleBaseTilingDataV2);36+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
37- GET_TILING_DATA_WITH_STRUCT(EleBaseTilingDataV2, tilingData, tiling);37+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
38 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);38 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
39- TPipe pipe;
40 if (TILING_KEY_IS(LOGICAL_NOT_DEFAULT_BOOL_TILING_KEY)) {39 if (TILING_KEY_IS(LOGICAL_NOT_DEFAULT_BOOL_TILING_KEY)) {
41- ElementwiseSch<0UL, LogicalNotOp::LogicalNotDag<int8_t>::OpDag> sch(&tilingData, &pipe);40+ ElementwiseSch16B<0UL, LogicalNotOp::LogicalNotDag<int8_t>::OpDag> sch(tilingData);
42 sch.Init(x, y);41 sch.Init(x, y);
43 sch.Process();42 sch.Process();
44 }43 }
45 return;44 return;
46-}45+}
@@ -22,7 +22,6 @@
22#include "util/math_util.h"22#include "util/math_util.h"
23#include "atvoss/elewise/elewise_tiling.h"23#include "atvoss/elewise/elewise_tiling.h"
24#include "atvoss/elewise/elewise_base_struct.h"24#include "atvoss/elewise/elewise_base_struct.h"
25-#include "math/neg/op_kernel/arch35/neg_tiling_struct.h"
26#include "math/neg/op_kernel/arch35/neg_dag.h"25#include "math/neg/op_kernel/arch35/neg_dag.h"
27#include "math/neg/op_kernel/arch35/neg_struct.h"26#include "math/neg/op_kernel/arch35/neg_struct.h"
28 27 
@@ -52,8 +51,7 @@ class NegTiling {
52public:51public:
53 explicit NegTiling(gert::TilingContext* context) : tilingContext(context) {};52 explicit NegTiling(gert::TilingContext* context) : tilingContext(context) {};
54 ge::graphStatus RunTiling();53 ge::graphStatus RunTiling();
55- ge::graphStatus SetTilingData();54+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
56- NegTilingData* tiling = nullptr;
57 55 
58protected:56protected:
59 ge::graphStatus CalcOutputDtype();57 ge::graphStatus CalcOutputDtype();
@@ -112,7 +110,7 @@ ge::graphStatus NegTiling::CheckOutputShape() const
112ge::graphStatus NegTiling::RunTiling()110ge::graphStatus NegTiling::RunTiling()
113{111{
114 ElewiseBaseTiling elewiseBaseTiling(tilingContext);112 ElewiseBaseTiling elewiseBaseTiling(tilingContext);
115- tiling = tilingContext->GetTilingData<NegTilingData>();113+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
116 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling);114 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling);
117 // 获取tiling计算所需的参数115 // 获取tiling计算所需的参数
118 ge::graphStatus status = CalcOutputDtype();116 ge::graphStatus status = CalcOutputDtype();
@@ -126,23 +124,23 @@ ge::graphStatus NegTiling::RunTiling()
126 return ge::GRAPH_FAILED);124 return ge::GRAPH_FAILED);
127 125 
128 if (this->outputDtype == ge::DT_FLOAT16) {126 if (this->outputDtype == ge::DT_FLOAT16) {
129- status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<half>::OpDag>(tiling->baseTiling);127+ status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<half>::OpDag>(*tiling);
130- tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, NEG_TPL_FP16);128+ tilingKey = GET_TPL_TILING_KEY(1, NEG_TPL_FP16);
131 } else if (this->outputDtype == ge::DT_BF16) {129 } else if (this->outputDtype == ge::DT_BF16) {
132- status = elewiseBaseTiling.DoTiling<NegDag::NegNeedCast<bfloat16_t>::OpDag>(tiling->baseTiling);130+ status = elewiseBaseTiling.DoTiling<NegDag::NegNeedCast<bfloat16_t>::OpDag>(*tiling);
133- tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, NEG_TPL_BF16);131+ tilingKey = GET_TPL_TILING_KEY(1, NEG_TPL_BF16);
134 } else if (this->outputDtype == ge::DT_FLOAT) {132 } else if (this->outputDtype == ge::DT_FLOAT) {
135- status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<float>::OpDag>(tiling->baseTiling);133+ status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<float>::OpDag>(*tiling);
136- tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, NEG_TPL_FP32);134+ tilingKey = GET_TPL_TILING_KEY(1, NEG_TPL_FP32);
137 } else if (this->outputDtype == ge::DT_INT32) {135 } else if (this->outputDtype == ge::DT_INT32) {
138- status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<int32_t>::OpDag>(tiling->baseTiling);136+ status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<int32_t>::OpDag>(*tiling);
139- tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, NEG_TPL_INT32);137+ tilingKey = GET_TPL_TILING_KEY(1, NEG_TPL_INT32);
140 } else if (this->outputDtype == ge::DT_INT8) {138 } else if (this->outputDtype == ge::DT_INT8) {
141- status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<int8_t>::OpDag>(tiling->baseTiling);139+ status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<int8_t>::OpDag>(*tiling);
142- tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, NEG_TPL_INT8);140+ tilingKey = GET_TPL_TILING_KEY(1, NEG_TPL_INT8);
143 } else if (this->outputDtype == ge::DT_INT64) {141 } else if (this->outputDtype == ge::DT_INT64) {
144- status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<int64_t>::OpDag>(tiling->baseTiling);142+ status = elewiseBaseTiling.DoTiling<NegDag::NegNoCast<int64_t>::OpDag>(*tiling);
145- tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, NEG_TPL_INT64);143+ tilingKey = GET_TPL_TILING_KEY(1, NEG_TPL_INT64);
146 } else {144 } else {
147 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",145 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",
148 ge::TypeUtils::DataTypeToSerialString(this->outputDtype),146 ge::TypeUtils::DataTypeToSerialString(this->outputDtype),
@@ -153,13 +151,13 @@ ge::graphStatus NegTiling::RunTiling()
153 OP_CHECK_IF(status != ge::GRAPH_SUCCESS,151 OP_CHECK_IF(status != ge::GRAPH_SUCCESS,
154 OP_LOGE(tilingContext->GetNodeName(), "ElewiseBaseTiling do tiling failed"), return ge::GRAPH_FAILED);152 OP_LOGE(tilingContext->GetNodeName(), "ElewiseBaseTiling do tiling failed"), return ge::GRAPH_FAILED);
155 153 
156- return SetTilingData();154+ return SetTilingData(elewiseBaseTiling);
157}155}
158 156 
159-ge::graphStatus NegTiling::SetTilingData()157+ge::graphStatus NegTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
160{158{
161 tilingContext->SetTilingKey(tilingKey);159 tilingContext->SetTilingKey(tilingKey);
162- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);160+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
163 auto rawTilingData = tilingContext->GetRawTilingData();161 auto rawTilingData = tilingContext->GetRawTilingData();
164 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, rawTilingData);162 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, rawTilingData);
165 size_t usrWorkspaceSize = 0;163 size_t usrWorkspaceSize = 0;
@@ -15,27 +15,26 @@
15 15 
16#include "kernel_operator.h"16#include "kernel_operator.h"
17#include "kernel_tiling/kernel_tiling.h"17#include "kernel_tiling/kernel_tiling.h"
18-#include "atvoss/elewise/elewise_sch.h"18+#include "atvoss/elewise/elewise_sch_16b.h"
19#include "arch35/neg_dag.h"19#include "arch35/neg_dag.h"
20#include "arch35/neg_struct.h"20#include "arch35/neg_struct.h"
21-#include "arch35/neg_tiling_struct.h"
22 21 
23using namespace AscendC;22using namespace AscendC;
23+using namespace Ops::Base;
24using namespace NegOp;24using namespace NegOp;
25 25 
26template <uint64_t scheMode, uint64_t dType>26template <uint64_t scheMode, uint64_t dType>
27__global__ __aicore__ void neg(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)27__global__ __aicore__ void neg(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
28{28{
29 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);29 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
30- REGISTER_TILING_DEFAULT(NegTilingData);30+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
31- GET_TILING_DATA_WITH_STRUCT(NegTilingData, tilingData, tiling);31+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
32- TPipe pipe;
33 if constexpr (IsSameType<DTYPE_X, bfloat16_t>::value) {32 if constexpr (IsSameType<DTYPE_X, bfloat16_t>::value) {
34- ElementwiseSch<scheMode, NegDag::NegNeedCast<bfloat16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);33+ ElementwiseSch16B<scheMode, NegDag::NegNeedCast<bfloat16_t>::OpDag> sch(tilingData);
35 sch.Init(x, y);34 sch.Init(x, y);
36 sch.Process();35 sch.Process();
37 } else {36 } else {
38- ElementwiseSch<scheMode, NegDag::NegNoCast<DTYPE_X>::OpDag> sch(&(tilingData.baseTiling), &pipe);37+ ElementwiseSch16B<scheMode, NegDag::NegNoCast<DTYPE_X>::OpDag> sch(tilingData);
39 sch.Init(x, y);38 sch.Init(x, y);
40 sch.Process();39 sch.Process();
41 }40 }
@@ -13,110 +13,133 @@
13#include "tiling_context_faker.h"13#include "tiling_context_faker.h"
14#include "tiling_case_executor.h"14#include "tiling_case_executor.h"
15#include "../../../../op_host/arch35/neg_tiling_arch35.h"15#include "../../../../op_host/arch35/neg_tiling_arch35.h"
16-#include "../../../../op_kernel/arch35/neg_tiling_struct.h"
17#include "atvoss/elewise/elewise_tiling.h"16#include "atvoss/elewise/elewise_tiling.h"
18 17 
19using namespace std;18using namespace std;
20using namespace ge;19using namespace ge;
21 20 
22class NegTiling : public testing::Test {21class NegTiling : public testing::Test {
23- protected:22+protected:
24- static void SetUpTestCase() {23+ static void SetUpTestCase() { std::cout << "NegTiling SetUp" << std::endl; }
25- std::cout << "NegTiling SetUp" << std::endl;
26- }
27 24 
28- static void TearDownTestCase() {25+ static void TearDownTestCase() { std::cout << "NegTiling TearDown" << std::endl; }
29- std::cout << "NegTiling TearDown" << std::endl;
30- }
31};26};
32 27 
33TEST_F(NegTiling, neg_test_tiling_float16_input)28TEST_F(NegTiling, neg_test_tiling_float16_input)
34{29{
35 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};30 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};
36 gert::TilingContextPara tilingContextPara("Neg",31 gert::TilingContextPara tilingContextPara("Neg",
37- {{{{8, 8}, {8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND},},32+ {
38- {{{{8, 8}, {8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND},},33+ {{{8, 8}, {8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND},
39- &compileInfo);34+ },
35+ {
36+ {{{8, 8}, {8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND},
37+ },
38+ &compileInfo);
40 uint64_t expectTilingKey = 3;39 uint64_t expectTilingKey = 3;
41- string expectTilingData = "64 1 32768 512 1 1 1 512 64 32768 1 ";40+ string expectTilingData = "64 140737488355329 ";
42 std::vector<size_t> expectWorkspaces = {16777216};41 std::vector<size_t> expectWorkspaces = {16777216};
43- ExecuteTestCaseForEle(tilingContextPara, ge::GRAPH_SUCCESS, true, expectTilingKey, true, expectTilingData, expectWorkspaces);42+ ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
44}43}
45 44 
46TEST_F(NegTiling, neg_test_tiling_float_input)45TEST_F(NegTiling, neg_test_tiling_float_input)
47{46{
48 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};47 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};
49 gert::TilingContextPara tilingContextPara("Neg",48 gert::TilingContextPara tilingContextPara("Neg",
50- {{{{8, 8}, {8, 8}}, ge::DT_FLOAT, ge::FORMAT_ND},},49+ {
51- {{{{8, 8}, {8, 8}}, ge::DT_FLOAT, ge::FORMAT_ND},},50+ {{{8, 8}, {8, 8}}, ge::DT_FLOAT, ge::FORMAT_ND},
52- &compileInfo);51+ },
52+ {
53+ {{{8, 8}, {8, 8}}, ge::DT_FLOAT, ge::FORMAT_ND},
54+ },
55+ &compileInfo);
53 uint64_t expectTilingKey = 7;56 uint64_t expectTilingKey = 7;
54- string expectTilingData = "64 1 16384 512 1 1 1 512 64 16384 1 ";57+ string expectTilingData = "64 70368744177665 ";
55 std::vector<size_t> expectWorkspaces = {16777216};58 std::vector<size_t> expectWorkspaces = {16777216};
56- ExecuteTestCaseForEle(tilingContextPara, ge::GRAPH_SUCCESS, true, expectTilingKey, true, expectTilingData, expectWorkspaces);59+ ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
57}60}
58 61 
59TEST_F(NegTiling, neg_test_tiling_INT32_input)62TEST_F(NegTiling, neg_test_tiling_INT32_input)
60{63{
61 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};64 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};
62 gert::TilingContextPara tilingContextPara("Neg",65 gert::TilingContextPara tilingContextPara("Neg",
63- {{{{8, 8}, {8, 8}}, ge::DT_INT32, ge::FORMAT_ND},},66+ {
64- {{{{8, 8}, {8, 8}}, ge::DT_INT32, ge::FORMAT_ND},},67+ {{{8, 8}, {8, 8}}, ge::DT_INT32, ge::FORMAT_ND},
65- &compileInfo);68+ },
69+ {
70+ {{{8, 8}, {8, 8}}, ge::DT_INT32, ge::FORMAT_ND},
71+ },
72+ &compileInfo);
66 uint64_t expectTilingKey = 11;73 uint64_t expectTilingKey = 11;
67- string expectTilingData = "64 1 16384 512 1 1 1 512 64 16384 1 ";74+ string expectTilingData = "64 70368744177665 ";
68 std::vector<size_t> expectWorkspaces = {16777216};75 std::vector<size_t> expectWorkspaces = {16777216};
69- ExecuteTestCaseForEle(tilingContextPara, ge::GRAPH_SUCCESS, true, expectTilingKey, true, expectTilingData, expectWorkspaces);76+ ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
70}77}
71 78 
72TEST_F(NegTiling, neg_test_tiling_int8_input)79TEST_F(NegTiling, neg_test_tiling_int8_input)
73{80{
74 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};81 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};
75 gert::TilingContextPara tilingContextPara("Neg",82 gert::TilingContextPara tilingContextPara("Neg",
76- {{{{8, 8}, {8, 8}}, ge::DT_INT8, ge::FORMAT_ND},},83+ {
77- {{{{8, 8}, {8, 8}}, ge::DT_INT8, ge::FORMAT_ND},},84+ {{{8, 8}, {8, 8}}, ge::DT_INT8, ge::FORMAT_ND},
78- &compileInfo);85+ },
86+ {
87+ {{{8, 8}, {8, 8}}, ge::DT_INT8, ge::FORMAT_ND},
88+ },
89+ &compileInfo);
79 uint64_t expectTilingKey = 9;90 uint64_t expectTilingKey = 9;
80- string expectTilingData = "64 1 65536 512 1 1 1 512 64 65536 1 ";91+ string expectTilingData = "64 281474976710657 ";
81 std::vector<size_t> expectWorkspaces = {16777216};92 std::vector<size_t> expectWorkspaces = {16777216};
82- ExecuteTestCaseForEle(tilingContextPara, ge::GRAPH_SUCCESS, true, expectTilingKey, true, expectTilingData, expectWorkspaces);93+ ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
83}94}
84 95 
85TEST_F(NegTiling, neg_test_tiling_int64_input)96TEST_F(NegTiling, neg_test_tiling_int64_input)
86{97{
87 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};98 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};
88 gert::TilingContextPara tilingContextPara("Neg",99 gert::TilingContextPara tilingContextPara("Neg",
89- {{{{8, 8}, {8, 8}}, ge::DT_INT64, ge::FORMAT_ND},},100+ {
90- {{{{8, 8}, {8, 8}}, ge::DT_INT64, ge::FORMAT_ND},},101+ {{{8, 8}, {8, 8}}, ge::DT_INT64, ge::FORMAT_ND},
91- &compileInfo);102+ },
103+ {
104+ {{{8, 8}, {8, 8}}, ge::DT_INT64, ge::FORMAT_ND},
105+ },
106+ &compileInfo);
92 uint64_t expectTilingKey = 13;107 uint64_t expectTilingKey = 13;
93- string expectTilingData = "64 1 8192 512 1 1 1 512 64 8192 1 ";108+ string expectTilingData = "64 35184372088833 ";
94 std::vector<size_t> expectWorkspaces = {16777216};109 std::vector<size_t> expectWorkspaces = {16777216};
95- ExecuteTestCaseForEle(tilingContextPara, ge::GRAPH_SUCCESS, true, expectTilingKey, true, expectTilingData, expectWorkspaces);110+ ExecuteTestCase(tilingContextPara, ge::GRAPH_SUCCESS, expectTilingKey, expectTilingData, expectWorkspaces);
96}111}
97 112 
98TEST_F(NegTiling, neg_test_tiling_invalid_shape)113TEST_F(NegTiling, neg_test_tiling_invalid_shape)
99{114{
100 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};115 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};
101 gert::TilingContextPara tilingContextPara("Neg",116 gert::TilingContextPara tilingContextPara("Neg",
102- {{{{8, 8}, {8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND},},117+ {
103- {{{{8}, {8}}, ge::DT_FLOAT16, ge::FORMAT_ND},},118+ {{{8, 8}, {8, 8}}, ge::DT_FLOAT16, ge::FORMAT_ND},
104- &compileInfo);119+ },
105- uint64_t expectTilingKey = 3;120+ {
106- string expectTilingData = "64 1 32768 512 1 1 1 512 64 32768 1 ";121+ {{{8}, {8}}, ge::DT_FLOAT16, ge::FORMAT_ND},
107- std::vector<size_t> expectWorkspaces = {16777216};122+ },
108- ExecuteTestCaseForEle(tilingContextPara, ge::GRAPH_FAILED, true, expectTilingKey, true, expectTilingData, expectWorkspaces);123+ &compileInfo);
124+ uint64_t expectTilingKey = 0;
125+ string expectTilingData = "";
126+ std::vector<size_t> expectWorkspaces = {0};
127+ ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED, expectTilingKey, expectTilingData, expectWorkspaces);
109}128}
110 129 
111TEST_F(NegTiling, neg_test_tiling_invalid_dtype)130TEST_F(NegTiling, neg_test_tiling_invalid_dtype)
112{131{
113 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};132 Ops::Base::ElewiseCompileInfo compileInfo = {64, 253952};
114 gert::TilingContextPara tilingContextPara("Neg",133 gert::TilingContextPara tilingContextPara("Neg",
115- {{{{8, 8}, {8, 8}}, ge::DT_COMPLEX32, ge::FORMAT_ND},},134+ {
116- {{{{8, 8}, {8, 8}}, ge::DT_COMPLEX32, ge::FORMAT_ND},},135+ {{{8, 8}, {8, 8}}, ge::DT_COMPLEX32, ge::FORMAT_ND},
117- &compileInfo);136+ },
118- uint64_t expectTilingKey = 3;137+ {
119- string expectTilingData = "64 1 32768 512 1 1 1 512 64 32768 1 ";138+ {{{8, 8}, {8, 8}}, ge::DT_COMPLEX32, ge::FORMAT_ND},
120- std::vector<size_t> expectWorkspaces = {16777216};139+ },
121- ExecuteTestCaseForEle(tilingContextPara, ge::GRAPH_FAILED, true, expectTilingKey, true, expectTilingData, expectWorkspaces);140+ &compileInfo);
122-}141+ uint64_t expectTilingKey = 0;
142+ string expectTilingData = "";
143+ std::vector<size_t> expectWorkspaces = {0};
144+ ExecuteTestCase(tilingContextPara, ge::GRAPH_FAILED, expectTilingKey, expectTilingData, expectWorkspaces);
145+}
@@ -45,39 +45,34 @@ class RoundTiling {
45public:45public:
46 explicit RoundTiling(gert::TilingContext* context) : tilingContext(context) {};46 explicit RoundTiling(gert::TilingContext* context) : tilingContext(context) {};
47 ge::graphStatus RunTiling();47 ge::graphStatus RunTiling();
48- RoundTilingData* tiling_ = nullptr;
49 48 
50protected:49protected:
51 ge::graphStatus CalcOutputDtype();50 ge::graphStatus CalcOutputDtype();
52 ge::graphStatus CalcInputDtype();51 ge::graphStatus CalcInputDtype();
53 ge::graphStatus CheckShape() const;52 ge::graphStatus CheckShape() const;
54- ge::graphStatus SetTilingData();53+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
55- ge::graphStatus DoTilingF(bool decimalsNeg, bool decimalsNan);54+ ge::graphStatus DoTilingF(bool decimalsNeg, bool decimalsNan, float decimalsVal);
56 ge::graphStatus DoTilingI(int64_t decimals);55 ge::graphStatus DoTilingI(int64_t decimals);
57 56 
58private:57private:
59 uint64_t dType = 0;58 uint64_t dType = 0;
60- uint64_t schMode = 0;
61 gert::TilingContext* tilingContext = nullptr;59 gert::TilingContext* tilingContext = nullptr;
62 ge::DataType outputDtype = ge::DT_UNDEFINED;60 ge::DataType outputDtype = ge::DT_UNDEFINED;
63 ge::DataType inputDtype = ge::DT_UNDEFINED;61 ge::DataType inputDtype = ge::DT_UNDEFINED;
64};62};
65 63 
66-ge::graphStatus RoundTiling::SetTilingData()64+ge::graphStatus RoundTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
67{65{
68 OP_LOGD(tilingContext->GetNodeName(), "RoundTiling SetTilingData enter.");66 OP_LOGD(tilingContext->GetNodeName(), "RoundTiling SetTilingData enter.");
69- auto rawTilingData = tilingContext->GetRawTilingData();
70- OP_CHECK_NULL_WITH_CONTEXT(tilingContext, rawTilingData);
71 67 
72 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);68 size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);
73 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);69 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
74 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);70 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);
75 71 
76- schMode = tiling_->baseTiling.scheMode;72+ const uint64_t tilingKey = GET_TPL_TILING_KEY(1, dType);
77- const uint64_t tilingKey = GET_TPL_TILING_KEY(schMode, dType);
78 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);73 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);
79 tilingContext->SetTilingKey(tilingKey);74 tilingContext->SetTilingKey(tilingKey);
80- tilingContext->SetBlockDim(tiling_->baseTiling.blockNum);75+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
81 return ge::GRAPH_SUCCESS;76 return ge::GRAPH_SUCCESS;
82}77}
83 78 
@@ -135,62 +130,62 @@ ge::graphStatus RoundTiling::CalcOutputDtype()
135 return ge::GRAPH_SUCCESS;130 return ge::GRAPH_SUCCESS;
136}131}
137 132 
138-ge::graphStatus RoundTiling::DoTilingF(bool decimalsNeg, bool decimalsNan)133+ge::graphStatus RoundTiling::DoTilingF(bool decimalsNeg, bool decimalsNan, float decimalsVal)
139{134{
140 ElewiseBaseTiling elewiseBaseTiling(tilingContext);135 ElewiseBaseTiling elewiseBaseTiling(tilingContext);
141 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;136 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;
142 if (this->outputDtype == ge::DT_FLOAT16) {137 if (this->outputDtype == ge::DT_FLOAT16) {
143- if (tiling_->decimals == DEFAULT_FP32_ZERO) {138+ if (decimalsVal == DEFAULT_FP32_ZERO) {
144 dType = static_cast<uint64_t>(ROUND_TPL_ZERO);139 dType = static_cast<uint64_t>(ROUND_TPL_ZERO);
145- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundZero<half>::OpDag>(tiling_->baseTiling);140+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundZero<half>::OpDag>();
146 } else if (decimalsNeg) {141 } else if (decimalsNeg) {
147 dType = static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS);142 dType = static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS);
148- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundNegativeDecimals<half>::OpDag>(143+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundNegativeDecimals<half>::OpDag>();
149- tiling_->baseTiling);144+ elewiseBaseTiling.SetScalar<float>(decimalsVal);
150 } else if (decimalsNan) {145 } else if (decimalsNan) {
151 dType = static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS);146 dType = static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS);
152- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundNan<half>::OpDag>(tiling_->baseTiling);147+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundNan<half>::OpDag>();
153 } else {148 } else {
154 dType = static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS);149 dType = static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS);
155- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundPositiveDecimals<half>::OpDag>(150+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundPositiveDecimals<half>::OpDag>();
156- tiling_->baseTiling);151+ elewiseBaseTiling.SetScalar<float>(decimalsVal);
157 }152 }
158 } else if (this->outputDtype == ge::DT_BF16) {153 } else if (this->outputDtype == ge::DT_BF16) {
159- if (tiling_->decimals == DEFAULT_FP32_ZERO) {154+ if (decimalsVal == DEFAULT_FP32_ZERO) {
160 dType = static_cast<uint64_t>(ROUND_TPL_ZERO);155 dType = static_cast<uint64_t>(ROUND_TPL_ZERO);
161- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundZero<bfloat16_t>::OpDag>(tiling_->baseTiling);156+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundZero<bfloat16_t>::OpDag>();
162 } else if (decimalsNeg) {157 } else if (decimalsNeg) {
163 dType = static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS);158 dType = static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS);
164- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundNegativeDecimals<bfloat16_t>::OpDag>(159+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundNegativeDecimals<bfloat16_t>::OpDag>();
165- tiling_->baseTiling);160+ elewiseBaseTiling.SetScalar<float>(decimalsVal);
166 } else if (decimalsNan) {161 } else if (decimalsNan) {
167 dType = static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS);162 dType = static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS);
168- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundNan<bfloat16_t>::OpDag>(tiling_->baseTiling);163+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundNan<bfloat16_t>::OpDag>();
169 } else {164 } else {
170 dType = static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS);165 dType = static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS);
171- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundPositiveDecimals<bfloat16_t>::OpDag>(166+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundPositiveDecimals<bfloat16_t>::OpDag>();
172- tiling_->baseTiling);167+ elewiseBaseTiling.SetScalar<float>(decimalsVal);
173 }168 }
174 } else if (this->outputDtype == ge::DT_FLOAT) {169 } else if (this->outputDtype == ge::DT_FLOAT) {
175- if (tiling_->decimals == DEFAULT_FP32_ZERO) {170+ if (decimalsVal == DEFAULT_FP32_ZERO) {
176 dType = static_cast<uint64_t>(ROUND_TPL_ZERO);171 dType = static_cast<uint64_t>(ROUND_TPL_ZERO);
177- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundZero<float>::OpDag>(tiling_->baseTiling);172+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundZero<float>::OpDag>();
178 } else if (decimalsNeg) {173 } else if (decimalsNeg) {
179 dType = static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS);174 dType = static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS);
180- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundNegativeDecimals<float>::OpDag>(175+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundNegativeDecimals<float>::OpDag>();
181- tiling_->baseTiling);176+ elewiseBaseTiling.SetScalar<float>(decimalsVal);
182 } else if (decimalsNan) {177 } else if (decimalsNan) {
183 dType = static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS);178 dType = static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS);
184- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundNan<float>::OpDag>(tiling_->baseTiling);179+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundNan<float>::OpDag>();
185 } else {180 } else {
186 dType = static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS);181 dType = static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS);
187- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundPositiveDecimals<float>::OpDag>(182+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundPositiveDecimals<float>::OpDag>();
188- tiling_->baseTiling);183+ elewiseBaseTiling.SetScalar<float>(decimalsVal);
189 }184 }
190 }185 }
191 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED,186 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED,
192 OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTilingF failed"), return ge::GRAPH_FAILED);187 OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTilingF failed"), return ge::GRAPH_FAILED);
193- return ge::GRAPH_SUCCESS;188+ return SetTilingData(elewiseBaseTiling);
194}189}
195 190 
196ge::graphStatus RoundTiling::DoTilingI(int64_t decimals)191ge::graphStatus RoundTiling::DoTilingI(int64_t decimals)
@@ -200,47 +195,46 @@ ge::graphStatus RoundTiling::DoTilingI(int64_t decimals)
200 195 
201 if (decimals >= DEFAULT_ZERO) {196 if (decimals >= DEFAULT_ZERO) {
202 dType = static_cast<uint64_t>(ROUND_TPL_INT32);197 dType = static_cast<uint64_t>(ROUND_TPL_INT32);
203- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundInt<int32_t>::OpDag>(tiling_->baseTiling);198+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundInt<int32_t>::OpDag>();
204 }199 }
205 if ((decimals < DEFAULT_ZERO) && (decimals > DEFAULT_NEG_NINE)) {200 if ((decimals < DEFAULT_ZERO) && (decimals > DEFAULT_NEG_NINE)) {
206- tiling_->power = powerArr[llabs(static_cast<int32_t>(decimals))];201+ int32_t power = powerArr[llabs(static_cast<int32_t>(decimals))];
207- tiling_->num = numArr[llabs(static_cast<int32_t>(decimals))];202+ int32_t num = numArr[llabs(static_cast<int32_t>(decimals))];
208 if (llabs(static_cast<int32_t>(decimals)) & 1) {203 if (llabs(static_cast<int32_t>(decimals)) & 1) {
209 dType = static_cast<uint64_t>(ROUND_TPL_INT32_NEGINF);204 dType = static_cast<uint64_t>(ROUND_TPL_INT32_NEGINF);
210- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundIntNegativeDecimalsInf<int32_t>::OpDag>(205+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundIntNegativeDecimalsInf<int32_t>::OpDag>();
211- tiling_->baseTiling);206+ elewiseBaseTiling.SetScalar<int32_t>(power);
207+ elewiseBaseTiling.SetScalar<int32_t>(num);
212 } else {208 } else {
213 dType = static_cast<uint64_t>(ROUND_TPL_INT32_NEG);209 dType = static_cast<uint64_t>(ROUND_TPL_INT32_NEG);
214- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundIntNegativeDecimals<int32_t>::OpDag>(210+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundIntNegativeDecimals<int32_t>::OpDag>();
215- tiling_->baseTiling);211+ elewiseBaseTiling.SetScalar<int32_t>(power);
212+ elewiseBaseTiling.SetScalar<int32_t>(num);
216 }213 }
217 }214 }
218 if (decimals == DEFAULT_NEG_NINE) {215 if (decimals == DEFAULT_NEG_NINE) {
219 dType = static_cast<uint64_t>(ROUND_TPL_INT32_NEG_NINE);216 dType = static_cast<uint64_t>(ROUND_TPL_INT32_NEG_NINE);
220- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundIntNegativeDecimalsNine<int32_t>::OpDag>(217+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundIntNegativeDecimalsNine<int32_t>::OpDag>();
221- tiling_->baseTiling);
222 }218 }
223 if (decimals < DEFAULT_NEG_NINE) {219 if (decimals < DEFAULT_NEG_NINE) {
224- tiling_->num = static_cast<int32_t>(DEFAULT_FP32_ZERO);220+ int32_t num = static_cast<int32_t>(DEFAULT_FP32_ZERO);
225 if (decimals < DEFAULT_NEG_MAX) {221 if (decimals < DEFAULT_NEG_MAX) {
226- tiling_->num = static_cast<int32_t>(DEFAULT_FP32_MIN);222+ num = static_cast<int32_t>(DEFAULT_FP32_MIN);
227 }223 }
228 dType = static_cast<uint64_t>(ROUND_TPL_INT32_CONST);224 dType = static_cast<uint64_t>(ROUND_TPL_INT32_CONST);
229- baseTilingResult = elewiseBaseTiling.DoTiling<RoundDag::RoundIntConst<int>::OpDag>(tiling_->baseTiling);225+ baseTilingResult = elewiseBaseTiling.DoTiling32B<RoundDag::RoundIntConst<int>::OpDag>();
226+ elewiseBaseTiling.SetScalar<int32_t>(num);
230 }227 }
231 228 
232 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED,229 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED,
233 OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTilingInt failed"), return ge::GRAPH_FAILED);230 OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTilingInt failed"), return ge::GRAPH_FAILED);
234- return ge::GRAPH_SUCCESS;231+ return SetTilingData(elewiseBaseTiling);
235}232}
236 233 
237ge::graphStatus RoundTiling::RunTiling()234ge::graphStatus RoundTiling::RunTiling()
238{235{
239 OP_LOGD(tilingContext->GetNodeName(), "RoundTiling RunTiling enter.");236 OP_LOGD(tilingContext->GetNodeName(), "RoundTiling RunTiling enter.");
240 237 
241- tiling_ = tilingContext->GetTilingData<RoundTilingData>();
242- OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tiling_);
243- 
244 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input x dtype failed"),238 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input x dtype failed"),
245 return ge::GRAPH_FAILED);239 return ge::GRAPH_FAILED);
246 OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED,240 OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED,
@@ -256,29 +250,25 @@ ge::graphStatus RoundTiling::RunTiling()
256 return ge::GRAPH_FAILED);250 return ge::GRAPH_FAILED);
257 251 
258 if (this->outputDtype == ge::DT_INT32) {252 if (this->outputDtype == ge::DT_INT32) {
259- DoTilingI(*decimalsPtr);253+ return DoTilingI(*decimalsPtr);
260 } else {254 } else {
261 bool decimalsNeg = false;255 bool decimalsNeg = false;
262 bool decimalsNan = false;256 bool decimalsNan = false;
263- 257+ float decimalsVal = DEFAULT_FP32_ZERO;
264- tiling_->decimals = DEFAULT_FP32_ZERO;
265 if (*decimalsPtr < DEFAULT_ZERO) {258 if (*decimalsPtr < DEFAULT_ZERO) {
266 decimalsNeg = true;259 decimalsNeg = true;
267 }260 }
268 if (*decimalsPtr != DEFAULT_ZERO) {261 if (*decimalsPtr != DEFAULT_ZERO) {
269 if (llabs(*decimalsPtr) > DEFAULT_THIRTY_EIGHT) {262 if (llabs(*decimalsPtr) > DEFAULT_THIRTY_EIGHT) {
270- tiling_->decimals = DEFAULT_INF;263+ decimalsVal = DEFAULT_INF;
271 decimalsNan = true;264 decimalsNan = true;
272 decimalsNeg = false;265 decimalsNeg = false;
273 } else {266 } else {
274- tiling_->decimals = pow(DEFAULT_TEN, llabs(*decimalsPtr));267+ decimalsVal = pow(DEFAULT_TEN, llabs(*decimalsPtr));
275 }268 }
276 }269 }
277- DoTilingF(decimalsNeg, decimalsNan);270+ return DoTilingF(decimalsNeg, decimalsNan, decimalsVal);
278 }271 }
279- 
280- SetTilingData();
281- return ge::GRAPH_SUCCESS;
282}272}
283 273 
284static ge::graphStatus TilingPrepareForRound(gert::TilingParseContext* context)274static ge::graphStatus TilingPrepareForRound(gert::TilingParseContext* context)
@@ -17,7 +17,6 @@
17 17 
18#include "register/tilingdata_base.h"18#include "register/tilingdata_base.h"
19#include "atvoss/elewise/elewise_tiling.h"19#include "atvoss/elewise/elewise_tiling.h"
20-#include "../../op_kernel/arch35/round_tiling_struct.h"
21 20 
22// 声明具有外部链接的对象或函数21// 声明具有外部链接的对象或函数
23namespace optiling {22namespace optiling {
@@ -10,50 +10,52 @@
10 10 
11/* !11/* !
12 * \file round.cpp12 * \file round.cpp
13- * \brief 13+ * \brief
14 */14 */
15#include "kernel_operator.h"15#include "kernel_operator.h"
16#include "kernel_tiling/kernel_tiling.h"16#include "kernel_tiling/kernel_tiling.h"
17#include "arch35/round_dag.h"17#include "arch35/round_dag.h"
18#include "arch35/round_struct.h"18#include "arch35/round_struct.h"
19-#include "atvoss/elewise/elewise_sch.h"19+#include "atvoss/elewise/elewise_sch_with_scalar.h"
20#include "atvoss/util/dfx.h"20#include "atvoss/util/dfx.h"
21-#include "arch35/round_tiling_struct.h"
22 21 
23using namespace AscendC;22using namespace AscendC;
24using namespace RoundOp;23using namespace RoundOp;
25using namespace Ops::Base;24using namespace Ops::Base;
26 25 
27template <uint64_t schMode, uint64_t dType, typename DtypeX>26template <uint64_t schMode, uint64_t dType, typename DtypeX>
28-__global__ __aicore__ void RoundKernelI(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) {27+__global__ __aicore__ void RoundKernelI(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
28+{
29 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);29 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
30- REGISTER_TILING_DEFAULT(RoundTilingData);30+ REGISTER_TILING_DEFAULT(EleBaseTilingData32B);
31- GET_TILING_DATA_WITH_STRUCT(RoundTilingData, tilingData, tiling);31+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData32B, tilingData, tiling);
32- TPipe pipe;
33 32 
34 if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32)) {33 if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32)) {
35- ElementwiseSch<schMode, typename RoundDag::RoundInt<int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);34+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode, typename RoundDag::RoundInt<int32_t>::OpDag> sch(
35+ tilingData);
36 sch.Init(x, y);36 sch.Init(x, y);
37 sch.Process();37 sch.Process();
38 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_CONST)) {38 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_CONST)) {
39- ElementwiseSch<schMode, typename RoundDag::RoundIntConst<int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);39+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode, typename RoundDag::RoundIntConst<int32_t>::OpDag> sch(
40- sch.template SetVar<int, 0>(tilingData.num);40+ tilingData);
41 sch.Init(y);41 sch.Init(y);
42 sch.Process();42 sch.Process();
43- } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_NEG_NINE)) {43+ } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_NEG_NINE)) {
44- ElementwiseSch<schMode, typename RoundDag::RoundIntNegativeDecimalsNine<int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);44+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode,
45+ typename RoundDag::RoundIntNegativeDecimalsNine<int32_t>::OpDag>
46+ sch(tilingData);
45 sch.Init(x, y);47 sch.Init(x, y);
46 sch.Process();48 sch.Process();
47 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_NEGINF)) {49 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_NEGINF)) {
48- ElementwiseSch<schMode, typename RoundDag::RoundIntNegativeDecimalsInf<int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);50+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode,
49- sch.template SetVar<int, 0>(tilingData.power);51+ typename RoundDag::RoundIntNegativeDecimalsInf<int32_t>::OpDag>
50- sch.template SetVar<int, 1>(tilingData.num);52+ sch(tilingData);
51 sch.Init(x, y);53 sch.Init(x, y);
52 sch.Process();54 sch.Process();
53 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_NEG)) {55 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_INT32_NEG)) {
54- ElementwiseSch<schMode, typename RoundDag::RoundIntNegativeDecimals<int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);56+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode,
55- sch.template SetVar<int, 0>(tilingData.power);57+ typename RoundDag::RoundIntNegativeDecimals<int32_t>::OpDag>
56- sch.template SetVar<int, 1>(tilingData.num);58+ sch(tilingData);
57 sch.Init(x, y);59 sch.Init(x, y);
58 sch.Process();60 sch.Process();
59 }61 }
@@ -62,28 +64,30 @@ __global__ __aicore__ void RoundKernelI(GM_ADDR x, GM_ADDR y, GM_ADDR workspace,
62}64}
63 65 
64template <uint64_t schMode, uint64_t dType, typename DtypeX>66template <uint64_t schMode, uint64_t dType, typename DtypeX>
65-__global__ __aicore__ void RoundKernelF(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) {67+__global__ __aicore__ void RoundKernelF(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
68+{
66 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);69 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
67- REGISTER_TILING_DEFAULT(RoundTilingData);70+ REGISTER_TILING_DEFAULT(EleBaseTilingData32B);
68- GET_TILING_DATA_WITH_STRUCT(RoundTilingData, tilingData, tiling);71+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData32B, tilingData, tiling);
69- TPipe pipe;
70 72 
71 if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_ZERO)) {73 if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_ZERO)) {
72- ElementwiseSch<schMode, typename RoundDag::RoundZero<DtypeX>::OpDag> sch(&(tilingData.baseTiling), &pipe);74+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode, typename RoundDag::RoundZero<DtypeX>::OpDag> sch(
75+ tilingData);
73 sch.Init(x, y);76 sch.Init(x, y);
74 sch.Process();77 sch.Process();
75 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS)) {78 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_POSITIVE_DECIMALS)) {
76- ElementwiseSch<schMode, typename RoundDag::RoundPositiveDecimals<DtypeX>::OpDag> sch(&(tilingData.baseTiling), &pipe);79+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode, typename RoundDag::RoundPositiveDecimals<DtypeX>::OpDag>
77- sch.template SetVar<float, 0>(tilingData.decimals);80+ sch(tilingData);
78 sch.Init(x, y);81 sch.Init(x, y);
79 sch.Process();82 sch.Process();
80 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS)) {83 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_NEGATIVE_DECIMALS)) {
81- ElementwiseSch<schMode, typename RoundDag::RoundNegativeDecimals<DtypeX>::OpDag> sch(&(tilingData.baseTiling), &pipe);84+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode, typename RoundDag::RoundNegativeDecimals<DtypeX>::OpDag>
82- sch.template SetVar<float, 0>(tilingData.decimals);85+ sch(tilingData);
83 sch.Init(x, y);86 sch.Init(x, y);
84 sch.Process();87 sch.Process();
85 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS)) {88 } else if constexpr (dType == static_cast<uint64_t>(ROUND_TPL_NAN_DECIMALS)) {
86- ElementwiseSch<schMode, typename RoundDag::RoundNan<DtypeX>::OpDag> sch(&(tilingData.baseTiling), &pipe);89+ ElementwiseSchWithScalar<EleBaseTilingData32B, schMode, typename RoundDag::RoundNan<DtypeX>::OpDag> sch(
90+ tilingData);
87 sch.Init(y);91 sch.Init(y);
88 sch.Process();92 sch.Process();
89 }93 }
@@ -92,11 +96,12 @@ __global__ __aicore__ void RoundKernelF(GM_ADDR x, GM_ADDR y, GM_ADDR workspace,
92}96}
93 97 
94template <uint64_t schMode, uint64_t dType>98template <uint64_t schMode, uint64_t dType>
95-__global__ __aicore__ void round(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) {99+__global__ __aicore__ void round(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
100+{
96 if constexpr (std::is_same<DTYPE_X, int32_t>::value) {101 if constexpr (std::is_same<DTYPE_X, int32_t>::value) {
97 RoundKernelI<schMode, dType, DTYPE_X>(x, y, workspace, tiling);102 RoundKernelI<schMode, dType, DTYPE_X>(x, y, workspace, tiling);
98 } else {103 } else {
99 RoundKernelF<schMode, dType, DTYPE_X>(x, y, workspace, tiling);104 RoundKernelF<schMode, dType, DTYPE_X>(x, y, workspace, tiling);
100 }105 }
101 return;106 return;
102-}107+}
@@ -31,7 +31,7 @@ const uint64_t RSQRT_KEY_FP16 = 101UL;
31const uint64_t RSQRT_KEY_BF16 = 102UL;31const uint64_t RSQRT_KEY_BF16 = 102UL;
32const uint64_t RSQRT_KEY_FP32 = 103UL;32const uint64_t RSQRT_KEY_FP32 = 103UL;
33 33 
34-ge::graphStatus RsqrtTiling::SetTilingData()34+ge::graphStatus RsqrtTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
35{35{
36 OP_LOGD(tilingContext->GetNodeName(), "RsqrtTiling SetTilingData enter.");36 OP_LOGD(tilingContext->GetNodeName(), "RsqrtTiling SetTilingData enter.");
37 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tilingContext->GetRawTilingData());37 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tilingContext->GetRawTilingData());
@@ -53,7 +53,7 @@ ge::graphStatus RsqrtTiling::SetTilingData()
53 }53 }
54 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);54 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);
55 tilingContext->SetTilingKey(tilingKey);55 tilingContext->SetTilingKey(tilingKey);
56- tilingContext->SetBlockDim(tiling_->blockNum);56+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
57 return ge::GRAPH_SUCCESS;57 return ge::GRAPH_SUCCESS;
58}58}
59 59 
@@ -114,7 +114,7 @@ ge::graphStatus RsqrtTiling::RunTiling()
114 OP_CHECK_IF(tilingContext == nullptr, OP_LOGE("RunTiling", "Tiling context is null"), return ge::GRAPH_FAILED);114 OP_CHECK_IF(tilingContext == nullptr, OP_LOGE("RunTiling", "Tiling context is null"), return ge::GRAPH_FAILED);
115 OP_LOGD(tilingContext->GetNodeName(), "RsqrtTiling RunTiling enter.");115 OP_LOGD(tilingContext->GetNodeName(), "RsqrtTiling RunTiling enter.");
116 Ops::Base::ElewiseBaseTiling elewiseBaseTiling(tilingContext);116 Ops::Base::ElewiseBaseTiling elewiseBaseTiling(tilingContext);
117- tiling_ = tilingContext->GetTilingData<EleBaseTilingDataV2>();117+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
118 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input dtype failed"),118 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input dtype failed"),
119 return ge::GRAPH_FAILED);119 return ge::GRAPH_FAILED);
120 OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get output dtype failed"),120 OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get output dtype failed"),
@@ -123,12 +123,12 @@ ge::graphStatus RsqrtTiling::RunTiling()
123 return ge::GRAPH_FAILED);123 return ge::GRAPH_FAILED);
124 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;124 ge::graphStatus baseTilingResult = ge::GRAPH_FAILED;
125 if (this->outputDtype == ge::DT_FLOAT16) {125 if (this->outputDtype == ge::DT_FLOAT16) {
126- baseTilingResult = elewiseBaseTiling.DoTiling<RsqrtDag::RsqrtOp<half>::OpDag>(*tiling_);126+ baseTilingResult = elewiseBaseTiling.DoTiling<RsqrtDag::RsqrtOp<half>::OpDag>(*tiling);
127 } else if (this->outputDtype == ge::DT_BF16) {127 } else if (this->outputDtype == ge::DT_BF16) {
128 baseTilingResult = elewiseBaseTiling.DoTiling<RsqrtDag::RsqrtOp<half>::OpDag>(128 baseTilingResult = elewiseBaseTiling.DoTiling<RsqrtDag::RsqrtOp<half>::OpDag>(
129- *tiling_); // bfloat16类型host没定义129+ *tiling); // bfloat16类型host没定义
130 } else if (this->outputDtype == ge::DT_FLOAT) {130 } else if (this->outputDtype == ge::DT_FLOAT) {
131- baseTilingResult = elewiseBaseTiling.DoTiling<RsqrtDag::RsqrtOp<float>::OpDag>(*tiling_);131+ baseTilingResult = elewiseBaseTiling.DoTiling<RsqrtDag::RsqrtOp<float>::OpDag>(*tiling);
132 } else {132 } else {
133 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",133 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",
134 ge::TypeUtils::DataTypeToSerialString(this->outputDtype), "FLOAT16, BF16, FLOAT");134 ge::TypeUtils::DataTypeToSerialString(this->outputDtype), "FLOAT16, BF16, FLOAT");
@@ -136,7 +136,7 @@ ge::graphStatus RsqrtTiling::RunTiling()
136 }136 }
137 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),137 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),
138 return ge::GRAPH_FAILED);138 return ge::GRAPH_FAILED);
139- baseTilingResult = SetTilingData();139+ baseTilingResult = SetTilingData(elewiseBaseTiling);
140 return baseTilingResult;140 return baseTilingResult;
141}141}
142 142 
@@ -34,10 +34,9 @@ protected:
34 ge::graphStatus CalcOutputDtype();34 ge::graphStatus CalcOutputDtype();
35 ge::graphStatus CalcInputDtype();35 ge::graphStatus CalcInputDtype();
36 ge::graphStatus CheckShape() const;36 ge::graphStatus CheckShape() const;
37- ge::graphStatus SetTilingData();37+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
38 38 
39private:39private:
40- EleBaseTilingDataV2* tiling_ = nullptr;
41 gert::TilingContext* tilingContext = nullptr;40 gert::TilingContext* tilingContext = nullptr;
42 ge::DataType outputDtype = ge::DT_UNDEFINED;41 ge::DataType outputDtype = ge::DT_UNDEFINED;
43 ge::DataType inputDtype = ge::DT_UNDEFINED;42 ge::DataType inputDtype = ge::DT_UNDEFINED;
@@ -8,49 +8,47 @@
8 * See LICENSE in the root of the software repository for the full text of the License.8 * See LICENSE in the root of the software repository for the full text of the License.
9 */9 */
10 10 
11- /*!11+/*!
12 * \file rsqrt.cpp12 * \file rsqrt.cpp
13 * \brief13 * \brief
14 */14 */
15#include "kernel_operator.h"15#include "kernel_operator.h"
16#include "kernel_tiling/kernel_tiling.h"16#include "kernel_tiling/kernel_tiling.h"
17-#include "atvoss/elewise/elewise_sch.h"17+#include "atvoss/elewise/elewise_sch_16b.h"
18#include "atvoss/util/dfx.h"18#include "atvoss/util/dfx.h"
19#include "arch35/rsqrt.h"19#include "arch35/rsqrt.h"
20 20 
21using namespace AscendC;21using namespace AscendC;
22using namespace Ops::Base;22using namespace Ops::Base;
23 23 
24-extern "C" __global__ __aicore__ void rsqrt(GM_ADDR x, GM_ADDR y,24+extern "C" __global__ __aicore__ void rsqrt(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
25- GM_ADDR workspace, GM_ADDR tiling)
26{25{
27- if (workspace == nullptr) {26+ if (workspace == nullptr) {
27+ return;
28+ }
29+ SetSysWorkspace(workspace);
30+ GM_ADDR userWS = GetUserWorkspace(workspace);
31+ if (userWS == nullptr) {
32+ return;
33+ }
34+ KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
35+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
36+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
37+ if (TILING_KEY_IS(101UL)) {
38+ ElementwiseSch16B<0UL, RsqrtDag::RsqrtOp<half>::OpDag> sch(tilingData);
39+ sch.Init(x, y);
40+ sch.Process();
41+ return;
42+ } else if (TILING_KEY_IS(102UL)) {
43+ ElementwiseSch16B<0UL, RsqrtDag::RsqrtOp<bfloat16_t>::OpDag> sch(tilingData);
44+ sch.Init(x, y);
45+ sch.Process();
46+ return;
47+ } else if (TILING_KEY_IS(103UL)) {
48+ ElementwiseSch16B<0UL, RsqrtDag::RsqrtOp<float>::OpDag> sch(tilingData);
49+ sch.Init(x, y);
50+ sch.Process();
51+ return;
52+ }
28 return;53 return;
29- }54+}
30- SetSysWorkspace(workspace);
31- GM_ADDR userWS = GetUserWorkspace(workspace);
32- if (userWS == nullptr) {
33- return;
34- }
35- KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
36- REGISTER_TILING_DEFAULT(EleBaseTilingDataV2);
37- GET_TILING_DATA_WITH_STRUCT(EleBaseTilingDataV2, tilingData, tiling);
38- TPipe pipe;
39- if (TILING_KEY_IS(101UL)) {
40- ElementwiseSch<0UL, RsqrtDag::RsqrtOp<half>::OpDag> sch(&tilingData, &pipe);
41- sch.Init(x, y);
42- sch.Process();
43- return;
44- } else if (TILING_KEY_IS(102UL)) {
45- ElementwiseSch<0UL, RsqrtDag::RsqrtOp<bfloat16_t>::OpDag> sch(&tilingData, &pipe);
46- sch.Init(x, y);
47- sch.Process();
48- return;
49- } else if (TILING_KEY_IS(103UL)) {
50- ElementwiseSch<0UL, RsqrtDag::RsqrtOp<float>::OpDag> sch(&tilingData, &pipe);
51- sch.Init(x, y);
52- sch.Process();
53- return;
54- }
55- return;
56-}
@@ -45,7 +45,7 @@ ge::graphStatus SignTiling::CalcOutputDtype()
45 return ge::GRAPH_SUCCESS;45 return ge::GRAPH_SUCCESS;
46}46}
47 47 
48-ge::graphStatus SignTiling::SetTilingData()48+ge::graphStatus SignTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
49{49{
50 OP_LOGD(tilingContext->GetNodeName(), "SignTiling SetTilingData enter.");50 OP_LOGD(tilingContext->GetNodeName(), "SignTiling SetTilingData enter.");
51 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tilingContext->GetRawTilingData());51 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tilingContext->GetRawTilingData());
@@ -84,7 +84,7 @@ ge::graphStatus SignTiling::SetTilingData()
84 }84 }
85 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);85 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);
86 tilingContext->SetTilingKey(tilingKey);86 tilingContext->SetTilingKey(tilingKey);
87- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);87+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
88 return ge::GRAPH_SUCCESS;88 return ge::GRAPH_SUCCESS;
89}89}
90 90 
@@ -133,7 +133,7 @@ ge::graphStatus SignTiling::RunTiling()
133{133{
134 OP_LOGD(tilingContext->GetNodeName(), "SignTiling RunTiling enter.");134 OP_LOGD(tilingContext->GetNodeName(), "SignTiling RunTiling enter.");
135 ElewiseBaseTiling elewiseBaseTiling(tilingContext);135 ElewiseBaseTiling elewiseBaseTiling(tilingContext);
136- tiling = tilingContext->GetTilingData<SignTilingData>();136+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
137 // 获取tiling计算所需参数137 // 获取tiling计算所需参数
138 ge::graphStatus baseTilingResult = CalcOutputDtype();138 ge::graphStatus baseTilingResult = CalcOutputDtype();
139 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input dtype failed"),139 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input dtype failed"),
@@ -143,27 +143,27 @@ ge::graphStatus SignTiling::RunTiling()
143 OP_CHECK_IF(CheckOutputShape() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),143 OP_CHECK_IF(CheckOutputShape() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),
144 return ge::GRAPH_FAILED);144 return ge::GRAPH_FAILED);
145 if (this->outputDtype == ge::DT_FLOAT16) {145 if (this->outputDtype == ge::DT_FLOAT16) {
146- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<half>::OpDag>(tiling->baseTiling);146+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<half>::OpDag>(*tiling);
147 } else if (this->outputDtype == ge::DT_FLOAT) {147 } else if (this->outputDtype == ge::DT_FLOAT) {
148- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<float>::OpDag>(tiling->baseTiling);148+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<float>::OpDag>(*tiling);
149 } else if (this->outputDtype == ge::DT_BF16) {149 } else if (this->outputDtype == ge::DT_BF16) {
150- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForBf<half>::OpDag>(tiling->baseTiling);150+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForBf<half>::OpDag>(*tiling);
151 } else if (this->outputDtype == ge::DT_INT32) {151 } else if (this->outputDtype == ge::DT_INT32) {
152- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<int32_t>::OpDag>(tiling->baseTiling);152+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<int32_t>::OpDag>(*tiling);
153 } else if (this->outputDtype == ge::DT_INT64) {153 } else if (this->outputDtype == ge::DT_INT64) {
154- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForInt64<int64_t>::OpDag>(tiling->baseTiling);154+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForInt64<int64_t>::OpDag>(*tiling);
155 } else if (this->outputDtype == ge::DT_INT8) {155 } else if (this->outputDtype == ge::DT_INT8) {
156- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<int8_t>::OpDag>(tiling->baseTiling);156+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<int8_t>::OpDag>(*tiling);
157 } else if (this->outputDtype == ge::DT_INT16) {157 } else if (this->outputDtype == ge::DT_INT16) {
158- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<int16_t>::OpDag>(tiling->baseTiling);158+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<int16_t>::OpDag>(*tiling);
159 } else if (this->outputDtype == ge::DT_UINT8) {159 } else if (this->outputDtype == ge::DT_UINT8) {
160- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<uint8_t>::OpDag>(tiling->baseTiling);160+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<uint8_t>::OpDag>(*tiling);
161 } else if (this->outputDtype == ge::DT_UINT16) {161 } else if (this->outputDtype == ge::DT_UINT16) {
162- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<uint16_t>::OpDag>(tiling->baseTiling);162+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<uint16_t>::OpDag>(*tiling);
163 } else if (this->outputDtype == ge::DT_UINT32) {163 } else if (this->outputDtype == ge::DT_UINT32) {
164- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<uint32_t>::OpDag>(tiling->baseTiling);164+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForNumber<uint32_t>::OpDag>(*tiling);
165 } else if (this->outputDtype == ge::DT_UINT64) {165 } else if (this->outputDtype == ge::DT_UINT64) {
166- baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForInt64<uint64_t>::OpDag>(tiling->baseTiling);166+ baseTilingResult = elewiseBaseTiling.DoTiling<SignDag::SignForInt64<uint64_t>::OpDag>(*tiling);
167 } else {167 } else {
168 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",168 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",
169 ge::TypeUtils::DataTypeToSerialString(this->outputDtype),169 ge::TypeUtils::DataTypeToSerialString(this->outputDtype),
@@ -172,7 +172,7 @@ ge::graphStatus SignTiling::RunTiling()
172 }172 }
173 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),173 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),
174 return ge::GRAPH_FAILED);174 return ge::GRAPH_FAILED);
175- baseTilingResult = SetTilingData();175+ baseTilingResult = SetTilingData(elewiseBaseTiling);
176 return baseTilingResult;176 return baseTilingResult;
177}177}
178 178 
@@ -18,11 +18,9 @@
18#include "register/tilingdata_base.h"18#include "register/tilingdata_base.h"
19#include "register/op_impl_registry.h"19#include "register/op_impl_registry.h"
20#include "atvoss/elewise/elewise_tiling.h"20#include "atvoss/elewise/elewise_tiling.h"
21-#include "../../op_kernel/sign_struct.h"
22 21 
23namespace optiling {22namespace optiling {
24using namespace Ops::Base;23using namespace Ops::Base;
25-using namespace SignNs;
26 24 
27class SignTiling {25class SignTiling {
28public:26public:
@@ -35,13 +33,12 @@ protected:
35 ge::graphStatus CalcInputDtype();33 ge::graphStatus CalcInputDtype();
36 ge::graphStatus CheckOutputShape() const;34 ge::graphStatus CheckOutputShape() const;
37 std::string DataTypeToSerialString(ge::DataType type);35 std::string DataTypeToSerialString(ge::DataType type);
38- ge::graphStatus SetTilingData();36+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
39 37 
40private:38private:
41 gert::TilingContext* tilingContext = nullptr;39 gert::TilingContext* tilingContext = nullptr;
42 ge::DataType outputDtype = ge::DT_UNDEFINED;40 ge::DataType outputDtype = ge::DT_UNDEFINED;
43 ge::DataType inputDtype = ge::DT_UNDEFINED;41 ge::DataType inputDtype = ge::DT_UNDEFINED;
44- SignTilingData* tiling = nullptr;
45};42};
46 43 
47} // namespace optiling44} // namespace optiling
@@ -15,12 +15,10 @@
15#include "arch35/sign.h"15#include "arch35/sign.h"
16#include "kernel_operator.h"16#include "kernel_operator.h"
17#include "kernel_tiling/kernel_tiling.h"17#include "kernel_tiling/kernel_tiling.h"
18-#include "atvoss/elewise/elewise_sch.h"18+#include "atvoss/elewise/elewise_sch_16b.h"
19#include "atvoss/elewise/elewise_base_struct.h"19#include "atvoss/elewise/elewise_base_struct.h"
20-#include "sign_struct.h"
21 20 
22using namespace Ops::Base;21using namespace Ops::Base;
23-using namespace SignNs;
24 22 
25extern "C" __global__ __aicore__ void sign(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)23extern "C" __global__ __aicore__ void sign(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
26{24{
@@ -33,61 +31,60 @@ extern "C" __global__ __aicore__ void sign(GM_ADDR x, GM_ADDR y, GM_ADDR workspa
33 return;31 return;
34 }32 }
35 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);33 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
36- REGISTER_TILING_DEFAULT(SignTilingData);34+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
37- GET_TILING_DATA_WITH_STRUCT(SignTilingData, tilingData, tiling);35+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
38- TPipe pipe;
39 if (TILING_KEY_IS(101UL)) {36 if (TILING_KEY_IS(101UL)) {
40- ElementwiseSch<0UL, SignDag::SignForNumber<half>::OpDag> sch(&(tilingData.baseTiling), &pipe);37+ ElementwiseSch16B<0UL, SignDag::SignForNumber<half>::OpDag> sch(tilingData);
41 sch.Init(x, y);38 sch.Init(x, y);
42 sch.Process();39 sch.Process();
43 return;40 return;
44 } else if (TILING_KEY_IS(102UL)) {41 } else if (TILING_KEY_IS(102UL)) {
45- ElementwiseSch<0UL, SignDag::SignForBf<bfloat16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);42+ ElementwiseSch16B<0UL, SignDag::SignForBf<bfloat16_t>::OpDag> sch(tilingData);
46 sch.Init(x, y);43 sch.Init(x, y);
47 sch.Process();44 sch.Process();
48 return;45 return;
49 } else if (TILING_KEY_IS(103UL)) {46 } else if (TILING_KEY_IS(103UL)) {
50- ElementwiseSch<0UL, SignDag::SignForNumber<float>::OpDag> sch(&(tilingData.baseTiling), &pipe);47+ ElementwiseSch16B<0UL, SignDag::SignForNumber<float>::OpDag> sch(tilingData);
51 sch.Init(x, y);48 sch.Init(x, y);
52 sch.Process();49 sch.Process();
53 return;50 return;
54 } else if (TILING_KEY_IS(104UL)) {51 } else if (TILING_KEY_IS(104UL)) {
55- ElementwiseSch<0UL, SignDag::SignForNumber<int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);52+ ElementwiseSch16B<0UL, SignDag::SignForNumber<int32_t>::OpDag> sch(tilingData);
56 sch.Init(x, y);53 sch.Init(x, y);
57 sch.Process();54 sch.Process();
58 return;55 return;
59 } else if (TILING_KEY_IS(105UL)) {56 } else if (TILING_KEY_IS(105UL)) {
60- ElementwiseSch<0UL, SignDag::SignForInt64<int64_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);57+ ElementwiseSch16B<0UL, SignDag::SignForInt64<int64_t>::OpDag> sch(tilingData);
61 sch.Init(x, y);58 sch.Init(x, y);
62 sch.Process();59 sch.Process();
63 return;60 return;
64 } else if (TILING_KEY_IS(106UL)) {61 } else if (TILING_KEY_IS(106UL)) {
65- ElementwiseSch<0UL, SignDag::SignForNumber<int8_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);62+ ElementwiseSch16B<0UL, SignDag::SignForNumber<int8_t>::OpDag> sch(tilingData);
66 sch.Init(x, y);63 sch.Init(x, y);
67 sch.Process();64 sch.Process();
68 return;65 return;
69 } else if (TILING_KEY_IS(111UL)) {66 } else if (TILING_KEY_IS(111UL)) {
70- ElementwiseSch<0UL, SignDag::SignForNumber<int16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);67+ ElementwiseSch16B<0UL, SignDag::SignForNumber<int16_t>::OpDag> sch(tilingData);
71 sch.Init(x, y);68 sch.Init(x, y);
72 sch.Process();69 sch.Process();
73 return;70 return;
74 } else if (TILING_KEY_IS(107UL)) {71 } else if (TILING_KEY_IS(107UL)) {
75- ElementwiseSch<0UL, SignDag::SignForNumber<uint8_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);72+ ElementwiseSch16B<0UL, SignDag::SignForNumber<uint8_t>::OpDag> sch(tilingData);
76 sch.Init(x, y);73 sch.Init(x, y);
77 sch.Process();74 sch.Process();
78 return;75 return;
79 } else if (TILING_KEY_IS(108UL)) {76 } else if (TILING_KEY_IS(108UL)) {
80- ElementwiseSch<0UL, SignDag::SignForNumber<uint16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);77+ ElementwiseSch16B<0UL, SignDag::SignForNumber<uint16_t>::OpDag> sch(tilingData);
81 sch.Init(x, y);78 sch.Init(x, y);
82 sch.Process();79 sch.Process();
83 return;80 return;
84 } else if (TILING_KEY_IS(109UL)) {81 } else if (TILING_KEY_IS(109UL)) {
85- ElementwiseSch<0UL, SignDag::SignForNumber<uint32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);82+ ElementwiseSch16B<0UL, SignDag::SignForNumber<uint32_t>::OpDag> sch(tilingData);
86 sch.Init(x, y);83 sch.Init(x, y);
87 sch.Process();84 sch.Process();
88 return;85 return;
89 } else if (TILING_KEY_IS(110UL)) {86 } else if (TILING_KEY_IS(110UL)) {
90- ElementwiseSch<0UL, SignDag::SignForInt64<uint64_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);87+ ElementwiseSch16B<0UL, SignDag::SignForInt64<uint64_t>::OpDag> sch(tilingData);
91 sch.Init(x, y);88 sch.Init(x, y);
92 sch.Process();89 sch.Process();
93 return;90 return;
@@ -30,7 +30,7 @@ const int64_t ASCEND_WORKSPACE = 16777216; // 16M
30const int64_t ASCEND_API_BUFFER = 122880; // 120K30const int64_t ASCEND_API_BUFFER = 122880; // 120K
31const int64_t DCACHE_SIZE = 32768;31const int64_t DCACHE_SIZE = 32768;
32 32 
33-ge::graphStatus SinTiling::SetTilingData()33+ge::graphStatus SinTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
34{34{
35 OP_LOGD(tilingContext->GetNodeName(), "SinTiling SetTilingData enter.");35 OP_LOGD(tilingContext->GetNodeName(), "SinTiling SetTilingData enter.");
36 36 
@@ -38,10 +38,10 @@ ge::graphStatus SinTiling::SetTilingData()
38 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);38 OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
39 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);39 currentWorkspace[0] = static_cast<size_t>(ASCEND_WORKSPACE);
40 40 
41- const uint64_t tilingKey = GET_TPL_TILING_KEY(tiling->baseTiling.scheMode, dType);41+ const uint64_t tilingKey = GET_TPL_TILING_KEY(1, dType);
42 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);42 OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%lu", tilingKey);
43 tilingContext->SetTilingKey(tilingKey);43 tilingContext->SetTilingKey(tilingKey);
44- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);44+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
45 45 
46 uint64_t ubSize = 0;46 uint64_t ubSize = 0;
47 auto platformInfo = tilingContext->GetPlatformInfo();47 auto platformInfo = tilingContext->GetPlatformInfo();
@@ -115,11 +115,8 @@ ge::graphStatus SinTiling::RunTiling()
115{115{
116 OP_LOGD(tilingContext->GetNodeName(), "SinTiling RunTiling enter.");116 OP_LOGD(tilingContext->GetNodeName(), "SinTiling RunTiling enter.");
117 Ops::Base::ElewiseBaseTiling elewiseBaseTiling(tilingContext);117 Ops::Base::ElewiseBaseTiling elewiseBaseTiling(tilingContext);
118- tiling = tilingContext->GetTilingData<SinNs::SinTilingData>();118+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
atomgit-bot
atomgit-botatomgit-bot8月3日

🟠 High Priority

变更前,SinTiling::RunTiling() 在调用 GetTilingData 后通过 OP_CHECK_IF(tiling == nullptr, ...) 显式检查返回值是否为空,空则提前返回 GRAPH_FAILED。变更后该检查被完全移除,紧接着第 131 行直接对 tiling 解引用(*tiling)传给 DoTiling。若 GetTilingData 返回 nullptr(例如 tiling 内存分配失败、类型转换失败等边界情况),将触发空指针解引用崩溃。触发条件:GetTilingData 返回 nullptr;失败模式:空指针解引用导致进程崩溃。

改动建议
118
- auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
118
+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
119
+ OP_CHECK_IF(tiling == nullptr,
120
+ OP_LOGE_FOR_INVALID_VALUE(tilingContext->GetNodeName(), "tiling_data", "nullptr", "not nullptr"),
121
+ return ge::GRAPH_FAILED);
应用建议
likedislike
119 119 
120- OP_CHECK_IF(tiling == nullptr,
121- OP_LOGE_FOR_INVALID_VALUE(tilingContext->GetNodeName(), "tiling_data", "nullptr", "not nullptr"),
122- return ge::GRAPH_FAILED);
123 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input dtype failed"),120 OP_CHECK_IF(CalcInputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get input dtype failed"),
124 return ge::GRAPH_FAILED);121 return ge::GRAPH_FAILED);
125 OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get output dtype failed"),122 OP_CHECK_IF(CalcOutputDtype() == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "get output dtype failed"),
@@ -131,14 +128,14 @@ ge::graphStatus SinTiling::RunTiling()
131 if (this->outputDtype == ge::DT_FLOAT16) {128 if (this->outputDtype == ge::DT_FLOAT16) {
132 dType = TPL_FP16;129 dType = TPL_FP16;
133 baseTilingResult = elewiseBaseTiling.DoTiling<SinOp::SinDAG<Ops::Base::half>::OpDag>(130 baseTilingResult = elewiseBaseTiling.DoTiling<SinOp::SinDAG<Ops::Base::half>::OpDag>(
134- tiling->baseTiling, ASCEND_API_BUFFER + DCACHE_SIZE);131+ *tiling, ASCEND_API_BUFFER + DCACHE_SIZE);
135 } else if (this->outputDtype == ge::DT_BF16) {132 } else if (this->outputDtype == ge::DT_BF16) {
136 dType = TPL_BF16;133 dType = TPL_BF16;
137 baseTilingResult = elewiseBaseTiling.DoTiling<SinOp::SinDAG<Ops::Base::bfloat16_t>::OpDag>(134 baseTilingResult = elewiseBaseTiling.DoTiling<SinOp::SinDAG<Ops::Base::bfloat16_t>::OpDag>(
138- tiling->baseTiling, ASCEND_API_BUFFER + DCACHE_SIZE);135+ *tiling, ASCEND_API_BUFFER + DCACHE_SIZE);
139 } else if (this->outputDtype == ge::DT_FLOAT) {136 } else if (this->outputDtype == ge::DT_FLOAT) {
140 dType = TPL_FP32;137 dType = TPL_FP32;
141- baseTilingResult = elewiseBaseTiling.DoTiling<SinOp::SinDAG<float>::OpDag>(tiling->baseTiling,138+ baseTilingResult = elewiseBaseTiling.DoTiling<SinOp::SinDAG<float>::OpDag>(*tiling,
142 ASCEND_API_BUFFER + DCACHE_SIZE);139 ASCEND_API_BUFFER + DCACHE_SIZE);
143 } else {140 } else {
144 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",141 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "y",
@@ -148,7 +145,7 @@ ge::graphStatus SinTiling::RunTiling()
148 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),145 OP_CHECK_IF(baseTilingResult == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),
149 return ge::GRAPH_FAILED);146 return ge::GRAPH_FAILED);
150 147 
151- return SetTilingData();148+ return SetTilingData(elewiseBaseTiling);
152}149}
153 150 
154static ge::graphStatus Tiling4Sin(gert::TilingContext* tilingContextGen)151static ge::graphStatus Tiling4Sin(gert::TilingContext* tilingContextGen)
@@ -17,9 +17,9 @@
17 17 
18#include "register/tilingdata_base.h"18#include "register/tilingdata_base.h"
19#include "atvoss/elewise/elewise_tiling.h"19#include "atvoss/elewise/elewise_tiling.h"
20-#include "../../op_kernel/arch35/sin_tiling_struct.h"
21 20 
22namespace optiling {21namespace optiling {
22+using namespace Ops::Base;
23 23 
24struct SinCompileInfo {24struct SinCompileInfo {
25 uint64_t coreNum = 0;25 uint64_t coreNum = 0;
@@ -35,10 +35,9 @@ protected:
35 ge::graphStatus CalcOutputDtype();35 ge::graphStatus CalcOutputDtype();
36 ge::graphStatus CalcInputDtype();36 ge::graphStatus CalcInputDtype();
37 ge::graphStatus CheckShape() const;37 ge::graphStatus CheckShape() const;
38- ge::graphStatus SetTilingData();38+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
39 39 
40private:40private:
41- SinNs::SinTilingData* tiling = nullptr;
42 gert::TilingContext* tilingContext = nullptr;41 gert::TilingContext* tilingContext = nullptr;
43 ge::DataType outputDtype = ge::DT_UNDEFINED;42 ge::DataType outputDtype = ge::DT_UNDEFINED;
44 ge::DataType inputDtype = ge::DT_UNDEFINED;43 ge::DataType inputDtype = ge::DT_UNDEFINED;
@@ -18,30 +18,30 @@
18#include "arch35/sin_struct.h"18#include "arch35/sin_struct.h"
19#include "atvoss/util/dfx.h"19#include "atvoss/util/dfx.h"
20#include "atvoss/elewise/elewise_base_struct.h"20#include "atvoss/elewise/elewise_base_struct.h"
21-#include "atvoss/elewise/elewise_sch.h"21+#include "atvoss/elewise/elewise_sch_16b.h"
22-#include "arch35/sin_tiling_struct.h"
23 22 
24using namespace AscendC;23using namespace AscendC;
24+using namespace Ops::Base;
25 25 
26template <uint64_t schMode, uint64_t dType>26template <uint64_t schMode, uint64_t dType>
27-__global__ __aicore__ void sin(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) {27+__global__ __aicore__ void sin(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
28+{
28 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);29 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
29- REGISTER_TILING_DEFAULT(SinNs::SinTilingData);30+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
30- GET_TILING_DATA_WITH_STRUCT(SinNs::SinTilingData, tilingData, tiling);31+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
31 32 
32- TPipe pipe;33+ if constexpr (dType == TPL_FP16) {
33- if constexpr(dType == TPL_FP16) {34+ ElementwiseSch16B<schMode, SinOp::SinDAG<half>::OpDag> sch(tilingData);
34- Ops::Base::ElementwiseSch<schMode, SinOp::SinDAG<half>::OpDag> sch(&(tilingData.baseTiling), &pipe);
35 sch.Init(x, y);35 sch.Init(x, y);
36 sch.Process();36 sch.Process();
37- } else if constexpr(dType == TPL_BF16) {37+ } else if constexpr (dType == TPL_BF16) {
38- Ops::Base::ElementwiseSch<schMode, SinOp::SinDAG<bfloat16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);38+ ElementwiseSch16B<schMode, SinOp::SinDAG<bfloat16_t>::OpDag> sch(tilingData);
39 sch.Init(x, y);39 sch.Init(x, y);
40 sch.Process();40 sch.Process();
41- } else if constexpr(dType == TPL_FP32) {41+ } else if constexpr (dType == TPL_FP32) {
42- Ops::Base::ElementwiseSch<schMode, SinOp::SinDAG<float>::OpDag> sch(&(tilingData.baseTiling), &pipe);42+ ElementwiseSch16B<schMode, SinOp::SinDAG<float>::OpDag> sch(tilingData);
43 sch.Init(x, y);43 sch.Init(x, y);
44 sch.Process();44 sch.Process();
45 }45 }
46 return;46 return;
47-}47+}
@@ -88,7 +88,7 @@ ge::graphStatus SquareTiling::CheckShape() const
88 return ge::GRAPH_SUCCESS;88 return ge::GRAPH_SUCCESS;
89}89}
90 90 
91-ge::graphStatus SquareTiling::SetTilingData()91+ge::graphStatus SquareTiling::SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling)
92{92{
93 if (this->outputDtype == ge::DT_FLOAT16) {93 if (this->outputDtype == ge::DT_FLOAT16) {
94 tilingKey = TILING_KEY_FP16;94 tilingKey = TILING_KEY_FP16;
@@ -106,6 +106,12 @@ ge::graphStatus SquareTiling::SetTilingData()
106 "FLOAT16, BF16, FLOAT, INT32, INT64");106 "FLOAT16, BF16, FLOAT, INT32, INT64");
107 return ge::GRAPH_FAILED;107 return ge::GRAPH_FAILED;
108 }108 }
109+ OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%ld.", tilingKey);
110+ tilingContext->SetTilingKey(tilingKey);
111+ tilingContext->SetBlockDim(elewiseBaseTiling.GetBlockDim());
112+ size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);
113+ OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
114+ currentWorkspace[0] = 16 * 1024 * 1024;
109 return ge::GRAPH_SUCCESS;115 return ge::GRAPH_SUCCESS;
110}116}
111 117 
@@ -123,20 +129,23 @@ ge::graphStatus SquareTiling::RunTiling()
123 status = CheckShape();129 status = CheckShape();
124 OP_CHECK_IF(status == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),130 OP_CHECK_IF(status == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "check shape failed"),
125 return ge::GRAPH_FAILED);131 return ge::GRAPH_FAILED);
126- status = SetTilingData();132+ 
127- OP_CHECK_IF(status == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "SetTilingData failed"),133+ auto tiling = tilingContext->GetTilingData<EleBaseTilingData16B>();
128- return ge::GRAPH_FAILED);134+ if (this->outputDtype == ge::DT_FLOAT16) {
129- tiling = (tilingContext->GetTilingData<SquareTilingData>());135+ tilingKey = TILING_KEY_FP16;
130- if (tilingKey == TILING_KEY_FP16) {136+ status = elewiseBaseTiling.DoTiling<SquareOp<half>::OpDag>(*tiling);
131- status = elewiseBaseTiling.DoTiling<SquareOp<half>::OpDag>(tiling->baseTiling);137+ } else if (this->outputDtype == ge::DT_BF16) {
132- } else if (tilingKey == TILING_KEY_BF16) {138+ tilingKey = TILING_KEY_BF16;
133- status = elewiseBaseTiling.DoTiling<SquareOp<half>::OpDag>(tiling->baseTiling);139+ status = elewiseBaseTiling.DoTiling<SquareOp<half>::OpDag>(*tiling);
134- } else if (tilingKey == TILING_KEY_FP32) {140+ } else if (this->outputDtype == ge::DT_FLOAT) {
135- status = elewiseBaseTiling.DoTiling<SquareOp<float>::OpDag>(tiling->baseTiling);141+ tilingKey = TILING_KEY_FP32;
136- } else if (tilingKey == TILING_KEY_INT32) {142+ status = elewiseBaseTiling.DoTiling<SquareOp<float>::OpDag>(*tiling);
137- status = elewiseBaseTiling.DoTiling<SquareOp<int32_t>::OpDag>(tiling->baseTiling);143+ } else if (this->outputDtype == ge::DT_INT32) {
138- } else if (tilingKey == TILING_KEY_INT64) {144+ tilingKey = TILING_KEY_INT32;
139- status = elewiseBaseTiling.DoTiling<SquareOp<int64_t>::OpDag>(tiling->baseTiling);145+ status = elewiseBaseTiling.DoTiling<SquareOp<int32_t>::OpDag>(*tiling);
146+ } else if (this->outputDtype == ge::DT_INT64) {
147+ tilingKey = TILING_KEY_INT64;
148+ status = elewiseBaseTiling.DoTiling<SquareOp<int64_t>::OpDag>(*tiling);
140 } else {149 } else {
141 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "z",150 OP_LOGE_FOR_INVALID_DTYPE(tilingContext->GetNodeName(), "z",
142 ge::TypeUtils::DataTypeToSerialString(this->outputDtype),151 ge::TypeUtils::DataTypeToSerialString(this->outputDtype),
@@ -146,16 +155,7 @@ ge::graphStatus SquareTiling::RunTiling()
146 OP_CHECK_IF(status == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),155 OP_CHECK_IF(status == ge::GRAPH_FAILED, OP_LOGE(tilingContext->GetNodeName(), "elewiseBaseTiling failed"),
147 return ge::GRAPH_FAILED);156 return ge::GRAPH_FAILED);
148 157 
149- OP_LOGD(tilingContext->GetNodeName(), "[TilingData] : tilingKey=%ld.", tilingKey);158+ return SetTilingData(elewiseBaseTiling);
150- tilingContext->SetTilingKey(tilingKey);
151- tilingContext->SetBlockDim(tiling->baseTiling.blockNum);
152- OP_CHECK_NULL_WITH_CONTEXT(tilingContext, tilingContext->GetRawTilingData());
153- size_t usrWorkspaceSize = 0;
154- size_t sysWorkspaceSize = 16 * 1024 * 1024;
155- size_t* currentWorkspace = tilingContext->GetWorkspaceSizes(1);
156- OP_CHECK_NULL_WITH_CONTEXT(tilingContext, currentWorkspace);
157- currentWorkspace[0] = sysWorkspaceSize + usrWorkspaceSize;
158- return ge::GRAPH_SUCCESS;
159}159}
160 160 
161static ge::graphStatus Tiling4Square(gert::TilingContext* tilingContext)161static ge::graphStatus Tiling4Square(gert::TilingContext* tilingContext)
@@ -17,10 +17,8 @@
17#define OPS_BUILD_IN_OP_TILING_RUNTIME_SQUARE_TILING_H17#define OPS_BUILD_IN_OP_TILING_RUNTIME_SQUARE_TILING_H
18#include "register/tilingdata_base.h"18#include "register/tilingdata_base.h"
19#include "atvoss/elewise/elewise_tiling.h"19#include "atvoss/elewise/elewise_tiling.h"
20-#include "../../op_kernel/arch35/square_tiling_struct.h"
21 20 
22namespace optiling {21namespace optiling {
23-using namespace SquareNs;
24using namespace Ops::Base;22using namespace Ops::Base;
25 23 
26class SquareTiling {24class SquareTiling {
@@ -29,7 +27,7 @@ public:
29 ge::graphStatus RunTiling();27 ge::graphStatus RunTiling();
30 28 
31protected:29protected:
32- ge::graphStatus SetTilingData();30+ ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling);
33 ge::graphStatus CalcInputDtype();31 ge::graphStatus CalcInputDtype();
34 ge::graphStatus CalcOutputDtype();32 ge::graphStatus CalcOutputDtype();
35 ge::graphStatus CheckShape() const;33 ge::graphStatus CheckShape() const;
@@ -39,7 +37,6 @@ private:
39 ge::DataType inputDtype = ge::DT_UNDEFINED;37 ge::DataType inputDtype = ge::DT_UNDEFINED;
40 ge::DataType outputDtype = ge::DT_UNDEFINED;38 ge::DataType outputDtype = ge::DT_UNDEFINED;
41 gert::TilingContext* tilingContext = nullptr;39 gert::TilingContext* tilingContext = nullptr;
42- SquareTilingData* tiling = nullptr;
43};40};
44} // namespace optiling41} // namespace optiling
45 42 
@@ -15,12 +15,10 @@
15#include "kernel_operator.h"15#include "kernel_operator.h"
16#include "kernel_tiling/kernel_tiling.h"16#include "kernel_tiling/kernel_tiling.h"
17#include "arch35/square_dag.h"17#include "arch35/square_dag.h"
18-#include "atvoss/elewise/elewise_sch.h"18+#include "atvoss/elewise/elewise_sch_16b.h"
19#include "atvoss/util/dfx.h"19#include "atvoss/util/dfx.h"
20-#include "arch35/square_tiling_struct.h"
21 20 
22using namespace Ops::Base;21using namespace Ops::Base;
23-using namespace SquareNs;
24 22 
25extern "C" __global__ __aicore__ void square(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)23extern "C" __global__ __aicore__ void square(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
26{24{
@@ -33,34 +31,33 @@ extern "C" __global__ __aicore__ void square(GM_ADDR x, GM_ADDR y, GM_ADDR works
33 return;31 return;
34 }32 }
35 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);33 KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY);
36- REGISTER_TILING_DEFAULT(SquareTilingData);34+ REGISTER_TILING_DEFAULT(EleBaseTilingData16B);
37- GET_TILING_DATA_WITH_STRUCT(SquareTilingData, tilingData, tiling);35+ GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling);
38- TPipe pipe;
39 if (TILING_KEY_IS(1UL)) {36 if (TILING_KEY_IS(1UL)) {
40- ElementwiseSch<0UL, SquareOp<half>::OpDag> sch(&(tilingData.baseTiling), &pipe);37+ ElementwiseSch16B<0UL, SquareOp<half>::OpDag> sch(tilingData);
41 sch.Init(x, y);38 sch.Init(x, y);
42 sch.Process();39 sch.Process();
43 return;40 return;
44 } else if (TILING_KEY_IS(2UL)) {41 } else if (TILING_KEY_IS(2UL)) {
45- ElementwiseSch<0UL, SquareOp<bfloat16_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);42+ ElementwiseSch16B<0UL, SquareOp<bfloat16_t>::OpDag> sch(tilingData);
46 sch.Init(x, y);43 sch.Init(x, y);
47 sch.Process();44 sch.Process();
48 return;45 return;
49 } else if (TILING_KEY_IS(3UL)) {46 } else if (TILING_KEY_IS(3UL)) {
50- ElementwiseSch<0UL, SquareOp<float>::OpDag> sch(&(tilingData.baseTiling), &pipe);47+ ElementwiseSch16B<0UL, SquareOp<float>::OpDag> sch(tilingData);
51 sch.Init(x, y);48 sch.Init(x, y);
52 sch.Process();49 sch.Process();
53 return;50 return;
54 } else if (TILING_KEY_IS(4UL)) {51 } else if (TILING_KEY_IS(4UL)) {
55- ElementwiseSch<0UL, SquareOp<int32_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);52+ ElementwiseSch16B<0UL, SquareOp<int32_t>::OpDag> sch(tilingData);
56 sch.Init(x, y);53 sch.Init(x, y);
57 sch.Process();54 sch.Process();
58 return;55 return;
59 } else if (TILING_KEY_IS(5UL)) {56 } else if (TILING_KEY_IS(5UL)) {
60- ElementwiseSch<0UL, SquareOp<int64_t>::OpDag> sch(&(tilingData.baseTiling), &pipe);57+ ElementwiseSch16B<0UL, SquareOp<int64_t>::OpDag> sch(tilingData);
61 sch.Init(x, y);58 sch.Init(x, y);
62 sch.Process();59 sch.Process();
63 return;60 return;
64 }61 }
65 return;62 return;
66-}63+}