| @@ -24,7 +24,6 @@ | |||
| 24 | 24 | ||
| 25 | 25 | ||
| 26 | using namespace ge; | 26 | using namespace ge; |
| 27 | -using namespace AbsNs; | ||
| 28 | 27 | ||
| 29 | namespace optiling { | 28 | namespace optiling { |
| 30 | constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_BF16 = 101; | 29 | constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_BF16 = 101; |
| @@ -32,19 +31,19 @@ constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_OTHER = 102; | |||
| 32 | constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_COMPLEX = 103; | 31 | constexpr uint64_t ABS_TILING_KEY_ELEMENTWISE_COMPLEX = 103; |
| 33 | constexpr uint64_t ABS_WORKSPACE_RESERVE_BYTE = 16777216; | 32 | constexpr 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() | |||
| 83 | ge::graphStatus AbsTiling::RunTiling() | 86 | ge::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>(); | ||
| 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 | // 走新的模板tiling | 130 | // 走新的模板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 | 18 | ||
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | namespace optiling { | 21 | namespace optiling { |
| 23 | -using namespace AbsNs; | ||
| 24 | using namespace Ops::Base; | 22 | using namespace Ops::Base; |
| 25 | 23 | ||
| 26 | struct AbsCompileInfo { | 24 | struct AbsCompileInfo { |
| @@ -35,13 +33,12 @@ public: | |||
| 35 | 33 | ||
| 36 | protected: | 34 | protected: |
| 37 | ge::graphStatus CalcOutputDtype(); | 35 | ge::graphStatus CalcOutputDtype(); |
| 38 | - ge::graphStatus SetTilingData(); | 36 | + ge::graphStatus SetTilingData(const ElewiseBaseTiling& elewiseBaseTiling); |
| 39 | 37 | ||
| 40 | private: | 38 | private: |
| 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 optiling | 44 | } // namespace optiling |
| @@ -16,60 +16,57 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | -#include "atvoss/elewise/elewise_sch.h" | 19 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 20 | - | ||
| 21 | 20 | ||
| 22 | using namespace AscendC; | 21 | using namespace AscendC; |
| 23 | -using namespace AbsNs; | ||
| 24 | using namespace AbsOp; | 22 | using namespace AbsOp; |
| 25 | 23 | ||
| 26 | extern "C" __global__ __aicore__ void abs(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) | 24 | extern "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 @@ | |||
| 22 | using namespace std; | 22 | using namespace std; |
| 23 | 23 | ||
| 24 | class AbsTilingTest : public testing::Test { | 24 | class 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 | ||
| 37 | TEST_F(AbsTilingTest, test_tiling_fp16_001) | 31 | TEST_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 | ||
| 173 | TEST_F(AbsTilingTest, test_tiling_failed_empty_tensor_009) | 167 | TEST_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 | ||
| 190 | TEST_F(AbsTilingTest, test_tiling_failed_empty_tensor_010) | 184 | TEST_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 @@ | |||
| 27 | namespace optiling { | 27 | namespace optiling { |
| 28 | const size_t ASCEND_WORKSPACE = 16777216; // 16M | 28 | const 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 | ||
| 129 | static ge::graphStatus Tiling4Ceil(gert::TilingContext* tilingContextGen) | 126 | static ge::graphStatus Tiling4Ceil(gert::TilingContext* tilingContextGen) |
| @@ -17,10 +17,8 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | namespace optiling { | 21 | namespace optiling { |
| 23 | -using namespace CeilNs; | ||
| 24 | using namespace Ops::Base; | 22 | using namespace Ops::Base; |
| 25 | 23 | ||
| 26 | struct CeilCompileInfo { | 24 | struct CeilCompileInfo { |
| @@ -32,13 +30,12 @@ class CeilTiling { | |||
| 32 | public: | 30 | public: |
| 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 | ||
| 37 | protected: | 34 | protected: |
| 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 | ||
| 43 | private: | 40 | private: |
| 44 | gert::TilingContext* tilingContext = nullptr; | 41 | gert::TilingContext* tilingContext = nullptr; |
| @@ -16,32 +16,29 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | -#include "atvoss/elewise/elewise_sch.h" | 19 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 20 | - | ||
| 21 | 20 | ||
| 22 | using namespace AscendC; | 21 | using namespace AscendC; |
| 23 | using namespace Ops::Base; | 22 | using namespace Ops::Base; |
| 24 | -using namespace CeilNs; | ||
| 25 | using namespace CeilOp; | 23 | using namespace CeilOp; |
| 26 | template <uint64_t schMode, uint64_t dType> | 24 | template <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 | |||||||||
| 26 | const int64_t ASCEND_API_BUFFER = 122880; // 120K | 26 | const int64_t ASCEND_API_BUFFER = 122880; // 120K | ||||||
| 27 | const int64_t DCACHE_SIZE = 32768; | 27 | const 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>(); | ||||||
🟠 High Priority 变更前, 改动建议
![]() ![]() | |||||||||
| 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 | ||||||||
| 148 | static ge::graphStatus TilingForCos(gert::TilingContext* tilingContextGen) | 147 | static ge::graphStatus TilingForCos(gert::TilingContext* tilingContextGen) | ||||||
| @@ -16,21 +16,20 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | - | ||
| 20 | 19 | ||
| 21 | namespace optiling { | 20 | namespace optiling { |
| 21 | +using namespace Ops::Base; | ||
| 22 | 22 | ||
| 23 | class CosTiling { | 23 | class CosTiling { |
| 24 | public: | 24 | public: |
| 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 | ||
| 29 | protected: | 28 | protected: |
| 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 | ||
| 35 | private: | 34 | private: |
| 36 | gert::TilingContext* tilingContext = nullptr; | 35 | gert::TilingContext* tilingContext = nullptr; |
| @@ -12,34 +12,34 @@ | |||
| 12 | * \file cos.cpp | 12 | * \file cos.cpp |
| 13 | * \brief z = cos(x) | 13 | * \brief z = cos(x) |
| 14 | */ | 14 | */ |
| 15 | - | 15 | + |
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | -#include "arch35/cos_tilingdata.h" | 20 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 21 | - | ||
| 22 | 21 | ||
| 23 | 22 | ||
| 24 | using namespace AscendC; | 23 | using namespace AscendC; |
| 24 | +using namespace Ops::Base; | ||
| 25 | template <uint64_t schMode, uint64_t dType> | 25 | template <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 | * \brief | 13 | * \brief | ||||||
| 14 | */ | 14 | */ | ||||||
| 15 | 15 | ||||||||
| 16 | + | ||||||||
| 16 | 17 | ||||||||
| 17 | 18 | ||||||||
| 19 | + | ||||||||
| 18 | 20 | ||||||||
| 19 | 21 | ||||||||
| 20 | 22 | ||||||||
| @@ -24,7 +26,7 @@ using namespace Ops::Base; | |||||||||
| 24 | namespace optiling { | 26 | namespace optiling { | ||||||
| 25 | const int64_t ASCEND_WORKSPACE = static_cast<int64_t>(16) * 1024 * 1024; | 27 | const 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>(); | ||||||
🟠 High Priority 变更前, 改动建议
![]() ![]() | |||||||||
| 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 | ||||||||
| 126 | static ge::graphStatus TilingPrepareForLog1p(gert::TilingParseContext* context) | 127 | static ge::graphStatus TilingPrepareForLog1p(gert::TilingParseContext* context) | ||||||
| @@ -15,10 +15,11 @@ | |||
| 15 | 15 | ||
| 16 | 16 | ||
| 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 | ||
| 21 | namespace optiling { | 21 | namespace optiling { |
| 22 | +using namespace Ops::Base; | ||
| 22 | 23 | ||
| 23 | struct Log1pCompileInfo { | 24 | struct Log1pCompileInfo { |
| 24 | uint64_t coreNum = 0; | 25 | uint64_t coreNum = 0; |
| @@ -29,13 +30,12 @@ class Log1pTiling { | |||
| 29 | public: | 30 | public: |
| 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 | ||
| 34 | protected: | 34 | 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 | ||
| 40 | private: | 40 | private: |
| 41 | gert::TilingContext* tilingContext = nullptr; | 41 | gert::TilingContext* tilingContext = nullptr; |
| @@ -16,31 +16,31 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | -#include "atvoss/elewise/elewise_sch.h" | 19 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 20 | 20 | ||
| 21 | - | ||
| 22 | 21 | ||
| 23 | using namespace AscendC; | 22 | using namespace AscendC; |
| 23 | +using namespace Ops::Base; | ||
| 24 | using namespace Log1pOp; | 24 | using namespace Log1pOp; |
| 25 | template <uint64_t schMode, uint64_t dType> | 25 | template <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 | ||
| 28 | class LogicalNotTiling { | 28 | class LogicalNotTiling { |
| 29 | public: | 29 | public: |
| 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 | + | ||
| 32 | private: | 33 | private: |
| 33 | - gert::TilingContext *tilingContext; | 34 | + gert::TilingContext* tilingContext; |
| 34 | }; | 35 | }; |
| 35 | 36 | ||
| 36 | ge::graphStatus LogicalNotTiling::RunTiling() | 37 | ge::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/blockdim | 51 | // 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 | 15 | ||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | -#include "atvoss/elewise/elewise_sch.h" | 18 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | using namespace AscendC; | 21 | using namespace AscendC; |
| @@ -23,8 +23,8 @@ using namespace Ops::Base; | |||
| 23 | 23 | ||
| 24 | 24 | ||
| 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 | 22 | ||
| 23 | 23 | ||
| 24 | 24 | ||
| 25 | - | ||
| 26 | 25 | ||
| 27 | 26 | ||
| 28 | 27 | ||
| @@ -52,8 +51,7 @@ class NegTiling { | |||
| 52 | public: | 51 | public: |
| 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 | ||
| 58 | protected: | 56 | protected: |
| 59 | ge::graphStatus CalcOutputDtype(); | 57 | ge::graphStatus CalcOutputDtype(); |
| @@ -112,7 +110,7 @@ ge::graphStatus NegTiling::CheckOutputShape() const | |||
| 112 | ge::graphStatus NegTiling::RunTiling() | 110 | ge::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 | 16 | ||
| 17 | 17 | ||
| 18 | -#include "atvoss/elewise/elewise_sch.h" | 18 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | - | ||
| 22 | 21 | ||
| 23 | using namespace AscendC; | 22 | using namespace AscendC; |
| 23 | +using namespace Ops::Base; | ||
| 24 | using namespace NegOp; | 24 | using namespace NegOp; |
| 25 | 25 | ||
| 26 | template <uint64_t scheMode, uint64_t dType> | 26 | template <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 | 13 | ||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | - | ||
| 17 | 16 | ||
| 18 | 17 | ||
| 19 | using namespace std; | 18 | using namespace std; |
| 20 | using namespace ge; | 19 | using namespace ge; |
| 21 | 20 | ||
| 22 | class NegTiling : public testing::Test { | 21 | class 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 | ||
| 33 | TEST_F(NegTiling, neg_test_tiling_float16_input) | 28 | TEST_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 | ||
| 46 | TEST_F(NegTiling, neg_test_tiling_float_input) | 45 | TEST_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 | ||
| 59 | TEST_F(NegTiling, neg_test_tiling_INT32_input) | 62 | TEST_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 | ||
| 72 | TEST_F(NegTiling, neg_test_tiling_int8_input) | 79 | TEST_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 | ||
| 85 | TEST_F(NegTiling, neg_test_tiling_int64_input) | 96 | TEST_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 | ||
| 98 | TEST_F(NegTiling, neg_test_tiling_invalid_shape) | 113 | TEST_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 | ||
| 111 | TEST_F(NegTiling, neg_test_tiling_invalid_dtype) | 130 | TEST_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 { | |||
| 45 | public: | 45 | public: |
| 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 | ||
| 50 | protected: | 49 | protected: |
| 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 | ||
| 58 | private: | 57 | private: |
| 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 | ||
| 196 | ge::graphStatus RoundTiling::DoTilingI(int64_t decimals) | 191 | ge::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 | ||
| 237 | ge::graphStatus RoundTiling::RunTiling() | 234 | ge::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 | ||
| 284 | static ge::graphStatus TilingPrepareForRound(gert::TilingParseContext* context) | 274 | static ge::graphStatus TilingPrepareForRound(gert::TilingParseContext* context) |
| @@ -17,7 +17,6 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | // 声明具有外部链接的对象或函数 | 21 | // 声明具有外部链接的对象或函数 |
| 23 | namespace optiling { | 22 | namespace optiling { |
| @@ -10,50 +10,52 @@ | |||
| 10 | 10 | ||
| 11 | /* ! | 11 | /* ! |
| 12 | * \file round.cpp | 12 | * \file round.cpp |
| 13 | - * \brief | 13 | + * \brief |
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | -#include "atvoss/elewise/elewise_sch.h" | 19 | +#include "atvoss/elewise/elewise_sch_with_scalar.h" |
| 20 | 20 | ||
| 21 | - | ||
| 22 | 21 | ||
| 23 | using namespace AscendC; | 22 | using namespace AscendC; |
| 24 | using namespace RoundOp; | 23 | using namespace RoundOp; |
| 25 | using namespace Ops::Base; | 24 | using namespace Ops::Base; |
| 26 | 25 | ||
| 27 | template <uint64_t schMode, uint64_t dType, typename DtypeX> | 26 | template <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 | ||
| 64 | template <uint64_t schMode, uint64_t dType, typename DtypeX> | 66 | template <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 | ||
| 94 | template <uint64_t schMode, uint64_t dType> | 98 | template <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; | |||
| 31 | const uint64_t RSQRT_KEY_BF16 = 102UL; | 31 | const uint64_t RSQRT_KEY_BF16 = 102UL; |
| 32 | const uint64_t RSQRT_KEY_FP32 = 103UL; | 32 | const 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 | ||
| 39 | private: | 39 | private: |
| 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.cpp | 12 | * \file rsqrt.cpp |
| 13 | * \brief | 13 | * \brief |
| 14 | */ | 14 | */ |
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | -#include "atvoss/elewise/elewise_sch.h" | 17 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | using namespace AscendC; | 21 | using namespace AscendC; |
| 22 | using namespace Ops::Base; | 22 | using 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 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | - | ||
| 22 | 21 | ||
| 23 | namespace optiling { | 22 | namespace optiling { |
| 24 | using namespace Ops::Base; | 23 | using namespace Ops::Base; |
| 25 | -using namespace SignNs; | ||
| 26 | 24 | ||
| 27 | class SignTiling { | 25 | class SignTiling { |
| 28 | public: | 26 | public: |
| @@ -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 | ||
| 40 | private: | 38 | private: |
| 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 optiling | 44 | } // namespace optiling |
| @@ -15,12 +15,10 @@ | |||
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | -#include "atvoss/elewise/elewise_sch.h" | 18 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | using namespace Ops::Base; | 21 | using namespace Ops::Base; |
| 23 | -using namespace SignNs; | ||
| 24 | 22 | ||
| 25 | extern "C" __global__ __aicore__ void sign(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) | 23 | extern "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 | |||||||||||||
| 30 | const int64_t ASCEND_API_BUFFER = 122880; // 120K | 30 | const int64_t ASCEND_API_BUFFER = 122880; // 120K | ||||||||||
| 31 | const int64_t DCACHE_SIZE = 32768; | 31 | const 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>(); | ||||||||||
🟠 High Priority 变更前, 改动建议
![]() ![]() | |||||||||||||
| 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 | ||||||||||||
| 154 | static ge::graphStatus Tiling4Sin(gert::TilingContext* tilingContextGen) | 151 | static ge::graphStatus Tiling4Sin(gert::TilingContext* tilingContextGen) | ||||||||||
| @@ -17,9 +17,9 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | namespace optiling { | 21 | namespace optiling { |
| 22 | +using namespace Ops::Base; | ||
| 23 | 23 | ||
| 24 | struct SinCompileInfo { | 24 | struct 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 | ||
| 40 | private: | 40 | private: |
| 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 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | -#include "atvoss/elewise/elewise_sch.h" | 21 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 22 | - | ||
| 23 | 22 | ||
| 24 | using namespace AscendC; | 23 | using namespace AscendC; |
| 24 | +using namespace Ops::Base; | ||
| 25 | 25 | ||
| 26 | template <uint64_t schMode, uint64_t dType> | 26 | template <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 | ||
| 161 | static ge::graphStatus Tiling4Square(gert::TilingContext* tilingContext) | 161 | static ge::graphStatus Tiling4Square(gert::TilingContext* tilingContext) |
| @@ -17,10 +17,8 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | namespace optiling { | 21 | namespace optiling { |
| 23 | -using namespace SquareNs; | ||
| 24 | using namespace Ops::Base; | 22 | using namespace Ops::Base; |
| 25 | 23 | ||
| 26 | class SquareTiling { | 24 | class SquareTiling { |
| @@ -29,7 +27,7 @@ public: | |||
| 29 | ge::graphStatus RunTiling(); | 27 | ge::graphStatus RunTiling(); |
| 30 | 28 | ||
| 31 | protected: | 29 | protected: |
| 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 optiling | 41 | } // namespace optiling |
| 45 | 42 | ||
| @@ -15,12 +15,10 @@ | |||
| 15 | 15 | ||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | -#include "atvoss/elewise/elewise_sch.h" | 18 | +#include "atvoss/elewise/elewise_sch_16b.h" |
| 19 | 19 | ||
| 20 | - | ||
| 21 | 20 | ||
| 22 | using namespace Ops::Base; | 21 | using namespace Ops::Base; |
| 23 | -using namespace SquareNs; | ||
| 24 | 22 | ||
| 25 | extern "C" __global__ __aicore__ void square(GM_ADDR x, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling) | 23 | extern "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 | +} |


🟡 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正常终止。这是一个可靠性缺陷。