已合并
feat: 算子910B~910E区间判断追加IsRegBase()并统一DAV_3510硬编码判断以兼容后续Regbase芯片 #5467
hahaha22创建于 28 天前
feat: 算子910B~910E区间判断追加IsRegBase()并统一DAV_3510硬编码判断以兼容后续Regbase芯片 #5467
已合并
共 24 个文件变更+555-523
| @@ -87,9 +87,11 @@ static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySo | |||
| 87 | { | 87 | { |
| 88 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); | 88 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); |
| 89 | OP_LOGI("AddAclnn", "curArch is %u", static_cast<uint32_t>(curArch)); | 89 | OP_LOGI("AddAclnn", "curArch is %u", static_cast<uint32_t>(curArch)); |
| 90 | + if (IsRegBase(curArch)) { | ||
H | |||
| 91 | + return ASCEND910B_DTYPE_SUPPORT_LIST; | ||
| 92 | + } | ||
| 90 | switch (curArch) { | 93 | switch (curArch) { |
| 91 | - case NpuArch::DAV_2201: | 94 | + case NpuArch::DAV_2201: { |
| 92 | - case NpuArch::DAV_3510: { | ||
| 93 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 95 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 94 | } | 96 | } |
| 95 | case NpuArch::DAV_1001: { | 97 | case NpuArch::DAV_1001: { |
| @@ -40,9 +40,11 @@ static inline const std::initializer_list<DataType>& GetAiCoreDtypeSupportListBy | |||
| 40 | { | 40 | { |
| 41 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); | 41 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); |
| 42 | OP_LOGI("AddL0", "curArch is %u", static_cast<uint32_t>(curArch)); | 42 | OP_LOGI("AddL0", "curArch is %u", static_cast<uint32_t>(curArch)); |
| 43 | + if (IsRegBase(curArch)) { | ||
| 44 | + return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; | ||
| 45 | + } | ||
| 43 | switch (curArch) { | 46 | switch (curArch) { |
| 44 | - case NpuArch::DAV_2201: | 47 | + case NpuArch::DAV_2201: { |
| 45 | - case NpuArch::DAV_3510: { | ||
| 46 | return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; | 48 | return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; |
| 47 | } | 49 | } |
| 48 | case NpuArch::DAV_1001: { | 50 | case NpuArch::DAV_1001: { |
| @@ -17,6 +17,7 @@ | |||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | 19 | ||
| 20 | + | ||
| 20 | 21 | ||
| 21 | using namespace op; | 22 | using namespace op; |
| 22 | 23 | ||
| @@ -35,9 +36,11 @@ static const std::initializer_list<op::DataType> ASCEND910B_AICORE_DTYPE_SUPPORT | |||
| 35 | static inline const std::initializer_list<op::DataType>& GetAiCoreDtypeSupportListBySocVersion() | 36 | static inline const std::initializer_list<op::DataType>& GetAiCoreDtypeSupportListBySocVersion() |
| 36 | { | 37 | { |
| 37 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); | 38 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); |
| 39 | + if (IsRegBase(npuArch)) { | ||
| 40 | + return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; | ||
| 41 | + } | ||
| 38 | switch (npuArch) { | 42 | switch (npuArch) { |
| 39 | - case NpuArch::DAV_2201: | 43 | + case NpuArch::DAV_2201: { |
| 40 | - case NpuArch::DAV_3510: { | ||
| 41 | return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; | 44 | return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; |
| 42 | } | 45 | } |
| 43 | case NpuArch::DAV_1001: | 46 | case NpuArch::DAV_1001: |
| @@ -55,22 +58,21 @@ static bool IsAiCoreSupport(const aclTensor* self) | |||
| 55 | } | 58 | } |
| 56 | 59 | ||
| 57 | // AICORE算子kernel | 60 | // AICORE算子kernel |
| 58 | -static const aclTensor* DivAiCore( | 61 | +static const aclTensor* DivAiCore(const aclTensor* self, const aclTensor* other, aclTensor* divOut, |
| 59 | - const aclTensor* self, const aclTensor* other, aclTensor* divOut, aclOpExecutor* executor) | 62 | + aclOpExecutor* executor) |
| 60 | { | 63 | { |
| 61 | L0_DFX(DivAiCore, self, other, divOut); | 64 | L0_DFX(DivAiCore, self, other, divOut); |
| 62 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICORE,将AiCore Div算子加入任务队列 | 65 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICORE,将AiCore Div算子加入任务队列 |
| 63 | // Div是算子的OpType,self、other是算子的输入,divOut是算子的输出 | 66 | // Div是算子的OpType,self、other是算子的输入,divOut是算子的输出 |
| 64 | auto ret = ADD_TO_LAUNCHER_LIST_AICORE(Div, OP_INPUT(self, other), OP_OUTPUT(divOut)); | 67 | auto ret = ADD_TO_LAUNCHER_LIST_AICORE(Div, OP_INPUT(self, other), OP_OUTPUT(divOut)); |
| 65 | - OP_CHECK( | 68 | + OP_CHECK(ret == ACL_SUCCESS, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "DivAiCore ADD_TO_LAUNCHER_LIST_AICORE failed."), |
| 66 | - ret == ACL_SUCCESS, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "DivAiCore ADD_TO_LAUNCHER_LIST_AICORE failed."), | 69 | + return nullptr); |
| 67 | - return nullptr); | ||
| 68 | return divOut; | 70 | return divOut; |
| 69 | } | 71 | } |
| 70 | 72 | ||
| 71 | // AICPU算子kernel | 73 | // AICPU算子kernel |
| 72 | -static const aclTensor* DivAiCpu( | 74 | +static const aclTensor* DivAiCpu(const aclTensor* self, const aclTensor* other, aclTensor* divOut, |
| 73 | - const aclTensor* self, const aclTensor* other, aclTensor* divOut, aclOpExecutor* executor) | 75 | + aclOpExecutor* executor) |
| 74 | { | 76 | { |
| 75 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICPU,将AiCpu Div算子加入任务队列 | 77 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICPU,将AiCpu Div算子加入任务队列 |
| 76 | // Div是算子的OpType,self、other是算子的输入,divOut是算子的输出 | 78 | // Div是算子的OpType,self、other是算子的输入,divOut是算子的输出 |
| @@ -78,9 +80,8 @@ static const aclTensor* DivAiCpu( | |||
| 78 | 80 | ||
| 79 | static internal::AicpuTaskSpace space("Div"); | 81 | static internal::AicpuTaskSpace space("Div"); |
| 80 | auto ret = ADD_TO_LAUNCHER_LIST_AICPU(Div, OP_ATTR_NAMES(), OP_INPUT(self, other), OP_OUTPUT(divOut)); | 82 | auto ret = ADD_TO_LAUNCHER_LIST_AICPU(Div, OP_ATTR_NAMES(), OP_INPUT(self, other), OP_OUTPUT(divOut)); |
| 81 | - OP_CHECK( | 83 | + OP_CHECK(ret == ACL_SUCCESS, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "DivAiCpu ADD_TO_LAUNCHER_LIST_AICPU failed."), |
| 82 | - ret == ACL_SUCCESS, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "DivAiCpu ADD_TO_LAUNCHER_LIST_AICPU failed."), | 84 | + return nullptr); |
| 83 | - return nullptr); | ||
| 84 | return divOut; | 85 | return divOut; |
| 85 | } | 86 | } |
| 86 | 87 | ||
| @@ -88,9 +89,8 @@ const aclTensor* Div(const aclTensor* self, const aclTensor* other, aclOpExecuto | |||
| 88 | { | 89 | { |
| 89 | op::Shape broadcastShape; | 90 | op::Shape broadcastShape; |
| 90 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { | 91 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { |
| 91 | - OP_LOGE( | 92 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), |
| 92 | - ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), | 93 | + op::ToString(other->GetViewShape()).GetString()); |
| 93 | - op::ToString(other->GetViewShape()).GetString()); | ||
| 94 | return nullptr; | 94 | return nullptr; |
| 95 | } | 95 | } |
| 96 | 96 | ||
| @@ -198,13 +198,13 @@ static bool HasEmptyTensor(const aclTensor* self) | |||
| 198 | static const std::initializer_list<op::DataType> GetInputDtypeSupportList() | 198 | static const std::initializer_list<op::DataType> GetInputDtypeSupportList() |
| 199 | { | 199 | { |
| 200 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); | 200 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); |
| 201 | + if (IsRegBase(npuArch)) { | ||
| 202 | + return REGBASE_DTYPE_SUPPORT_LIST; | ||
| 203 | + } | ||
| 201 | switch (npuArch) { | 204 | switch (npuArch) { |
| 202 | case NpuArch::DAV_2201: { | 205 | case NpuArch::DAV_2201: { |
| 203 | return DTYPE_SUPPORT_910B_LIST; | 206 | return DTYPE_SUPPORT_910B_LIST; |
| 204 | } | 207 | } |
| 205 | - case NpuArch::DAV_3510: { | ||
| 206 | - return REGBASE_DTYPE_SUPPORT_LIST; | ||
| 207 | - } | ||
| 208 | case NpuArch::DAV_1001: { | 208 | case NpuArch::DAV_1001: { |
| 209 | return DTYPE_SUPPORT_910_LIST; | 209 | return DTYPE_SUPPORT_910_LIST; |
| 210 | } | 210 | } |
| @@ -102,13 +102,13 @@ static const size_t DIM_SUPPORT_MAX = 8; | |||
| 102 | static const std::initializer_list<op::DataType> GetInputDtypeSupportList() | 102 | static const std::initializer_list<op::DataType> GetInputDtypeSupportList() |
| 103 | { | 103 | { |
| 104 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); | 104 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); |
| 105 | + if (IsRegBase(npuArch)) { | ||
| 106 | + return REGBASE_DTYPE_SUPPORT_LIST; | ||
| 107 | + } | ||
| 105 | switch (npuArch) { | 108 | switch (npuArch) { |
| 106 | case NpuArch::DAV_2201: { | 109 | case NpuArch::DAV_2201: { |
| 107 | return DTYPE_SUPPORT_910B_LIST; | 110 | return DTYPE_SUPPORT_910B_LIST; |
| 108 | } | 111 | } |
| 109 | - case NpuArch::DAV_3510: { | ||
| 110 | - return REGBASE_DTYPE_SUPPORT_LIST; | ||
| 111 | - } | ||
| 112 | case NpuArch::DAV_1001: { | 112 | case NpuArch::DAV_1001: { |
| 113 | return DTYPE_SUPPORT_910_LIST; | 113 | return DTYPE_SUPPORT_910_LIST; |
| 114 | } | 114 | } |
| @@ -16,6 +16,7 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | + | ||
| 19 | 20 | ||
| 20 | using namespace op; | 21 | using namespace op; |
| 21 | 22 | ||
| @@ -34,8 +35,10 @@ static const std::initializer_list<op::DataType> REGBASE_DTYPE_SUPPORT_LIST = { | |||
| 34 | static inline const std::initializer_list<op::DataType>& GetAiCoreDtypeSupportListBySocVersion() | 35 | static inline const std::initializer_list<op::DataType>& GetAiCoreDtypeSupportListBySocVersion() |
| 35 | { | 36 | { |
| 36 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); | 37 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); |
| 38 | + if (IsRegBase(npuArch)) { | ||
| 39 | + return REGBASE_DTYPE_SUPPORT_LIST; | ||
| 40 | + } | ||
| 37 | switch (npuArch) { | 41 | switch (npuArch) { |
| 38 | - case NpuArch::DAV_3510: | ||
| 39 | case NpuArch::DAV_2201: { | 42 | case NpuArch::DAV_2201: { |
| 40 | return REGBASE_DTYPE_SUPPORT_LIST; | 43 | return REGBASE_DTYPE_SUPPORT_LIST; |
| 41 | } | 44 | } |
| @@ -52,27 +55,25 @@ static bool IsAiCoreSupport(const aclTensor* self) | |||
| 52 | } | 55 | } |
| 53 | 56 | ||
| 54 | // AICORE算子kernel | 57 | // AICORE算子kernel |
| 55 | -static const aclTensor* FloorDivAiCore( | 58 | +static const aclTensor* FloorDivAiCore(const aclTensor* self, const aclTensor* other, aclTensor* floorDivOut, |
| 56 | - const aclTensor* self, const aclTensor* other, aclTensor* floorDivOut, aclOpExecutor* executor) | 59 | + aclOpExecutor* executor) |
| 57 | { | 60 | { |
| 58 | L0_DFX(FloorDivAiCore); | 61 | L0_DFX(FloorDivAiCore); |
| 59 | auto ret = ADD_TO_LAUNCHER_LIST_AICORE(FloorDiv, OP_INPUT(self, other), OP_OUTPUT(floorDivOut)); | 62 | auto ret = ADD_TO_LAUNCHER_LIST_AICORE(FloorDiv, OP_INPUT(self, other), OP_OUTPUT(floorDivOut)); |
| 60 | - OP_CHECK( | 63 | + OP_CHECK(ret == ACLNN_SUCCESS, |
| 61 | - ret == ACLNN_SUCCESS, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "FloorDivAiCore ADD_TO_LAUNCHER_LIST_AICORE failed."), | 64 | + OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "FloorDivAiCore ADD_TO_LAUNCHER_LIST_AICORE failed."), return nullptr); |
| 62 | - return nullptr); | ||
| 63 | return floorDivOut; | 65 | return floorDivOut; |
| 64 | } | 66 | } |
| 65 | 67 | ||
| 66 | // AICPU-tf算子kernel | 68 | // AICPU-tf算子kernel |
| 67 | -static const aclTensor* FloorDivAiCpu( | 69 | +static const aclTensor* FloorDivAiCpu(const aclTensor* self, const aclTensor* other, aclTensor* floorDivOut, |
| 68 | - const aclTensor* self, const aclTensor* other, aclTensor* floorDivOut, aclOpExecutor* executor) | 70 | + aclOpExecutor* executor) |
| 69 | { | 71 | { |
| 70 | L0_DFX(FloorDivAiCpu); | 72 | L0_DFX(FloorDivAiCpu); |
| 71 | static internal::AicpuTaskSpace space("FloorDiv", ge::DEPEND_IN_SHAPE, true); | 73 | static internal::AicpuTaskSpace space("FloorDiv", ge::DEPEND_IN_SHAPE, true); |
| 72 | auto ret = ADD_TO_LAUNCHER_LIST_AICPU(FloorDiv, OP_ATTR_NAMES(), OP_INPUT(self, other), OP_OUTPUT(floorDivOut)); | 74 | auto ret = ADD_TO_LAUNCHER_LIST_AICPU(FloorDiv, OP_ATTR_NAMES(), OP_INPUT(self, other), OP_OUTPUT(floorDivOut)); |
| 73 | - OP_CHECK( | 75 | + OP_CHECK(ret == ACLNN_SUCCESS, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "FloorDivAiCpu ADD_TO_LAUNCHER_LIST_AICPU failed."), |
| 74 | - ret == ACLNN_SUCCESS, OP_LOGE(ACLNN_ERR_INNER_NULLPTR, "FloorDivAiCpu ADD_TO_LAUNCHER_LIST_AICPU failed."), | 76 | + return nullptr); |
| 75 | - return nullptr); | ||
| 76 | return floorDivOut; | 77 | return floorDivOut; |
| 77 | } | 78 | } |
| 78 | 79 | ||
| @@ -80,9 +81,8 @@ const aclTensor* FloorDiv(const aclTensor* self, const aclTensor* other, aclOpEx | |||
| 80 | { | 81 | { |
| 81 | op::Shape broadcastShape; | 82 | op::Shape broadcastShape; |
| 82 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { | 83 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { |
| 83 | - OP_LOGE( | 84 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), |
| 84 | - ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), | 85 | + op::ToString(other->GetViewShape()).GetString()); |
| 85 | - op::ToString(other->GetViewShape()).GetString()); | ||
| 86 | return nullptr; | 86 | return nullptr; |
| 87 | } | 87 | } |
| 88 | auto out = executor->AllocTensor(broadcastShape, self->GetDataType()); | 88 | auto out = executor->AllocTensor(broadcastShape, self->GetDataType()); |
| @@ -97,12 +97,11 @@ const aclTensor* FloorDiv(const aclTensor* self, const aclTensor* other, bool is | |||
| 97 | { | 97 | { |
| 98 | op::Shape broadcastShape; | 98 | op::Shape broadcastShape; |
| 99 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { | 99 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { |
| 100 | - OP_LOGE( | 100 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), |
| 101 | - ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), | 101 | + op::ToString(other->GetViewShape()).GetString()); |
| 102 | - op::ToString(other->GetViewShape()).GetString()); | ||
| 103 | return nullptr; | 102 | return nullptr; |
| 104 | } | 103 | } |
| 105 | - | 104 | + |
| 106 | aclTensor* out; | 105 | aclTensor* out; |
| 107 | if (isScalar || self->GetDataType() == other->GetDataType()) { | 106 | if (isScalar || self->GetDataType() == other->GetDataType()) { |
| 108 | out = executor->AllocTensor(broadcastShape, self->GetDataType()); | 107 | out = executor->AllocTensor(broadcastShape, self->GetDataType()); |
| @@ -45,9 +45,11 @@ static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST_COMPLEX = {o | |||
| 45 | 45 | ||
| 46 | static const std::initializer_list<DataType>& GetDtypeSupportList(NpuArch npuArch, SocVersion socVersion) | 46 | static const std::initializer_list<DataType>& GetDtypeSupportList(NpuArch npuArch, SocVersion socVersion) |
| 47 | { | 47 | { |
| 48 | + if (IsRegBase(npuArch)) { | ||
| 49 | + return ASCEND910B_DTYPE_DTYPE_SUPPORT_LIST; | ||
| 50 | + } | ||
| 48 | switch (npuArch) { | 51 | switch (npuArch) { |
| 49 | - case NpuArch::DAV_2201: | 52 | + case NpuArch::DAV_2201: { |
| 50 | - case NpuArch::DAV_3510: { | ||
| 51 | return ASCEND910B_DTYPE_DTYPE_SUPPORT_LIST; | 53 | return ASCEND910B_DTYPE_DTYPE_SUPPORT_LIST; |
| 52 | } | 54 | } |
| 53 | case NpuArch::DAV_1001: { | 55 | case NpuArch::DAV_1001: { |
| @@ -143,7 +143,7 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* other, const ac | |||
| 143 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | 143 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() |
| 144 | { | 144 | { |
| 145 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); | 145 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); |
| 146 | - if (socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) { | 146 | + if ((socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) || IsRegBase()) { |
| 147 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 147 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 148 | } else { | 148 | } else { |
| 149 | return ASCEND910_DTYPE_SUPPORT_LIST; | 149 | return ASCEND910_DTYPE_SUPPORT_LIST; |
| @@ -153,7 +153,7 @@ static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | |||
| 153 | static inline const std::initializer_list<op::DataType>& GetOutDtypeSupportList() | 153 | static inline const std::initializer_list<op::DataType>& GetOutDtypeSupportList() |
| 154 | { | 154 | { |
| 155 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); | 155 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); |
| 156 | - if (socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) { | 156 | + if ((socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) || IsRegBase()) { |
| 157 | return ASCEND910B_OUT_DTYPE_SUPPORT_LIST; | 157 | return ASCEND910B_OUT_DTYPE_SUPPORT_LIST; |
| 158 | } | 158 | } |
| 159 | 159 | ||
| @@ -179,9 +179,8 @@ static bool CheckPromoteType(const aclTensor* self, const aclTensor* other, cons | |||
| 179 | // 检查self和other能否做数据类型推导 | 179 | // 检查self和other能否做数据类型推导 |
| 180 | op::DataType promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); | 180 | op::DataType promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); |
| 181 | if (promoteType == DataType::DT_UNDEFINED) { | 181 | if (promoteType == DataType::DT_UNDEFINED) { |
| 182 | - OP_LOGE( | 182 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", |
| 183 | - ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", | 183 | + op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); |
| 184 | - op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); | ||
| 185 | return false; | 184 | return false; |
| 186 | } | 185 | } |
| 187 | 186 | ||
| @@ -199,9 +198,8 @@ static bool CheckShape(const aclTensor* self, const aclTensor* other, const aclT | |||
| 199 | OP_CHECK_BROADCAST_AND_INFER_SHAPE(self, other, broadcastShape, return false); | 198 | OP_CHECK_BROADCAST_AND_INFER_SHAPE(self, other, broadcastShape, return false); |
| 200 | 199 | ||
| 201 | if (broadcastShape != out->GetViewShape()) { | 200 | if (broadcastShape != out->GetViewShape()) { |
| 202 | - OP_LOGE( | 201 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Shape of out should be %s, but current is %s.", |
| 203 | - ACLNN_ERR_PARAM_INVALID, "Shape of out should be %s, but current is %s.", | 202 | + op::ToString(broadcastShape).GetString(), op::ToString(out->GetViewShape()).GetString()); |
| 204 | - op::ToString(broadcastShape).GetString(), op::ToString(out->GetViewShape()).GetString()); | ||
| 205 | return false; | 203 | return false; |
| 206 | } | 204 | } |
| 207 | return true; | 205 | return true; |
| @@ -247,23 +245,21 @@ static bool CheckDtypeValidScalar(const aclTensor* self, const aclScalar* other, | |||
| 247 | return true; | 245 | return true; |
| 248 | } | 246 | } |
| 249 | 247 | ||
| 250 | -static bool CheckPromoteTypeScalar( | 248 | +static bool CheckPromoteTypeScalar(const aclTensor* self, const aclScalar* other, const aclTensor* out, |
| 251 | - const aclTensor* self, const aclScalar* other, const aclTensor* out, DataType promoteType) | 249 | + DataType promoteType) |
| 252 | { | 250 | { |
| 253 | // 检查self和other能否做数据类型推导 | 251 | // 检查self和other能否做数据类型推导 |
| 254 | if (promoteType == DataType::DT_UNDEFINED) { | 252 | if (promoteType == DataType::DT_UNDEFINED) { |
| 255 | - OP_LOGE( | 253 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", |
| 256 | - ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", | 254 | + op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); |
| 257 | - op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); | ||
| 258 | return false; | 255 | return false; |
| 259 | } | 256 | } |
| 260 | 257 | ||
| 261 | // 检查promote后的数据类型是否在Less算子的支持列表内 | 258 | // 检查promote后的数据类型是否在Less算子的支持列表内 |
| 262 | auto supportList = GetDtypeSupportList(); | 259 | auto supportList = GetDtypeSupportList(); |
| 263 | if (!CheckType(promoteType, supportList)) { | 260 | if (!CheckType(promoteType, supportList)) { |
| 264 | - OP_LOGE( | 261 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "aclnnLtScalar not implemented for input promote dtype %s.", |
| 265 | - ACLNN_ERR_PARAM_INVALID, "aclnnLtScalar not implemented for input promote dtype %s.", | 262 | + ToString(promoteType).GetString()); |
| 266 | - ToString(promoteType).GetString()); | ||
| 267 | return false; | 263 | return false; |
| 268 | } | 264 | } |
| 269 | 265 | ||
| @@ -284,8 +280,8 @@ static bool CheckShapeScalar(const aclTensor* self, const aclTensor* out) | |||
| 284 | return true; | 280 | return true; |
| 285 | } | 281 | } |
| 286 | 282 | ||
| 287 | -static aclnnStatus CheckParamsScalar( | 283 | +static aclnnStatus CheckParamsScalar(const aclTensor* self, const aclScalar* other, const aclTensor* out, |
| 288 | - const aclTensor* self, const aclScalar* other, const aclTensor* out, DataType promote) | 284 | + DataType promote) |
| 289 | { | 285 | { |
| 290 | // 1. 检查参数是否为空指针 | 286 | // 1. 检查参数是否为空指针 |
| 291 | CHECK_RET(CheckNotNullScalar(self, other, out), ACLNN_ERR_PARAM_NULLPTR); | 287 | CHECK_RET(CheckNotNullScalar(self, other, out), ACLNN_ERR_PARAM_NULLPTR); |
| @@ -302,8 +298,8 @@ static aclnnStatus CheckParamsScalar( | |||
| 302 | return ACLNN_SUCCESS; | 298 | return ACLNN_SUCCESS; |
| 303 | } | 299 | } |
| 304 | 300 | ||
| 305 | -aclnnStatus aclnnLtScalarGetWorkspaceSizeV35( | 301 | +aclnnStatus aclnnLtScalarGetWorkspaceSizeV35(const aclTensor* self, const aclScalar* other, aclTensor* out, |
| 306 | - const aclTensor* self, const aclScalar* other, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor) | 302 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 307 | { | 303 | { |
| 308 | // 固定写法,创建OpExecutor | 304 | // 固定写法,创建OpExecutor |
| 309 | auto uniqueExecutor = CREATE_EXECUTOR(); | 305 | auto uniqueExecutor = CREATE_EXECUTOR(); |
| @@ -365,8 +361,8 @@ aclnnStatus aclnnLtScalarGetWorkspaceSizeV35( | |||
| 365 | return ACLNN_SUCCESS; | 361 | return ACLNN_SUCCESS; |
| 366 | } | 362 | } |
| 367 | 363 | ||
| 368 | -aclnnStatus aclnnLtScalarGetWorkspaceSize( | 364 | +aclnnStatus aclnnLtScalarGetWorkspaceSize(const aclTensor* self, const aclScalar* other, aclTensor* out, |
| 369 | - const aclTensor* self, const aclScalar* other, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor) | 365 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 370 | { | 366 | { |
| 371 | L2_DFX_PHASE_1(aclnnLtScalar, DFX_IN(self, other), DFX_OUT(out)); | 367 | L2_DFX_PHASE_1(aclnnLtScalar, DFX_IN(self, other), DFX_OUT(out)); |
| 372 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); | 368 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); |
| @@ -446,8 +442,8 @@ aclnnStatus aclnnLtScalar(void* workspace, uint64_t workspaceSize, aclOpExecutor | |||
| 446 | } | 442 | } |
| 447 | 443 | ||
| 448 | // inplace lt scalar | 444 | // inplace lt scalar |
| 449 | -aclnnStatus aclnnInplaceLtScalarGetWorkspaceSize( | 445 | +aclnnStatus aclnnInplaceLtScalarGetWorkspaceSize(const aclTensor* selfRef, const aclScalar* other, |
| 450 | - const aclTensor* selfRef, const aclScalar* other, uint64_t* workspaceSize, aclOpExecutor** executor) | 446 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 451 | { | 447 | { |
| 452 | auto out = const_cast<aclTensor*>(selfRef); | 448 | auto out = const_cast<aclTensor*>(selfRef); |
| 453 | return aclnnLtScalarGetWorkspaceSize(selfRef, other, out, workspaceSize, executor); | 449 | return aclnnLtScalarGetWorkspaceSize(selfRef, other, out, workspaceSize, executor); |
| @@ -85,7 +85,7 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* other, const ac | |||
| 85 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | 85 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() |
| 86 | { | 86 | { |
| 87 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); | 87 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); |
| 88 | - if (socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) { | 88 | + if ((socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) || IsRegBase()) { |
| 89 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 89 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 90 | } else { | 90 | } else { |
| 91 | return ASCEND910_DTYPE_SUPPORT_LIST; | 91 | return ASCEND910_DTYPE_SUPPORT_LIST; |
| @@ -95,7 +95,7 @@ static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | |||
| 95 | static inline const std::initializer_list<op::DataType>& GetOutDtypeSupportList() | 95 | static inline const std::initializer_list<op::DataType>& GetOutDtypeSupportList() |
| 96 | { | 96 | { |
| 97 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); | 97 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); |
| 98 | - if (socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) { | 98 | + if ((socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) || IsRegBase()) { |
| 99 | return ASCEND910B_OUT_DTYPE_SUPPORT_LIST; | 99 | return ASCEND910B_OUT_DTYPE_SUPPORT_LIST; |
| 100 | } | 100 | } |
| 101 | 101 | ||
| @@ -128,9 +128,8 @@ static bool CheckPromoteType(const aclTensor* self, const aclTensor* other, cons | |||
| 128 | promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); | 128 | promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); |
| 129 | } | 129 | } |
| 130 | if (promoteType == DataType::DT_UNDEFINED) { | 130 | if (promoteType == DataType::DT_UNDEFINED) { |
| 131 | - OP_LOGE( | 131 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", |
| 132 | - ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", | 132 | + op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); |
| 133 | - op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); | ||
| 134 | return false; | 133 | return false; |
| 135 | } | 134 | } |
| 136 | 135 | ||
| @@ -138,9 +137,8 @@ static bool CheckPromoteType(const aclTensor* self, const aclTensor* other, cons | |||
| 138 | if (IsRegBase(npuArch)) { | 137 | if (IsRegBase(npuArch)) { |
| 139 | auto supportList = GetDtypeSupportList(); | 138 | auto supportList = GetDtypeSupportList(); |
| 140 | if (!CheckType(promoteType, supportList)) { | 139 | if (!CheckType(promoteType, supportList)) { |
| 141 | - OP_LOGE( | 140 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "aclnnLtTensor not implemented for input promote dtype %s.", |
| 142 | - ACLNN_ERR_PARAM_INVALID, "aclnnLtTensor not implemented for input promote dtype %s.", | 141 | + ToString(promoteType).GetString()); |
| 143 | - ToString(promoteType).GetString()); | ||
| 144 | return false; | 142 | return false; |
| 145 | } | 143 | } |
| 146 | } | 144 | } |
| @@ -158,16 +156,15 @@ static bool CheckShape(const aclTensor* self, const aclTensor* other, const aclT | |||
| 158 | op::Shape broadcastShape; | 156 | op::Shape broadcastShape; |
| 159 | OP_CHECK_BROADCAST_AND_INFER_SHAPE(self, other, broadcastShape, return false); | 157 | OP_CHECK_BROADCAST_AND_INFER_SHAPE(self, other, broadcastShape, return false); |
| 160 | if (broadcastShape != out->GetViewShape()) { | 158 | if (broadcastShape != out->GetViewShape()) { |
| 161 | - OP_LOGE( | 159 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Shape of out should be %s, but current is %s.", |
| 162 | - ACLNN_ERR_PARAM_INVALID, "Shape of out should be %s, but current is %s.", | 160 | + op::ToString(broadcastShape).GetString(), op::ToString(out->GetViewShape()).GetString()); |
| 163 | - op::ToString(broadcastShape).GetString(), op::ToString(out->GetViewShape()).GetString()); | ||
| 164 | return false; | 161 | return false; |
| 165 | } | 162 | } |
| 166 | return true; | 163 | return true; |
| 167 | } | 164 | } |
| 168 | 165 | ||
| 169 | -static aclnnStatus CheckParams( | 166 | +static aclnnStatus CheckParams(const aclTensor* self, const aclTensor* other, const aclTensor* out, |
| 170 | - const aclTensor* self, const aclTensor* other, const aclTensor* out, DataType& promoteType) | 167 | + DataType& promoteType) |
| 171 | { | 168 | { |
| 172 | // 1. 检查参数是否为空指针 | 169 | // 1. 检查参数是否为空指针 |
| 173 | CHECK_RET(CheckNotNull(self, other, out), ACLNN_ERR_PARAM_NULLPTR); | 170 | CHECK_RET(CheckNotNull(self, other, out), ACLNN_ERR_PARAM_NULLPTR); |
| @@ -184,8 +181,8 @@ static aclnnStatus CheckParams( | |||
| 184 | return ACLNN_SUCCESS; | 181 | return ACLNN_SUCCESS; |
| 185 | } | 182 | } |
| 186 | 183 | ||
| 187 | -aclnnStatus aclnnLtTensorGetWorkspaceSize( | 184 | +aclnnStatus aclnnLtTensorGetWorkspaceSize(const aclTensor* self, const aclTensor* other, aclTensor* out, |
| 188 | - const aclTensor* self, const aclTensor* other, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor) | 185 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 189 | { | 186 | { |
| 190 | L2_DFX_PHASE_1(aclnnLtTensor, DFX_IN(self, other), DFX_OUT(out)); | 187 | L2_DFX_PHASE_1(aclnnLtTensor, DFX_IN(self, other), DFX_OUT(out)); |
| 191 | 188 | ||
| @@ -252,8 +249,8 @@ aclnnStatus aclnnLtTensor(void* workspace, uint64_t workspaceSize, aclOpExecutor | |||
| 252 | } | 249 | } |
| 253 | 250 | ||
| 254 | // InplaceLt | 251 | // InplaceLt |
| 255 | -aclnnStatus aclnnInplaceLtTensorGetWorkspaceSize( | 252 | +aclnnStatus aclnnInplaceLtTensorGetWorkspaceSize(const aclTensor* selfRef, const aclTensor* other, |
| 256 | - const aclTensor* selfRef, const aclTensor* other, uint64_t* workspaceSize, aclOpExecutor** executor) | 253 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 257 | { | 254 | { |
| 258 | auto out = const_cast<aclTensor*>(selfRef); | 255 | auto out = const_cast<aclTensor*>(selfRef); |
| 259 | return aclnnLtTensorGetWorkspaceSize(selfRef, other, out, workspaceSize, executor); | 256 | return aclnnLtTensorGetWorkspaceSize(selfRef, other, out, workspaceSize, executor); |
| @@ -1,5 +1,5 @@ | |||
| 1 | /** | 1 | /** |
| 2 | - * Copyright (c) 2025 Huawei Technologies Co., Ltd. | 2 | + * Copyright (c) 2025 Huawei Technologies Co., Ltd. |
| 3 | * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | 3 | * This program is free software, you can redistribute it and/or modify it under the terms and conditions of |
| 4 | * CANN Open Software License Agreement Version 2.0 (the "License"). | 4 | * CANN Open Software License Agreement Version 2.0 (the "License"). |
| 5 | * Please refer to the License for details. You may not use this file except in compliance with the License. | 5 | * Please refer to the License for details. You may not use this file except in compliance with the License. |
| @@ -49,7 +49,7 @@ static const std::initializer_list<op::DataType> OUT_DTYPE_SUPPORT_910_LIST = { | |||
| 49 | op::DataType::DT_UINT32, op::DataType::DT_UINT64, op::DataType::DT_BOOL, op::DataType::DT_UINT16, | 49 | op::DataType::DT_UINT32, op::DataType::DT_UINT64, op::DataType::DT_BOOL, op::DataType::DT_UINT16, |
| 50 | op::DataType::DT_COMPLEX64, op::DataType::DT_COMPLEX128}; | 50 | op::DataType::DT_COMPLEX64, op::DataType::DT_COMPLEX128}; |
| 51 | 51 | ||
| 52 | -static const std::initializer_list<op::DataType> REGBASE_OUT_DTYPE_SUPPORT_LIST = { | 52 | +static const std::initializer_list<op::DataType> REGBASE_OUT_DTYPE_SUPPORT_LIST = { |
| 53 | op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, | 53 | op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, |
| 54 | op::DataType::DT_INT16, op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, | 54 | op::DataType::DT_INT16, op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, |
| 55 | op::DataType::DT_UINT32, op::DataType::DT_UINT64, op::DataType::DT_BOOL, op::DataType::DT_UINT16, | 55 | op::DataType::DT_UINT32, op::DataType::DT_UINT64, op::DataType::DT_BOOL, op::DataType::DT_UINT16, |
| @@ -147,8 +147,8 @@ static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | |||
| 147 | static inline const std::initializer_list<op::DataType>& GetOutputDtypeSupportList() | 147 | static inline const std::initializer_list<op::DataType>& GetOutputDtypeSupportList() |
| 148 | { | 148 | { |
| 149 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); | 149 | auto socVersion = GetCurrentPlatformInfo().GetSocVersion(); |
| 150 | - if (socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) { | 150 | + if ((socVersion >= SocVersion::ASCEND910B && socVersion <= SocVersion::ASCEND910E) || IsRegBase()) { |
H 这里加的 ![]() ![]() | |||
| 151 | - return REGBASE_OUT_DTYPE_SUPPORT_LIST ; | 151 | + return REGBASE_OUT_DTYPE_SUPPORT_LIST; |
| 152 | } | 152 | } |
| 153 | return OUT_DTYPE_SUPPORT_910_LIST; | 153 | return OUT_DTYPE_SUPPORT_910_LIST; |
| 154 | } | 154 | } |
| @@ -159,14 +159,11 @@ static bool CheckDtypeValid(const aclTensor* self, const aclScalar* other, const | |||
| 159 | OP_CHECK_DTYPE_NOT_SUPPORT(self, supportList, return false); | 159 | OP_CHECK_DTYPE_NOT_SUPPORT(self, supportList, return false); |
| 160 | OP_CHECK_DTYPE_NOT_SUPPORT(other, supportList, return false); | 160 | OP_CHECK_DTYPE_NOT_SUPPORT(other, supportList, return false); |
| 161 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); | 161 | auto npuArch = op::GetCurrentPlatformInfo().GetCurNpuArch(); |
| 162 | - auto outSuportList = IsRegBase(npuArch) ? | 162 | + auto outSuportList = IsRegBase(npuArch) ? GetOutputDtypeSupportList() : supportList; |
| 163 | - GetOutputDtypeSupportList() : | ||
| 164 | - supportList; | ||
| 165 | op::DataType outType = out->GetDataType(); | 163 | op::DataType outType = out->GetDataType(); |
| 166 | if ((!CheckType(outType, outSuportList)) && (outType != DataType::DT_BOOL)) { | 164 | if ((!CheckType(outType, outSuportList)) && (outType != DataType::DT_BOOL)) { |
| 167 | - OP_LOGE( | 165 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "out dtype %s should be in dtype support list [%s].", |
| 168 | - ACLNN_ERR_PARAM_INVALID, "out dtype %s should be in dtype support list [%s].", | 166 | + op::ToString(out->GetDataType()).GetString(), op::ToString(outSuportList).GetString()); |
| 169 | - op::ToString(out->GetDataType()).GetString(), op::ToString(outSuportList).GetString()); | ||
| 170 | return false; | 167 | return false; |
| 171 | } | 168 | } |
| 172 | 169 | ||
| @@ -197,9 +194,8 @@ static bool CheckPromoteType(const aclTensor* self, const aclScalar* other, cons | |||
| 197 | // 检查self和other能否做数据类型推导 | 194 | // 检查self和other能否做数据类型推导 |
| 198 | op::DataType promoteType = PromoteTypeScalar(self, other); | 195 | op::DataType promoteType = PromoteTypeScalar(self, other); |
| 199 | if (promoteType == DataType::DT_UNDEFINED) { | 196 | if (promoteType == DataType::DT_UNDEFINED) { |
| 200 | - OP_LOGE( | 197 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", |
| 201 | - ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", | 198 | + op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); |
| 202 | - op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); | ||
| 203 | return false; | 199 | return false; |
| 204 | } | 200 | } |
| 205 | 201 | ||
| @@ -210,11 +206,10 @@ static bool CheckPromoteType(const aclTensor* self, const aclScalar* other, cons | |||
| 210 | if (IsRegBase(npuArch)) { | 206 | if (IsRegBase(npuArch)) { |
| 211 | const auto& supportList = GetDtypeSupportList(); | 207 | const auto& supportList = GetDtypeSupportList(); |
| 212 | if (!CheckType(promoteType, supportList)) { | 208 | if (!CheckType(promoteType, supportList)) { |
| 213 | - OP_LOGE( | 209 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, |
| 214 | - ACLNN_ERR_PARAM_INVALID, | 210 | + "aclnnLeScalar not implemented for input dtype %s," |
| 215 | - "aclnnLeScalar not implemented for input dtype %s," | 211 | + "should be in dtype support list [%s].", |
| 216 | - "should be in dtype support list [%s].", | 212 | + ToString(promoteType).GetString(), op::ToString(supportList).GetString()); |
| 217 | - ToString(promoteType).GetString(), op::ToString(supportList).GetString()); | ||
| 218 | return false; | 213 | return false; |
| 219 | } | 214 | } |
| 220 | } | 215 | } |
| @@ -232,9 +227,8 @@ static bool CheckShape(const aclTensor* self, const aclTensor* out) | |||
| 232 | 227 | ||
| 233 | // self和out的shape必须一致 | 228 | // self和out的shape必须一致 |
| 234 | if (self->GetViewShape() != out->GetViewShape()) { | 229 | if (self->GetViewShape() != out->GetViewShape()) { |
| 235 | - OP_LOGE( | 230 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "self shape is different with out shape, self [%s], out [%s].", |
| 236 | - ACLNN_ERR_PARAM_INVALID, "self shape is different with out shape, self [%s], out [%s].", | 231 | + op::ToString(self->GetViewShape()).GetString(), op::ToString(out->GetViewShape()).GetString()); |
| 237 | - op::ToString(self->GetViewShape()).GetString(), op::ToString(out->GetViewShape()).GetString()); | ||
| 238 | return false; | 232 | return false; |
| 239 | } | 233 | } |
| 240 | 234 | ||
| @@ -258,8 +252,8 @@ static aclnnStatus CheckParams(const aclTensor* self, const aclScalar* other, co | |||
| 258 | return ACLNN_SUCCESS; | 252 | return ACLNN_SUCCESS; |
| 259 | } | 253 | } |
| 260 | 254 | ||
| 261 | -static aclnnStatus aclnnLeScalarCommon( | 255 | +static aclnnStatus aclnnLeScalarCommon(const aclTensor* self, const aclScalar* other, aclTensor* out, |
| 262 | - const aclTensor* self, const aclScalar* other, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor) | 256 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 263 | { | 257 | { |
| 264 | // 固定写法,参数检查 | 258 | // 固定写法,参数检查 |
| 265 | auto ret = CheckParams(self, other, out); | 259 | auto ret = CheckParams(self, other, out); |
| @@ -309,8 +303,8 @@ static aclnnStatus aclnnLeScalarCommon( | |||
| 309 | return ACLNN_SUCCESS; | 303 | return ACLNN_SUCCESS; |
| 310 | } | 304 | } |
| 311 | 305 | ||
| 312 | -aclnnStatus aclnnLeScalarGetWorkspaceSize( | 306 | +aclnnStatus aclnnLeScalarGetWorkspaceSize(const aclTensor* self, const aclScalar* other, aclTensor* out, |
| 313 | - const aclTensor* self, const aclScalar* other, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor) | 307 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 314 | { | 308 | { |
| 315 | L2_DFX_PHASE_1(aclnnLeScalar, DFX_IN(self, other), DFX_OUT(out)); | 309 | L2_DFX_PHASE_1(aclnnLeScalar, DFX_IN(self, other), DFX_OUT(out)); |
| 316 | return aclnnLeScalarCommon(self, other, out, workspaceSize, executor); | 310 | return aclnnLeScalarCommon(self, other, out, workspaceSize, executor); |
| @@ -323,8 +317,8 @@ aclnnStatus aclnnLeScalar(void* workspace, uint64_t workspaceSize, aclOpExecutor | |||
| 323 | return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | 317 | return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); |
| 324 | } | 318 | } |
| 325 | 319 | ||
| 326 | -aclnnStatus aclnnInplaceLeScalarGetWorkspaceSize( | 320 | +aclnnStatus aclnnInplaceLeScalarGetWorkspaceSize(aclTensor* selfRef, const aclScalar* other, uint64_t* workspaceSize, |
| 327 | - aclTensor* selfRef, const aclScalar* other, uint64_t* workspaceSize, aclOpExecutor** executor) | 321 | + aclOpExecutor** executor) |
| 328 | { | 322 | { |
| 329 | L2_DFX_PHASE_1(aclnnInplaceLeScalar, DFX_IN(selfRef, other), DFX_OUT(selfRef)); | 323 | L2_DFX_PHASE_1(aclnnInplaceLeScalar, DFX_IN(selfRef, other), DFX_OUT(selfRef)); |
| 330 | return aclnnLeScalarCommon(selfRef, other, selfRef, workspaceSize, executor); | 324 | return aclnnLeScalarCommon(selfRef, other, selfRef, workspaceSize, executor); |
| @@ -71,8 +71,9 @@ inline static bool CheckNotNull(const aclTensor* self, const aclTensor* other, c | |||
| 71 | 71 | ||
| 72 | inline static bool CheckSocVersionIsSupportBf16(void) | 72 | inline static bool CheckSocVersionIsSupportBf16(void) |
| 73 | { | 73 | { |
| 74 | - return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 74 | + return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 75 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E; | 75 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 76 | + IsRegBase(); | ||
| 76 | } | 77 | } |
| 77 | 78 | ||
| 78 | inline static bool CheckDtypeValid(const aclTensor* self, const aclTensor* other, const aclTensor* out) | 79 | inline static bool CheckDtypeValid(const aclTensor* self, const aclTensor* other, const aclTensor* out) |
| @@ -58,8 +58,9 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* other, const ac | |||
| 58 | 58 | ||
| 59 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | 59 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() |
| 60 | { | 60 | { |
| 61 | - if (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 61 | + if ((GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 62 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) { | 62 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 63 | + IsRegBase()) { | ||
| 63 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 64 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 64 | } else { | 65 | } else { |
| 65 | return ASCEND910_DTYPE_SUPPORT_LIST; | 66 | return ASCEND910_DTYPE_SUPPORT_LIST; |
| @@ -54,8 +54,9 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* other, const ac | |||
| 54 | 54 | ||
| 55 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | 55 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() |
| 56 | { | 56 | { |
| 57 | - if (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 57 | + if ((GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 58 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) { | 58 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 59 | + IsRegBase()) { | ||
| 59 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 60 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 60 | } else { | 61 | } else { |
| 61 | return ASCEND910_DTYPE_SUPPORT_LIST; | 62 | return ASCEND910_DTYPE_SUPPORT_LIST; |
| @@ -121,8 +121,9 @@ static op::DataType GetScalarDefaultDtype(const op::DataType input) | |||
| 121 | 121 | ||
| 122 | static const std::initializer_list<DataType>& GetDtypeSupportList() | 122 | static const std::initializer_list<DataType>& GetDtypeSupportList() |
| 123 | { | 123 | { |
| 124 | - if (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 124 | + if ((GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 125 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) { | 125 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 126 | + IsRegBase()) { | ||
| 126 | return ASCEND910B_DTYPE_DTYPE_SUPPORT_LIST; | 127 | return ASCEND910B_DTYPE_DTYPE_SUPPORT_LIST; |
| 127 | } else { | 128 | } else { |
| 128 | return ASCEND910_DTYPE_DTYPE_SUPPORT_LIST; | 129 | return ASCEND910_DTYPE_DTYPE_SUPPORT_LIST; |
| @@ -47,13 +47,13 @@ static constexpr int64_t DIM_FOUR = 4; | |||
| 47 | static inline const std::initializer_list<DataType>& GetAiCoreDtypeSupportListBySocVersion() | 47 | static inline const std::initializer_list<DataType>& GetAiCoreDtypeSupportListBySocVersion() |
| 48 | { | 48 | { |
| 49 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); | 49 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); |
| 50 | + if (IsRegBase(curArch)) { | ||
| 51 | + return REGBASE_AICORE_DTYPE_SUPPORT_LIST; | ||
| 52 | + } | ||
| 50 | switch (curArch) { | 53 | switch (curArch) { |
| 51 | case NpuArch::DAV_2201: { | 54 | case NpuArch::DAV_2201: { |
| 52 | return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; | 55 | return ASCEND910B_AICORE_DTYPE_SUPPORT_LIST; |
| 53 | } | 56 | } |
| 54 | - case NpuArch::DAV_3510: { | ||
| 55 | - return REGBASE_AICORE_DTYPE_SUPPORT_LIST; | ||
| 56 | - } | ||
| 57 | case NpuArch::DAV_1001: { | 57 | case NpuArch::DAV_1001: { |
| 58 | return ASCEND910_AICORE_DTYPE_SUPPORT_LIST; | 58 | return ASCEND910_AICORE_DTYPE_SUPPORT_LIST; |
| 59 | } | 59 | } |
| @@ -24,6 +24,7 @@ | |||
| 24 | 24 | ||
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | + | ||
| 27 | 28 | ||
| 28 | using namespace op; | 29 | using namespace op; |
| 29 | 30 | ||
| @@ -66,8 +67,9 @@ inline static bool CheckNotNull(const aclTensor* self, const aclTensor* other, c | |||
| 66 | 67 | ||
| 67 | inline static bool CheckSocVersionIsSupportBf16(void) | 68 | inline static bool CheckSocVersionIsSupportBf16(void) |
| 68 | { | 69 | { |
| 69 | - return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 70 | + return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 70 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E; | 71 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 72 | + IsRegBase(); | ||
| 71 | } | 73 | } |
| 72 | 74 | ||
| 73 | // 检查输入的数据类型是否在算子的支持列表内 | 75 | // 检查输入的数据类型是否在算子的支持列表内 |
| @@ -132,8 +132,9 @@ static bool CheckPowScalarTensorNotNull(const aclScalar* self, const aclTensor* | |||
| 132 | 132 | ||
| 133 | static inline bool CheckSocVersionIsSupportBf16(void) | 133 | static inline bool CheckSocVersionIsSupportBf16(void) |
| 134 | { | 134 | { |
| 135 | - return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 135 | + return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 136 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E; | 136 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 137 | + IsRegBase(); | ||
| 137 | } | 138 | } |
| 138 | 139 | ||
| 139 | // 判断910B芯片上,pow是否走AICPU路径 | 140 | // 判断910B芯片上,pow是否走AICPU路径 |
| @@ -22,6 +22,7 @@ | |||
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | 24 | ||
| 25 | + | ||
| 25 | 26 | ||
| 26 | using namespace op; | 27 | using namespace op; |
| 27 | 28 | ||
| @@ -49,174 +50,181 @@ constexpr size_t MAX_DIM_LEN = 8; | |||
| 49 | 50 | ||
| 50 | // 根据API定义,需要列出所能支持的所有dtype | 51 | // 根据API定义,需要列出所能支持的所有dtype |
| 51 | static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST = { | 52 | static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST = { |
| 52 | - op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, | 53 | + op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, |
| 53 | - op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, op::DataType::DT_BOOL, | 54 | + op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, op::DataType::DT_BOOL, |
| 54 | - op::DataType::DT_COMPLEX64, op::DataType::DT_COMPLEX128, op::DataType::DT_INT16, op::DataType::DT_BF16}; | 55 | + op::DataType::DT_COMPLEX64, op::DataType::DT_COMPLEX128, op::DataType::DT_INT16, op::DataType::DT_BF16}; |
| 55 | 56 | ||
| 56 | -static bool CheckNotNull(const aclTensor *self, const aclTensor *exponent, const aclTensor *out) { | 57 | +static bool CheckNotNull(const aclTensor* self, const aclTensor* exponent, const aclTensor* out) |
| 57 | - OP_CHECK_NULL(self, return false); | 58 | +{ |
| 58 | - OP_CHECK_NULL(exponent, return false); | 59 | + OP_CHECK_NULL(self, return false); |
| 59 | - OP_CHECK_NULL(out, return false); | 60 | + OP_CHECK_NULL(exponent, return false); |
| 60 | - return true; | 61 | + OP_CHECK_NULL(out, return false); |
| 62 | + return true; | ||
| 61 | } | 63 | } |
| 62 | 64 | ||
| 63 | -static inline bool CheckSocVersionIsSupportBf16(void) { | 65 | +static inline bool CheckSocVersionIsSupportBf16(void) |
| 64 | - return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 66 | +{ |
| 65 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E; | 67 | + return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 68 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || | ||
| 69 | + IsRegBase(); | ||
| 66 | } | 70 | } |
| 67 | 71 | ||
| 68 | -static bool CheckDtypeValid(const aclTensor *self, const aclTensor *exponent) { | 72 | +static bool CheckDtypeValid(const aclTensor* self, const aclTensor* exponent) |
| 69 | - if (!CheckSocVersionIsSupportBf16() && | 73 | +{ |
| 70 | - (self->GetDataType() == op::DataType::DT_BF16 || exponent->GetDataType() == op::DataType::DT_BF16)) { | 74 | + if (!CheckSocVersionIsSupportBf16() && |
| 71 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Input dtype of pow is not support bfloat16 in current socversion."); | 75 | + (self->GetDataType() == op::DataType::DT_BF16 || exponent->GetDataType() == op::DataType::DT_BF16)) { |
| 72 | - return false; | 76 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Input dtype of pow is not support bfloat16 in current socversion."); |
| 73 | - } | 77 | + return false; |
| 74 | - if ((self->GetDataType() == op::DataType::DT_BOOL) && (exponent->GetDataType() == op::DataType::DT_BOOL)) { | 78 | + } |
| 75 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self and exponent dtype are bool is not supported."); | 79 | + if ((self->GetDataType() == op::DataType::DT_BOOL) && (exponent->GetDataType() == op::DataType::DT_BOOL)) { |
| 76 | - return false; | 80 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self and exponent dtype are bool is not supported."); |
| 77 | - } | 81 | + return false; |
| 82 | + } | ||
| 78 | 83 | ||
| 79 | - OP_CHECK_DTYPE_NOT_SUPPORT(self, DTYPE_SUPPORT_LIST, return false); | 84 | + OP_CHECK_DTYPE_NOT_SUPPORT(self, DTYPE_SUPPORT_LIST, return false); |
| 80 | - OP_CHECK_DTYPE_NOT_SUPPORT(exponent, DTYPE_SUPPORT_LIST, return false); | 85 | + OP_CHECK_DTYPE_NOT_SUPPORT(exponent, DTYPE_SUPPORT_LIST, return false); |
| 81 | - return true; | 86 | + return true; |
| 82 | } | 87 | } |
| 83 | 88 | ||
| 84 | -static bool CheckPromoteType(const aclTensor *self, const aclTensor *exponent, const aclTensor *out) { | 89 | +static bool CheckPromoteType(const aclTensor* self, const aclTensor* exponent, const aclTensor* out) |
| 85 | - // 检查self和exponent能否做数据类型推导 | 90 | +{ |
| 86 | - op::DataType promoteType = op::PromoteType(self->GetDataType(), exponent->GetDataType()); | 91 | + // 检查self和exponent能否做数据类型推导 |
| 87 | - if (promoteType == DataType::DT_UNDEFINED) { | 92 | + op::DataType promoteType = op::PromoteType(self->GetDataType(), exponent->GetDataType()); |
| 88 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "self dtype %s and exponent dtype %s can not promote dtype.", | 93 | + if (promoteType == DataType::DT_UNDEFINED) { |
| 89 | - op::ToString(self->GetDataType()).GetString(), op::ToString(exponent->GetDataType()).GetString()); | 94 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "self dtype %s and exponent dtype %s can not promote dtype.", |
| 90 | - return false; | 95 | + op::ToString(self->GetDataType()).GetString(), op::ToString(exponent->GetDataType()).GetString()); |
| 91 | - } | 96 | + return false; |
| 97 | + } | ||
| 92 | 98 | ||
| 93 | - OP_CHECK_RESULT_DTYPE_CAST_FAILED(promoteType, out->GetDataType(), return false); | 99 | + OP_CHECK_RESULT_DTYPE_CAST_FAILED(promoteType, out->GetDataType(), return false); |
| 94 | - return true; | 100 | + return true; |
| 95 | } | 101 | } |
| 96 | 102 | ||
| 97 | -static bool CheckShape(const aclTensor *self, const aclTensor *exponent, const aclTensor *out) { | 103 | +static bool CheckShape(const aclTensor* self, const aclTensor* exponent, const aclTensor* out) |
| 104 | +{ | ||
| 105 | + OP_CHECK_MAX_DIM(self, MAX_DIM_LEN, return false); | ||
| 106 | + OP_CHECK_MAX_DIM(exponent, MAX_DIM_LEN, return false); | ||
| 98 | 107 | ||
| 99 | - OP_CHECK_MAX_DIM(self, MAX_DIM_LEN, return false); | 108 | + op::Shape broadcastShape; |
| 100 | - OP_CHECK_MAX_DIM(exponent, MAX_DIM_LEN, return false); | 109 | + OP_CHECK_BROADCAST_AND_INFER_SHAPE(self, exponent, broadcastShape, return false); |
| 101 | 110 | ||
| 102 | - op::Shape broadcastShape; | 111 | + if (broadcastShape != out->GetViewShape()) { |
| 103 | - OP_CHECK_BROADCAST_AND_INFER_SHAPE(self, exponent, broadcastShape, return false); | 112 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Shape of out should be %s, but current is %s.", |
| 104 | - | 113 | + op::ToString(broadcastShape).GetString(), op::ToString(out->GetViewShape()).GetString()); |
| 105 | - if (broadcastShape != out->GetViewShape()) { | 114 | + return false; |
| 106 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Shape of out should be %s, but current is %s.", | 115 | + } |
| 107 | - op::ToString(broadcastShape).GetString(), op::ToString(out->GetViewShape()).GetString()); | 116 | + return true; |
| 108 | - return false; | ||
| 109 | - } | ||
| 110 | - return true; | ||
| 111 | } | 117 | } |
| 112 | 118 | ||
| 113 | -static void CheckFormat(const aclTensor* self, const aclTensor* exponent) { | 119 | +static void CheckFormat(const aclTensor* self, const aclTensor* exponent) |
| 120 | +{ | ||
| 114 | ge::Format selfStorageFormat = self->GetStorageFormat(); | 121 | ge::Format selfStorageFormat = self->GetStorageFormat(); |
| 115 | ge::Format exponentStorageFormat = exponent->GetStorageFormat(); | 122 | ge::Format exponentStorageFormat = exponent->GetStorageFormat(); |
| 116 | - if (selfStorageFormat == ge::Format::FORMAT_FRACTAL_NZ || | 123 | + if (selfStorageFormat == ge::Format::FORMAT_FRACTAL_NZ || exponentStorageFormat == ge::Format::FORMAT_FRACTAL_NZ) { |
| 117 | - exponentStorageFormat == ge::Format::FORMAT_FRACTAL_NZ) { | ||
| 118 | OP_LOGW("aclnnPowTensorTensor doesn't support format NZ."); | 124 | OP_LOGW("aclnnPowTensorTensor doesn't support format NZ."); |
| 119 | } | 125 | } |
| 120 | } | 126 | } |
| 121 | 127 | ||
| 122 | -static aclnnStatus CheckParams(const aclTensor *self, const aclTensor *exponent, const aclTensor *out) { | 128 | +static aclnnStatus CheckParams(const aclTensor* self, const aclTensor* exponent, const aclTensor* out) |
| 123 | - // 1. 检查参数是否为空指针 | 129 | +{ |
| 124 | - CHECK_RET(CheckNotNull(self, exponent, out), ACLNN_ERR_PARAM_NULLPTR); | 130 | + // 1. 检查参数是否为空指针 |
| 131 | + CHECK_RET(CheckNotNull(self, exponent, out), ACLNN_ERR_PARAM_NULLPTR); | ||
| 125 | 132 | ||
| 126 | - // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验 | 133 | + // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验 |
| 127 | - CHECK_RET(CheckDtypeValid(self, exponent), ACLNN_ERR_PARAM_INVALID); | 134 | + CHECK_RET(CheckDtypeValid(self, exponent), ACLNN_ERR_PARAM_INVALID); |
| 128 | 135 | ||
| 129 | - // 3. 检查self和exponent能否做数据类型推导以及推导的数据类型能否转换为输出数据类型 | 136 | + // 3. 检查self和exponent能否做数据类型推导以及推导的数据类型能否转换为输出数据类型 |
| 130 | - CHECK_RET(CheckPromoteType(self, exponent, out), ACLNN_ERR_PARAM_INVALID); | 137 | + CHECK_RET(CheckPromoteType(self, exponent, out), ACLNN_ERR_PARAM_INVALID); |
| 131 | 138 | ||
| 132 | - // 4. 检查双输入是否能broadcast | 139 | + // 4. 检查双输入是否能broadcast |
| 133 | - CHECK_RET(CheckShape(self, exponent, out), ACLNN_ERR_PARAM_INVALID); | 140 | + CHECK_RET(CheckShape(self, exponent, out), ACLNN_ERR_PARAM_INVALID); |
| 134 | 141 | ||
| 135 | - return ACLNN_SUCCESS; | ||
| 136 | -} | ||
| 137 | - | ||
| 138 | -aclnnStatus aclnnPowTensorTensorGetWorkspaceSize(const aclTensor *self, const aclTensor *exponent, aclTensor *out, | ||
| 139 | - uint64_t *workspaceSize, aclOpExecutor **executor) { | ||
| 140 | - L2_DFX_PHASE_1(aclnnPowTensorTensor, DFX_IN(self, exponent), DFX_OUT(out)); | ||
| 141 | - | ||
| 142 | -// 固定写法,创建OpExecutor | ||
| 143 | - auto uniqueExecutor = CREATE_EXECUTOR(); | ||
| 144 | - CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); | ||
| 145 | - | ||
| 146 | - // 固定写法,参数检查 | ||
| 147 | - auto ret = CheckParams(self, exponent, out); | ||
| 148 | - CHECK_RET(ret == ACLNN_SUCCESS, ret); | ||
| 149 | - | ||
| 150 | - // 检查格式 | ||
| 151 | - CheckFormat(self, exponent); | ||
| 152 | - | ||
| 153 | - // pow算子的空tensor在kernel中支持,对标竞品根据算子实际情况补充 | ||
| 154 | - if (self->IsEmpty() || exponent->IsEmpty()) { | ||
| 155 | - // 根据实际支持情况补充 | ||
| 156 | - *workspaceSize = 0; | ||
| 157 | - uniqueExecutor.ReleaseTo(executor); | ||
| 158 | return ACLNN_SUCCESS; | 142 | return ACLNN_SUCCESS; |
| 159 | - } | ||
| 160 | - | ||
| 161 | - // Pow算子需要对self和exponent两个输入做隐式数据类型转换,根据具体算子语义按需调用 | ||
| 162 | - auto promoteType = op::PromoteType(self->GetDataType(), exponent->GetDataType()); | ||
| 163 | - | ||
| 164 | - // 固定写法,将输入self转换成连续的tensor | ||
| 165 | - auto selfContiguous = l0op::Contiguous(self, uniqueExecutor.get()); | ||
| 166 | - CHECK_RET(selfContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 167 | - | ||
| 168 | - // 将输入self的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 169 | - auto selfCasted = l0op::Cast(selfContiguous, promoteType, uniqueExecutor.get()); | ||
| 170 | - CHECK_RET(selfCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 171 | - | ||
| 172 | - // 固定写法,将输入exponent转换成连续的tensor | ||
| 173 | - auto exponentContiguous = l0op::Contiguous(exponent, uniqueExecutor.get()); | ||
| 174 | - CHECK_RET(exponentContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 175 | - | ||
| 176 | - // 将输入exponent的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 177 | - auto exponentCasted = l0op::Cast(exponentContiguous, promoteType, uniqueExecutor.get()); | ||
| 178 | - CHECK_RET(exponentCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 179 | - | ||
| 180 | - // 调用Pow算子kernel | ||
| 181 | - auto powOpOut = l0op::Pow(selfCasted, exponentCasted, uniqueExecutor.get()); | ||
| 182 | - CHECK_RET(powOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 183 | - | ||
| 184 | - // 固定写法,将计算结果转换成输出out的数据类型 | ||
| 185 | - auto castOut = l0op::Cast(powOpOut, out->GetDataType(), uniqueExecutor.get()); | ||
| 186 | - CHECK_RET(castOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 187 | - | ||
| 188 | - // 固定写法,将计算结果拷贝到输出out上,out可能是非连续的tensor | ||
| 189 | - auto viewCopyResult = l0op::ViewCopy(castOut, out, uniqueExecutor.get()); | ||
| 190 | - CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 191 | - | ||
| 192 | - // 固定写法,获取计算过程中需要使用的workspace大小 | ||
| 193 | - *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | ||
| 194 | - uniqueExecutor.ReleaseTo(executor); // 需要把 uniqueExecutor持有executor转移给executor | ||
| 195 | - return ACLNN_SUCCESS; | ||
| 196 | } | 143 | } |
| 197 | 144 | ||
| 198 | -aclnnStatus aclnnInplacePowTensorTensorGetWorkspaceSize(const aclTensor *self, | 145 | +aclnnStatus aclnnPowTensorTensorGetWorkspaceSize(const aclTensor* self, const aclTensor* exponent, aclTensor* out, |
| 199 | - const aclTensor *exponent, | 146 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 200 | - uint64_t *workspaceSize, | 147 | +{ |
| 201 | - aclOpExecutor **executor) { | 148 | + L2_DFX_PHASE_1(aclnnPowTensorTensor, DFX_IN(self, exponent), DFX_OUT(out)); |
| 202 | - auto out = const_cast<aclTensor*>(self); | 149 | + |
| 203 | - return aclnnPowTensorTensorGetWorkspaceSize(self, exponent, out, workspaceSize, executor); | 150 | + // 固定写法,创建OpExecutor |
| 151 | + auto uniqueExecutor = CREATE_EXECUTOR(); | ||
| 152 | + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); | ||
| 153 | + | ||
| 154 | + // 固定写法,参数检查 | ||
| 155 | + auto ret = CheckParams(self, exponent, out); | ||
| 156 | + CHECK_RET(ret == ACLNN_SUCCESS, ret); | ||
| 157 | + | ||
| 158 | + // 检查格式 | ||
| 159 | + CheckFormat(self, exponent); | ||
| 160 | + | ||
| 161 | + // pow算子的空tensor在kernel中支持,对标竞品根据算子实际情况补充 | ||
| 162 | + if (self->IsEmpty() || exponent->IsEmpty()) { | ||
| 163 | + // 根据实际支持情况补充 | ||
| 164 | + *workspaceSize = 0; | ||
| 165 | + uniqueExecutor.ReleaseTo(executor); | ||
| 166 | + return ACLNN_SUCCESS; | ||
| 167 | + } | ||
| 168 | + | ||
| 169 | + // Pow算子需要对self和exponent两个输入做隐式数据类型转换,根据具体算子语义按需调用 | ||
| 170 | + auto promoteType = op::PromoteType(self->GetDataType(), exponent->GetDataType()); | ||
| 171 | + | ||
| 172 | + // 固定写法,将输入self转换成连续的tensor | ||
| 173 | + auto selfContiguous = l0op::Contiguous(self, uniqueExecutor.get()); | ||
| 174 | + CHECK_RET(selfContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 175 | + | ||
| 176 | + // 将输入self的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 177 | + auto selfCasted = l0op::Cast(selfContiguous, promoteType, uniqueExecutor.get()); | ||
| 178 | + CHECK_RET(selfCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 179 | + | ||
| 180 | + // 固定写法,将输入exponent转换成连续的tensor | ||
| 181 | + auto exponentContiguous = l0op::Contiguous(exponent, uniqueExecutor.get()); | ||
| 182 | + CHECK_RET(exponentContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 183 | + | ||
| 184 | + // 将输入exponent的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 185 | + auto exponentCasted = l0op::Cast(exponentContiguous, promoteType, uniqueExecutor.get()); | ||
| 186 | + CHECK_RET(exponentCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 187 | + | ||
| 188 | + // 调用Pow算子kernel | ||
| 189 | + auto powOpOut = l0op::Pow(selfCasted, exponentCasted, uniqueExecutor.get()); | ||
| 190 | + CHECK_RET(powOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 191 | + | ||
| 192 | + // 固定写法,将计算结果转换成输出out的数据类型 | ||
| 193 | + auto castOut = l0op::Cast(powOpOut, out->GetDataType(), uniqueExecutor.get()); | ||
| 194 | + CHECK_RET(castOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 195 | + | ||
| 196 | + // 固定写法,将计算结果拷贝到输出out上,out可能是非连续的tensor | ||
| 197 | + auto viewCopyResult = l0op::ViewCopy(castOut, out, uniqueExecutor.get()); | ||
| 198 | + CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 199 | + | ||
| 200 | + // 固定写法,获取计算过程中需要使用的workspace大小 | ||
| 201 | + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | ||
| 202 | + uniqueExecutor.ReleaseTo(executor); // 需要把 uniqueExecutor持有executor转移给executor | ||
| 203 | + return ACLNN_SUCCESS; | ||
| 204 | } | 204 | } |
| 205 | 205 | ||
| 206 | -aclnnStatus aclnnInplacePowTensorTensor(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, | 206 | +aclnnStatus aclnnInplacePowTensorTensorGetWorkspaceSize(const aclTensor* self, const aclTensor* exponent, |
| 207 | - aclrtStream stream) { | 207 | + uint64_t* workspaceSize, aclOpExecutor** executor) |
| 208 | - L2_DFX_PHASE_2(aclnnInplacePowTensorTensor); | 208 | +{ |
| 209 | - // 固定写法,调用框架能力,完成计算 | 209 | + auto out = const_cast<aclTensor*>(self); |
| 210 | - return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | 210 | + return aclnnPowTensorTensorGetWorkspaceSize(self, exponent, out, workspaceSize, executor); |
| 211 | } | 211 | } |
| 212 | 212 | ||
| 213 | -aclnnStatus aclnnPowTensorTensor(void *workspace, uint64_t workspaceSize, | 213 | +aclnnStatus aclnnInplacePowTensorTensor(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, |
| 214 | - aclOpExecutor *executor, aclrtStream stream) { | 214 | + aclrtStream stream) |
| 215 | - L2_DFX_PHASE_2(aclnnPowTensorTensor); | 215 | +{ |
| 216 | - // 固定写法,调用框架能力,完成计算 | 216 | + L2_DFX_PHASE_2(aclnnInplacePowTensorTensor); |
| 217 | - return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | 217 | + // 固定写法,调用框架能力,完成计算 |
| 218 | + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | ||
| 219 | +} | ||
| 220 | + | ||
| 221 | +aclnnStatus aclnnPowTensorTensor(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream) | ||
| 222 | +{ | ||
| 223 | + L2_DFX_PHASE_2(aclnnPowTensorTensor); | ||
| 224 | + // 固定写法,调用框架能力,完成计算 | ||
| 225 | + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | ||
| 218 | } | 226 | } |
| 219 | 227 | ||
| 220 | 228 | ||
| 221 | } | 229 | } |
| 222 | -#endif | 230 | +#endif |
| @@ -47,9 +47,9 @@ static const std::initializer_list<op::DataType> ASCEND910B_DTYPE_SUPPORT_LIST = | |||
| 47 | op::DataType::DT_BF16}; | 47 | op::DataType::DT_BF16}; |
| 48 | 48 | ||
| 49 | static const std::initializer_list<op::DataType> SIGNBIT_DTYPE_SUPPORT_LIST = { | 49 | static const std::initializer_list<op::DataType> SIGNBIT_DTYPE_SUPPORT_LIST = { |
| 50 | - op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, | 50 | + op::DataType::DT_FLOAT, op::DataType::DT_INT32, op::DataType::DT_INT64, op::DataType::DT_FLOAT16, |
| 51 | - op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, op::DataType::DT_UINT64, | 51 | + op::DataType::DT_INT8, op::DataType::DT_UINT8, op::DataType::DT_DOUBLE, op::DataType::DT_UINT64, |
| 52 | - op::DataType::DT_BOOL, op::DataType::DT_BF16}; | 52 | + op::DataType::DT_BOOL, op::DataType::DT_BF16}; |
| 53 | 53 | ||
| 54 | static bool CanUseSignbit(const aclTensor* self) | 54 | static bool CanUseSignbit(const aclTensor* self) |
| 55 | { | 55 | { |
| @@ -65,8 +65,9 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* out) | |||
| 65 | 65 | ||
| 66 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() | 66 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportList() |
| 67 | { | 67 | { |
| 68 | - if (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 68 | + if ((GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 69 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) { | 69 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 70 | + IsRegBase()) { | ||
| 70 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 71 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 71 | } else { | 72 | } else { |
| 72 | return ASCEND910_DTYPE_SUPPORT_LIST; | 73 | return ASCEND910_DTYPE_SUPPORT_LIST; |
| @@ -137,8 +138,8 @@ static aclnnStatus FillScalar(aclTensor* out, bool val, aclOpExecutor* executor) | |||
| 137 | return ACLNN_SUCCESS; | 138 | return ACLNN_SUCCESS; |
| 138 | } | 139 | } |
| 139 | 140 | ||
| 140 | -aclnnStatus aclnnSignbitGetWorkspaceSize( | 141 | +aclnnStatus aclnnSignbitGetWorkspaceSize(const aclTensor* self, aclTensor* out, uint64_t* workspaceSize, |
| 141 | - const aclTensor* self, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor) | 142 | + aclOpExecutor** executor) |
| 142 | { | 143 | { |
| 143 | L2_DFX_PHASE_1(aclnnSignbit, DFX_IN(self), DFX_OUT(out)); | 144 | L2_DFX_PHASE_1(aclnnSignbit, DFX_IN(self), DFX_OUT(out)); |
| 144 | 145 | ||
| @@ -174,7 +175,8 @@ aclnnStatus aclnnSignbitGetWorkspaceSize( | |||
| 174 | } else { | 175 | } else { |
| 175 | // 创建数据为0的tensor | 176 | // 创建数据为0的tensor |
| 176 | FVector<float> zeroVector = {0}; | 177 | FVector<float> zeroVector = {0}; |
| 177 | - auto zeroTensor = uniqueExecutor.get()->ConvertToTensor(zeroVector.data(), zeroVector.size(), self->GetDataType()); | 178 | + auto zeroTensor = uniqueExecutor.get()->ConvertToTensor(zeroVector.data(), zeroVector.size(), |
| 179 | + self->GetDataType()); | ||
| 178 | CHECK_RET(zeroTensor != nullptr, ACLNN_ERR_INNER_NULLPTR); | 180 | CHECK_RET(zeroTensor != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 179 | signBitOpOut = l0op::Less(selfContiguous, zeroTensor, uniqueExecutor.get()); | 181 | signBitOpOut = l0op::Less(selfContiguous, zeroTensor, uniqueExecutor.get()); |
| 180 | } | 182 | } |
| @@ -114,9 +114,11 @@ static op::DataType CombineCategoriesWithComplex(const op::DataType higher, cons | |||
| 114 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion() | 114 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion() |
| 115 | { | 115 | { |
| 116 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); | 116 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); |
| 117 | + if (IsRegBase(curArch)) { | ||
| 118 | + return ASCEND910B_DTYPE_SUPPORT_LIST; | ||
| 119 | + } | ||
| 117 | switch (curArch) { | 120 | switch (curArch) { |
| 118 | - case NpuArch::DAV_2201: | 121 | + case NpuArch::DAV_2201: { |
| 119 | - case NpuArch::DAV_3510: { | ||
| 120 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 122 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 121 | } | 123 | } |
| 122 | default: { | 124 | default: { |
| @@ -140,9 +140,11 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* other, const ac | |||
| 140 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion() | 140 | static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion() |
| 141 | { | 141 | { |
| 142 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); | 142 | auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); |
| 143 | + if (IsRegBase(curArch)) { | ||
| 144 | + return ASCEND910B_DTYPE_SUPPORT_LIST; | ||
| 145 | + } | ||
| 143 | switch (curArch) { | 146 | switch (curArch) { |
| 144 | - case NpuArch::DAV_2201: | 147 | + case NpuArch::DAV_2201: { |
| 145 | - case NpuArch::DAV_3510: { | ||
| 146 | return ASCEND910B_DTYPE_SUPPORT_LIST; | 148 | return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 147 | } | 149 | } |
| 148 | case NpuArch::DAV_1001: { | 150 | case NpuArch::DAV_1001: { |
| @@ -19,6 +19,7 @@ | |||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | + | ||
| 22 | 23 | ||
| 23 | using namespace op; | 24 | using namespace op; |
| 24 | 25 | ||
| @@ -42,8 +43,9 @@ static const std::initializer_list<op::DataType> ASCEND610LITE_DTYPE_SUPPORT_LIS | |||
| 42 | static bool IsAiCoreSupport(const aclTensor* self) | 43 | static bool IsAiCoreSupport(const aclTensor* self) |
| 43 | { | 44 | { |
| 44 | // 获取芯片类型,判断是1971还是1980 | 45 | // 获取芯片类型,判断是1971还是1980 |
| 45 | - if (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 46 | + if ((GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 46 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) { | 47 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || |
| 48 | + IsRegBase()) { | ||
| 47 | return CheckType(self->GetDataType(), ASCEND910B_AICORE_DTYPE_SUPPORT_LIST); | 49 | return CheckType(self->GetDataType(), ASCEND910B_AICORE_DTYPE_SUPPORT_LIST); |
| 48 | } | 50 | } |
| 49 | 51 | ||
| @@ -55,8 +57,8 @@ static bool IsAiCoreSupport(const aclTensor* self) | |||
| 55 | } | 57 | } |
| 56 | 58 | ||
| 57 | // AICORE算子kernel | 59 | // AICORE算子kernel |
| 58 | -static const aclTensor* SubAiCore( | 60 | +static const aclTensor* SubAiCore(const aclTensor* self, const aclTensor* other, aclTensor* subOut, |
| 59 | - const aclTensor* self, const aclTensor* other, aclTensor* subOut, aclOpExecutor* executor) | 61 | + aclOpExecutor* executor) |
| 60 | { | 62 | { |
| 61 | L0_DFX(SubAiCore, self, other, subOut); | 63 | L0_DFX(SubAiCore, self, other, subOut); |
| 62 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICORE,将AiCore Sub算子加入任务队列 | 64 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICORE,将AiCore Sub算子加入任务队列 |
| @@ -66,8 +68,8 @@ static const aclTensor* SubAiCore( | |||
| 66 | } | 68 | } |
| 67 | 69 | ||
| 68 | // AICPU算子kernel | 70 | // AICPU算子kernel |
| 69 | -static const aclTensor* SubAiCpu( | 71 | +static const aclTensor* SubAiCpu(const aclTensor* self, const aclTensor* other, aclTensor* subOut, |
| 70 | - const aclTensor* self, const aclTensor* other, aclTensor* subOut, aclOpExecutor* executor) | 72 | + aclOpExecutor* executor) |
| 71 | { | 73 | { |
| 72 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICPU,将AiCpu Sub算子加入任务队列 | 74 | // 使用框架宏ADD_TO_LAUNCHER_LIST_AICPU,将AiCpu Sub算子加入任务队列 |
| 73 | // Sub是算子的OpType,self、other是算子的输入,subOut是算子的输出 | 75 | // Sub是算子的OpType,self、other是算子的输入,subOut是算子的输出 |
| @@ -84,9 +86,8 @@ const aclTensor* Sub(const aclTensor* self, const aclTensor* other, aclOpExecuto | |||
| 84 | // 通过输入shape推导算子输出shape | 86 | // 通过输入shape推导算子输出shape |
| 85 | op::Shape broadcastShape; | 87 | op::Shape broadcastShape; |
| 86 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { | 88 | if (!BroadcastInferShape(self->GetViewShape(), other->GetViewShape(), broadcastShape)) { |
| 87 | - OP_LOGE( | 89 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), |
| 88 | - ACLNN_ERR_PARAM_INVALID, "Broadcast %s and %s failed.", op::ToString(self->GetViewShape()).GetString(), | 90 | + op::ToString(other->GetViewShape()).GetString()); |
| 89 | - op::ToString(other->GetViewShape()).GetString()); | ||
| 90 | return nullptr; | 91 | return nullptr; |
| 91 | } | 92 | } |
| 92 | 93 | ||
| @@ -99,4 +100,4 @@ const aclTensor* Sub(const aclTensor* self, const aclTensor* other, aclOpExecuto | |||
| 99 | return SubAiCpu(self, other, subOut, executor); | 100 | return SubAiCpu(self, other, subOut, executor); |
| 100 | } | 101 | } |
| 101 | } | 102 | } |
| 102 | -} // namespace l0op | 103 | +} // namespace l0op |
| @@ -40,131 +40,139 @@ static const std::initializer_list<op::DataType> ASCEND910_DTYPE_SUPPORT_LIST = | |||
| 40 | static const std::initializer_list<op::DataType> ASCEND910B_DTYPE_SUPPORT_LIST = { | 40 | static const std::initializer_list<op::DataType> ASCEND910B_DTYPE_SUPPORT_LIST = { |
| 41 | op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_DOUBLE, op::DataType::DT_BF16}; | 41 | op::DataType::DT_FLOAT, op::DataType::DT_FLOAT16, op::DataType::DT_DOUBLE, op::DataType::DT_BF16}; |
| 42 | 42 | ||
| 43 | -static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion() { | 43 | +static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion() |
| 44 | - auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); | 44 | +{ |
| 45 | - switch (curArch) { | 45 | + auto curArch = GetCurrentPlatformInfo().GetCurNpuArch(); |
| 46 | - case NpuArch::DAV_2201: | 46 | + if (IsRegBase(curArch)) { |
| 47 | - case NpuArch::DAV_3510: { | 47 | + return ASCEND910B_DTYPE_SUPPORT_LIST; |
| 48 | - return ASCEND910B_DTYPE_SUPPORT_LIST; | ||
| 49 | } | 48 | } |
| 50 | - case NpuArch::DAV_1001: { | 49 | + switch (curArch) { |
| 51 | - return ASCEND910_DTYPE_SUPPORT_LIST; | 50 | + case NpuArch::DAV_2201: { |
| 51 | + return ASCEND910B_DTYPE_SUPPORT_LIST; | ||
| 52 | + } | ||
| 53 | + case NpuArch::DAV_1001: { | ||
| 54 | + return ASCEND910_DTYPE_SUPPORT_LIST; | ||
| 55 | + } | ||
| 56 | + default: { | ||
| 57 | + return ASCEND910_DTYPE_SUPPORT_LIST; | ||
| 58 | + } | ||
| 52 | } | 59 | } |
| 53 | - default: { | ||
| 54 | - return ASCEND910_DTYPE_SUPPORT_LIST; | ||
| 55 | - } | ||
| 56 | - } | ||
| 57 | } | 60 | } |
| 58 | 61 | ||
| 59 | -static bool CheckNotNull(const aclTensor *gradOutput, const aclTensor *output, const aclTensor *gradInput) { | 62 | +static bool CheckNotNull(const aclTensor* gradOutput, const aclTensor* output, const aclTensor* gradInput) |
| 60 | - OP_CHECK_NULL(gradOutput, return false); | 63 | +{ |
| 61 | - OP_CHECK_NULL(output, return false); | 64 | + OP_CHECK_NULL(gradOutput, return false); |
| 62 | - OP_CHECK_NULL(gradInput, return false); | 65 | + OP_CHECK_NULL(output, return false); |
| 63 | - return true; | 66 | + OP_CHECK_NULL(gradInput, return false); |
| 67 | + return true; | ||
| 64 | } | 68 | } |
| 65 | 69 | ||
| 66 | -static bool CheckDtypeValid(const aclTensor *gradOutput, const aclTensor *output, const aclTensor *gradInput) { | 70 | +static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output, const aclTensor* gradInput) |
| 67 | - const std::initializer_list<op::DataType> dtypeSupportList = GetDtypeSupportListBySocVersion(); | 71 | +{ |
| 68 | - // 检查gradOutput的数据类型是否在TanhGrad算子的支持列表内 | 72 | + const std::initializer_list<op::DataType> dtypeSupportList = GetDtypeSupportListBySocVersion(); |
| 69 | - OP_CHECK_DTYPE_NOT_SUPPORT(gradOutput, dtypeSupportList, return false); | 73 | + // 检查gradOutput的数据类型是否在TanhGrad算子的支持列表内 |
| 74 | + OP_CHECK_DTYPE_NOT_SUPPORT(gradOutput, dtypeSupportList, return false); | ||
| 70 | 75 | ||
| 71 | - // 检查output的数据类型是否在TanhGrad算子的支持列表内 | 76 | + // 检查output的数据类型是否在TanhGrad算子的支持列表内 |
| 72 | - OP_CHECK_DTYPE_NOT_SUPPORT(output, dtypeSupportList, return false); | 77 | + OP_CHECK_DTYPE_NOT_SUPPORT(output, dtypeSupportList, return false); |
| 73 | 78 | ||
| 74 | - // 检查gradInput的数据类型是否在支持列表内 | 79 | + // 检查gradInput的数据类型是否在支持列表内 |
| 75 | - OP_CHECK_DTYPE_NOT_SUPPORT(gradInput, dtypeSupportList, return false); | 80 | + OP_CHECK_DTYPE_NOT_SUPPORT(gradInput, dtypeSupportList, return false); |
| 76 | 81 | ||
| 77 | - // 检查gradOutput和output/gradInput是否dtype一致 | 82 | + // 检查gradOutput和output/gradInput是否dtype一致 |
| 78 | - OP_CHECK_DTYPE_NOT_MATCH(gradOutput, output->GetDataType(), return false); | 83 | + OP_CHECK_DTYPE_NOT_MATCH(gradOutput, output->GetDataType(), return false); |
| 79 | - OP_CHECK_DTYPE_NOT_MATCH(gradOutput, gradInput->GetDataType(), return false); | 84 | + OP_CHECK_DTYPE_NOT_MATCH(gradOutput, gradInput->GetDataType(), return false); |
| 80 | 85 | ||
| 81 | - return true; | 86 | + return true; |
| 82 | } | 87 | } |
| 83 | 88 | ||
| 84 | -static bool CheckShapeValid(const aclTensor *gradOutput, const aclTensor *output, const aclTensor *gradInput) { | 89 | +static bool CheckShapeValid(const aclTensor* gradOutput, const aclTensor* output, const aclTensor* gradInput) |
| 85 | - OP_CHECK_MAX_DIM(gradOutput, MAX_SUPPORT_DIMS_NUMS, return false); | 90 | +{ |
| 86 | - OP_CHECK_MAX_DIM(output, MAX_SUPPORT_DIMS_NUMS, return false); | 91 | + OP_CHECK_MAX_DIM(gradOutput, MAX_SUPPORT_DIMS_NUMS, return false); |
| 92 | + OP_CHECK_MAX_DIM(output, MAX_SUPPORT_DIMS_NUMS, return false); | ||
| 87 | 93 | ||
| 88 | - // 检查gradOutput和output broadcast之后与gradInput是否shape一致 | 94 | + // 检查gradOutput和output broadcast之后与gradInput是否shape一致 |
| 89 | - Shape dstShape; | 95 | + Shape dstShape; |
| 90 | - OP_CHECK_BROADCAST_AND_INFER_SHAPE(gradOutput, output, dstShape, return false); | 96 | + OP_CHECK_BROADCAST_AND_INFER_SHAPE(gradOutput, output, dstShape, return false); |
| 91 | - OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gradInput, dstShape, return false); | 97 | + OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(gradInput, dstShape, return false); |
| 92 | 98 | ||
| 93 | - return true; | 99 | + return true; |
| 94 | } | 100 | } |
| 95 | 101 | ||
| 96 | -static aclnnStatus CheckParams(const aclTensor *gradOutput, const aclTensor *output, const aclTensor *gradInput) { | 102 | +static aclnnStatus CheckParams(const aclTensor* gradOutput, const aclTensor* output, const aclTensor* gradInput) |
| 97 | - // 1. 检查参数是否为空指针 | 103 | +{ |
| 98 | - CHECK_RET(CheckNotNull(gradOutput, output, gradInput), ACLNN_ERR_PARAM_NULLPTR); | 104 | + // 1. 检查参数是否为空指针 |
| 105 | + CHECK_RET(CheckNotNull(gradOutput, output, gradInput), ACLNN_ERR_PARAM_NULLPTR); | ||
| 99 | 106 | ||
| 100 | - // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验 | 107 | + // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验 |
| 101 | - CHECK_RET(CheckDtypeValid(gradOutput, output, gradInput), ACLNN_ERR_PARAM_INVALID); | 108 | + CHECK_RET(CheckDtypeValid(gradOutput, output, gradInput), ACLNN_ERR_PARAM_INVALID); |
| 102 | 109 | ||
| 103 | - // 3. 检查是否shape一致 | 110 | + // 3. 检查是否shape一致 |
| 104 | - CHECK_RET(CheckShapeValid(gradOutput, output, gradInput), ACLNN_ERR_PARAM_INVALID); | 111 | + CHECK_RET(CheckShapeValid(gradOutput, output, gradInput), ACLNN_ERR_PARAM_INVALID); |
| 105 | 112 | ||
| 106 | - return ACLNN_SUCCESS; | ||
| 107 | -} | ||
| 108 | - | ||
| 109 | -aclnnStatus aclnnTanhBackwardGetWorkspaceSize(const aclTensor *gradOutput, const aclTensor *output, | ||
| 110 | - aclTensor *gradInput, uint64_t *workspaceSize, | ||
| 111 | - aclOpExecutor **executor) { | ||
| 112 | - L2_DFX_PHASE_1(aclnnTanhBackward, DFX_IN(gradOutput, output), DFX_OUT(gradInput)); | ||
| 113 | - | ||
| 114 | - // 固定写法,创建OpExecutor | ||
| 115 | - auto uniqueExecutor = CREATE_EXECUTOR(); | ||
| 116 | - CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); | ||
| 117 | - | ||
| 118 | - // 固定写法,参数检查 | ||
| 119 | - auto ret = CheckParams(gradOutput, output, gradInput); | ||
| 120 | - CHECK_RET(ret == ACLNN_SUCCESS, ret); | ||
| 121 | - | ||
| 122 | - // TanhGrad算子的空tensor在kernel中支持 | ||
| 123 | - if (gradOutput->IsEmpty() || output->IsEmpty()) { | ||
| 124 | - // 根据实际支持情况补充 | ||
| 125 | - *workspaceSize = 0; | ||
| 126 | - uniqueExecutor.ReleaseTo(executor); | ||
| 127 | return ACLNN_SUCCESS; | 113 | return ACLNN_SUCCESS; |
| 128 | - } | ||
| 129 | - | ||
| 130 | - // 固定写法,将输入gradOutputContiguous转换成连续的tensor | ||
| 131 | - auto gradOutputContiguous = l0op::Contiguous(gradOutput, uniqueExecutor.get()); | ||
| 132 | - CHECK_RET(gradOutputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 133 | - | ||
| 134 | - // 固定写法,将输入output转换成连续的tensor | ||
| 135 | - auto outputContiguous = l0op::Contiguous(output, uniqueExecutor.get()); | ||
| 136 | - CHECK_RET(outputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 137 | - | ||
| 138 | - const aclTensor* tanhBackwardOpOut = nullptr; | ||
| 139 | - if (IsRegBase()) { | ||
| 140 | - // 调用TanhGrad算子kernel | ||
| 141 | - auto tanhBackwardOpOutBeforeCast = l0op::TanhGrad(gradOutputContiguous, outputContiguous, uniqueExecutor.get()); | ||
| 142 | - CHECK_RET(tanhBackwardOpOutBeforeCast != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 143 | - | ||
| 144 | - // 调用Cast算子kernel | ||
| 145 | - tanhBackwardOpOut = l0op::Cast(tanhBackwardOpOutBeforeCast, gradInput->GetDataType(), uniqueExecutor.get()); | ||
| 146 | - CHECK_RET(tanhBackwardOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 147 | - } else { | ||
| 148 | - // 调用TanhGrad算子kernel | ||
| 149 | - tanhBackwardOpOut = l0op::TanhGrad(gradOutputContiguous, outputContiguous, uniqueExecutor.get()); | ||
| 150 | - CHECK_RET(tanhBackwardOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 151 | - } | ||
| 152 | - | ||
| 153 | - // 固定写法,将计算结果拷贝到输出gradInput上,gradInput可能是非连续的tensor | ||
| 154 | - auto viewCopyResult = l0op::ViewCopy(tanhBackwardOpOut, gradInput, uniqueExecutor.get()); | ||
| 155 | - CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 156 | - | ||
| 157 | - // 固定写法,获取计算过程中需要使用的workspace大小 | ||
| 158 | - *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | ||
| 159 | - uniqueExecutor.ReleaseTo(executor); // 需要把 uniqueExecutor持有executor转移给executor | ||
| 160 | - return ACLNN_SUCCESS; | ||
| 161 | } | 114 | } |
| 162 | 115 | ||
| 163 | -aclnnStatus aclnnTanhBackward(void *workspace, uint64_t workspaceSize, | 116 | +aclnnStatus aclnnTanhBackwardGetWorkspaceSize(const aclTensor* gradOutput, const aclTensor* output, |
| 164 | - aclOpExecutor *executor, const aclrtStream stream) { | 117 | + aclTensor* gradInput, uint64_t* workspaceSize, aclOpExecutor** executor) |
| 165 | - L2_DFX_PHASE_2(aclnnTanhBackward); | 118 | +{ |
| 166 | - // 固定写法,调用框架能力,完成计算 | 119 | + L2_DFX_PHASE_1(aclnnTanhBackward, DFX_IN(gradOutput, output), DFX_OUT(gradInput)); |
| 167 | - return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | 120 | + |
| 121 | + // 固定写法,创建OpExecutor | ||
| 122 | + auto uniqueExecutor = CREATE_EXECUTOR(); | ||
| 123 | + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); | ||
| 124 | + | ||
| 125 | + // 固定写法,参数检查 | ||
| 126 | + auto ret = CheckParams(gradOutput, output, gradInput); | ||
| 127 | + CHECK_RET(ret == ACLNN_SUCCESS, ret); | ||
| 128 | + | ||
| 129 | + // TanhGrad算子的空tensor在kernel中支持 | ||
| 130 | + if (gradOutput->IsEmpty() || output->IsEmpty()) { | ||
| 131 | + // 根据实际支持情况补充 | ||
| 132 | + *workspaceSize = 0; | ||
| 133 | + uniqueExecutor.ReleaseTo(executor); | ||
| 134 | + return ACLNN_SUCCESS; | ||
| 135 | + } | ||
| 136 | + | ||
| 137 | + // 固定写法,将输入gradOutputContiguous转换成连续的tensor | ||
| 138 | + auto gradOutputContiguous = l0op::Contiguous(gradOutput, uniqueExecutor.get()); | ||
| 139 | + CHECK_RET(gradOutputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 140 | + | ||
| 141 | + // 固定写法,将输入output转换成连续的tensor | ||
| 142 | + auto outputContiguous = l0op::Contiguous(output, uniqueExecutor.get()); | ||
| 143 | + CHECK_RET(outputContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 144 | + | ||
| 145 | + const aclTensor* tanhBackwardOpOut = nullptr; | ||
| 146 | + if (IsRegBase()) { | ||
| 147 | + // 调用TanhGrad算子kernel | ||
| 148 | + auto tanhBackwardOpOutBeforeCast = l0op::TanhGrad(gradOutputContiguous, outputContiguous, uniqueExecutor.get()); | ||
| 149 | + CHECK_RET(tanhBackwardOpOutBeforeCast != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 150 | + | ||
| 151 | + // 调用Cast算子kernel | ||
| 152 | + tanhBackwardOpOut = l0op::Cast(tanhBackwardOpOutBeforeCast, gradInput->GetDataType(), uniqueExecutor.get()); | ||
| 153 | + CHECK_RET(tanhBackwardOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 154 | + } else { | ||
| 155 | + // 调用TanhGrad算子kernel | ||
| 156 | + tanhBackwardOpOut = l0op::TanhGrad(gradOutputContiguous, outputContiguous, uniqueExecutor.get()); | ||
| 157 | + CHECK_RET(tanhBackwardOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 158 | + } | ||
| 159 | + | ||
| 160 | + // 固定写法,将计算结果拷贝到输出gradInput上,gradInput可能是非连续的tensor | ||
| 161 | + auto viewCopyResult = l0op::ViewCopy(tanhBackwardOpOut, gradInput, uniqueExecutor.get()); | ||
| 162 | + CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 163 | + | ||
| 164 | + // 固定写法,获取计算过程中需要使用的workspace大小 | ||
| 165 | + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | ||
| 166 | + uniqueExecutor.ReleaseTo(executor); // 需要把 uniqueExecutor持有executor转移给executor | ||
| 167 | + return ACLNN_SUCCESS; | ||
| 168 | +} | ||
| 169 | + | ||
| 170 | +aclnnStatus aclnnTanhBackward(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, | ||
| 171 | + const aclrtStream stream) | ||
| 172 | +{ | ||
| 173 | + L2_DFX_PHASE_2(aclnnTanhBackward); | ||
| 174 | + // 固定写法,调用框架能力,完成计算 | ||
| 175 | + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | ||
| 168 | } | 176 | } |
| 169 | 177 | ||
| 170 | 178 | ||
| @@ -24,6 +24,7 @@ | |||
| 24 | 24 | ||
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | + | ||
| 27 | 28 | ||
| 28 | using namespace op; | 29 | using namespace op; |
| 29 | 30 | ||
| @@ -55,161 +56,170 @@ static const std::initializer_list<op::DataType> DTYPE_SUPPORT_910_LIST = { | |||
| 55 | op::DataType::DT_DOUBLE, op::DataType::DT_UINT16, op::DataType::DT_UINT32, op::DataType::DT_UINT64}; | 56 | op::DataType::DT_DOUBLE, op::DataType::DT_UINT16, op::DataType::DT_UINT32, op::DataType::DT_UINT64}; |
| 56 | 57 | ||
| 57 | // 列出out所能支持的所有dtype | 58 | // 列出out所能支持的所有dtype |
| 58 | -static const std::initializer_list<op::DataType> OUT_DTYPE_SUPPORT_LIST = { | 59 | +static const std::initializer_list<op::DataType> OUT_DTYPE_SUPPORT_LIST = {op::DataType::DT_BOOL}; |
| 59 | - op::DataType::DT_BOOL}; | ||
| 60 | 60 | ||
| 61 | // 算子支持的最大维度 | 61 | // 算子支持的最大维度 |
| 62 | static const size_t DIM_SUPPORT_MAX = 8; | 62 | static const size_t DIM_SUPPORT_MAX = 8; |
| 63 | 63 | ||
| 64 | -static bool CheckNotNull(const aclTensor *self, const aclTensor *other, const aclTensor *out) { | 64 | +static bool CheckNotNull(const aclTensor* self, const aclTensor* other, const aclTensor* out) |
| 65 | - OP_CHECK_NULL(self, return false); | 65 | +{ |
| 66 | - OP_CHECK_NULL(other, return false); | 66 | + OP_CHECK_NULL(self, return false); |
| 67 | - OP_CHECK_NULL(out, return false); | 67 | + OP_CHECK_NULL(other, return false); |
| 68 | - return true; | 68 | + OP_CHECK_NULL(out, return false); |
| 69 | + return true; | ||
| 69 | } | 70 | } |
| 70 | 71 | ||
| 71 | -static inline bool CheckSocVersionGe910B(void) { | 72 | +static inline bool CheckSocVersionGe910B(void) |
| 72 | - return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && | 73 | +{ |
| 73 | - GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E; | 74 | + return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B && |
| 75 | + GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) || | ||
| 76 | + IsRegBase(); | ||
| 74 | } | 77 | } |
| 75 | 78 | ||
| 76 | -static bool CheckDtypeValid(const aclTensor *self, const aclTensor *other, const aclTensor *out) { | 79 | +static bool CheckDtypeValid(const aclTensor* self, const aclTensor* other, const aclTensor* out) |
| 77 | - // 如果soc是1980芯片,则不支持DT_BF16,需要校验拦截,否则Cast报错 | 80 | +{ |
| 78 | - bool is910BSocVersion = CheckSocVersionGe910B(); | 81 | + // 如果soc是1980芯片,则不支持DT_BF16,需要校验拦截,否则Cast报错 |
| 79 | - const std::initializer_list<DataType> DTYPE_SUPPORT_LIST = | 82 | + bool is910BSocVersion = CheckSocVersionGe910B(); |
| 80 | - is910BSocVersion ? DTYPE_SUPPORT_910B_LIST : DTYPE_SUPPORT_910_LIST; | 83 | + const std::initializer_list<DataType> DTYPE_SUPPORT_LIST = is910BSocVersion ? DTYPE_SUPPORT_910B_LIST : |
| 84 | + DTYPE_SUPPORT_910_LIST; | ||
| 81 | 85 | ||
| 82 | - // 检查self和other能否做数据类型推导 | 86 | + // 检查self和other能否做数据类型推导 |
| 83 | - op::DataType promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); | 87 | + op::DataType promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); |
| 84 | - if (promoteType == DataType::DT_UNDEFINED) { | 88 | + if (promoteType == DataType::DT_UNDEFINED) { |
| 85 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", | 89 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s can not promote dtype.", |
| 86 | - op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); | 90 | + op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString()); |
| 87 | - return false; | 91 | + return false; |
| 88 | - } | ||
| 89 | - | ||
| 90 | - // 检查promoteType的数据类型是否在equal算子的支持列表内 | ||
| 91 | - if (!CheckType(promoteType, DTYPE_SUPPORT_LIST)) { | ||
| 92 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Self dtype %s and other dtype %s get promoteType dtype %s should be in " \ | ||
| 93 | - "dtype support list [%s].", op::ToString(self->GetDataType()).GetString(), | ||
| 94 | - op::ToString(other->GetDataType()).GetString(), op::ToString(promoteType).GetString(), | ||
| 95 | - op::ToString(DTYPE_SUPPORT_LIST).GetString()); | ||
| 96 | - return false; | ||
| 97 | - } | ||
| 98 | - | ||
| 99 | - // 检查out的数据类型是否是BOOL | ||
| 100 | - OP_CHECK_DTYPE_NOT_SUPPORT(out, OUT_DTYPE_SUPPORT_LIST, return false); | ||
| 101 | - return true; | ||
| 102 | -} | ||
| 103 | - | ||
| 104 | -static bool CheckMaxShape(const aclTensor *self, const aclTensor *other, const aclTensor *out) { | ||
| 105 | - OP_CHECK_MAX_DIM(self, DIM_SUPPORT_MAX, return false); | ||
| 106 | - OP_CHECK_MAX_DIM(other, DIM_SUPPORT_MAX, return false); | ||
| 107 | - OP_CHECK_MAX_DIM(out, DIM_SUPPORT_MAX, return false); | ||
| 108 | - return true; | ||
| 109 | -} | ||
| 110 | - | ||
| 111 | -static bool CheckOutShape(const aclTensor *out) { | ||
| 112 | - op::Shape outShape; | ||
| 113 | - outShape.SetDimNum(1); | ||
| 114 | - outShape.SetDim(0, 1); | ||
| 115 | - OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(out, outShape, return false); | ||
| 116 | - return true; | ||
| 117 | -} | ||
| 118 | - | ||
| 119 | -static aclnnStatus CheckParams(const aclTensor *self, const aclTensor *other, const aclTensor *out) { | ||
| 120 | - // 1. 检查两个入参参数是否为空指针;out为空指针时不报错,结果输出None(python中的空指针) | ||
| 121 | - CHECK_RET(CheckNotNull(self, other, out), ACLNN_ERR_PARAM_NULLPTR); | ||
| 122 | - | ||
| 123 | - // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验 | ||
| 124 | - CHECK_RET(CheckDtypeValid(self, other, out), ACLNN_ERR_PARAM_INVALID); | ||
| 125 | - | ||
| 126 | - // 3. 入参tensor最大维度检查 | ||
| 127 | - CHECK_RET(CheckMaxShape(self, other, out), ACLNN_ERR_PARAM_INVALID); | ||
| 128 | - | ||
| 129 | - // 4. 输出tensor形状检查 | ||
| 130 | - CHECK_RET(CheckOutShape(out), ACLNN_ERR_PARAM_INVALID); | ||
| 131 | - | ||
| 132 | - return ACLNN_SUCCESS; | ||
| 133 | -} | ||
| 134 | - | ||
| 135 | -aclnnStatus aclnnEqualGetWorkspaceSize(const aclTensor *self, const aclTensor *other, aclTensor *out, | ||
| 136 | - uint64_t *workspaceSize, aclOpExecutor **executor) { | ||
| 137 | - OP_CHECK_COMM_INPUT(workspaceSize, executor); | ||
| 138 | - | ||
| 139 | - L2_DFX_PHASE_1(aclnnEqual, DFX_IN(self, other), DFX_OUT(out)); | ||
| 140 | - // 固定写法,创建OpExecutor | ||
| 141 | - auto uniqueExecutor = CREATE_EXECUTOR(); | ||
| 142 | - CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); | ||
| 143 | - | ||
| 144 | - // 固定写法,参数检查 | ||
| 145 | - auto ret = CheckParams(self, other, out); | ||
| 146 | - CHECK_RET(ret == ACLNN_SUCCESS, ret); | ||
| 147 | - | ||
| 148 | - if ((self->GetViewShape() != other->GetViewShape()) || (self->IsEmpty() && other->IsEmpty())) { | ||
| 149 | - int64_t dim = 1; | ||
| 150 | - const aclTensor *dims = (uniqueExecutor.get())->ConvertToTensor(&dim, 1, op::DataType::DT_INT64); | ||
| 151 | - aclIntArray *outShape = (uniqueExecutor.get())->AllocIntArray(&dim, 1); | ||
| 152 | - CHECK_RET(dims != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 153 | - CHECK_RET(outShape != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 154 | - // [False] | ||
| 155 | - int64_t val = 0; | ||
| 156 | - // [True] | ||
| 157 | - if ((self->IsEmpty() && other->IsEmpty()) && (self->GetViewShape() == other->GetViewShape())) { | ||
| 158 | - val = 1; | ||
| 159 | } | 92 | } |
| 160 | - const aclTensor *value = (uniqueExecutor.get())->ConvertToTensor(&val, 1, out->GetDataType()); | ||
| 161 | - CHECK_RET(value != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 162 | 93 | ||
| 163 | - // 调用Fill算子kernel,对一维一元张量赋予bool值 | 94 | + // 检查promoteType的数据类型是否在equal算子的支持列表内 |
| 164 | - auto equalOpOut = l0op::Fill(dims, value, outShape, uniqueExecutor.get()); | 95 | + if (!CheckType(promoteType, DTYPE_SUPPORT_LIST)) { |
| 96 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, | ||
| 97 | + "Self dtype %s and other dtype %s get promoteType dtype %s should be in " | ||
| 98 | + "dtype support list [%s].", | ||
| 99 | + op::ToString(self->GetDataType()).GetString(), op::ToString(other->GetDataType()).GetString(), | ||
| 100 | + op::ToString(promoteType).GetString(), op::ToString(DTYPE_SUPPORT_LIST).GetString()); | ||
| 101 | + return false; | ||
| 102 | + } | ||
| 103 | + | ||
| 104 | + // 检查out的数据类型是否是BOOL | ||
| 105 | + OP_CHECK_DTYPE_NOT_SUPPORT(out, OUT_DTYPE_SUPPORT_LIST, return false); | ||
| 106 | + return true; | ||
| 107 | +} | ||
| 108 | + | ||
| 109 | +static bool CheckMaxShape(const aclTensor* self, const aclTensor* other, const aclTensor* out) | ||
| 110 | +{ | ||
| 111 | + OP_CHECK_MAX_DIM(self, DIM_SUPPORT_MAX, return false); | ||
| 112 | + OP_CHECK_MAX_DIM(other, DIM_SUPPORT_MAX, return false); | ||
| 113 | + OP_CHECK_MAX_DIM(out, DIM_SUPPORT_MAX, return false); | ||
| 114 | + return true; | ||
| 115 | +} | ||
| 116 | + | ||
| 117 | +static bool CheckOutShape(const aclTensor* out) | ||
| 118 | +{ | ||
| 119 | + op::Shape outShape; | ||
| 120 | + outShape.SetDimNum(1); | ||
| 121 | + outShape.SetDim(0, 1); | ||
| 122 | + OP_CHECK_SHAPE_NOT_EQUAL_WITH_EXPECTED_SIZE(out, outShape, return false); | ||
| 123 | + return true; | ||
| 124 | +} | ||
| 125 | + | ||
| 126 | +static aclnnStatus CheckParams(const aclTensor* self, const aclTensor* other, const aclTensor* out) | ||
| 127 | +{ | ||
| 128 | + // 1. 检查两个入参参数是否为空指针;out为空指针时不报错,结果输出None(python中的空指针) | ||
| 129 | + CHECK_RET(CheckNotNull(self, other, out), ACLNN_ERR_PARAM_NULLPTR); | ||
| 130 | + | ||
| 131 | + // 2. 检查输入的数据类型是否在API支持的数据类型范围之内,需要根据api定义校验 | ||
| 132 | + CHECK_RET(CheckDtypeValid(self, other, out), ACLNN_ERR_PARAM_INVALID); | ||
| 133 | + | ||
| 134 | + // 3. 入参tensor最大维度检查 | ||
| 135 | + CHECK_RET(CheckMaxShape(self, other, out), ACLNN_ERR_PARAM_INVALID); | ||
| 136 | + | ||
| 137 | + // 4. 输出tensor形状检查 | ||
| 138 | + CHECK_RET(CheckOutShape(out), ACLNN_ERR_PARAM_INVALID); | ||
| 139 | + | ||
| 140 | + return ACLNN_SUCCESS; | ||
| 141 | +} | ||
| 142 | + | ||
| 143 | +aclnnStatus aclnnEqualGetWorkspaceSize(const aclTensor* self, const aclTensor* other, aclTensor* out, | ||
| 144 | + uint64_t* workspaceSize, aclOpExecutor** executor) | ||
| 145 | +{ | ||
| 146 | + OP_CHECK_COMM_INPUT(workspaceSize, executor); | ||
| 147 | + | ||
| 148 | + L2_DFX_PHASE_1(aclnnEqual, DFX_IN(self, other), DFX_OUT(out)); | ||
| 149 | + // 固定写法,创建OpExecutor | ||
| 150 | + auto uniqueExecutor = CREATE_EXECUTOR(); | ||
| 151 | + CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR); | ||
| 152 | + | ||
| 153 | + // 固定写法,参数检查 | ||
| 154 | + auto ret = CheckParams(self, other, out); | ||
| 155 | + CHECK_RET(ret == ACLNN_SUCCESS, ret); | ||
| 156 | + | ||
| 157 | + if ((self->GetViewShape() != other->GetViewShape()) || (self->IsEmpty() && other->IsEmpty())) { | ||
| 158 | + int64_t dim = 1; | ||
| 159 | + const aclTensor* dims = (uniqueExecutor.get())->ConvertToTensor(&dim, 1, op::DataType::DT_INT64); | ||
| 160 | + aclIntArray* outShape = (uniqueExecutor.get())->AllocIntArray(&dim, 1); | ||
| 161 | + CHECK_RET(dims != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 162 | + CHECK_RET(outShape != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 163 | + // [False] | ||
| 164 | + int64_t val = 0; | ||
| 165 | + // [True] | ||
| 166 | + if ((self->IsEmpty() && other->IsEmpty()) && (self->GetViewShape() == other->GetViewShape())) { | ||
| 167 | + val = 1; | ||
| 168 | + } | ||
| 169 | + const aclTensor* value = (uniqueExecutor.get())->ConvertToTensor(&val, 1, out->GetDataType()); | ||
| 170 | + CHECK_RET(value != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 171 | + | ||
| 172 | + // 调用Fill算子kernel,对一维一元张量赋予bool值 | ||
| 173 | + auto equalOpOut = l0op::Fill(dims, value, outShape, uniqueExecutor.get()); | ||
| 174 | + CHECK_RET(equalOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 175 | + | ||
| 176 | + // 固定写法,将计算结果拷贝到输出out上,out可能是非连续的tensor | ||
| 177 | + auto viewCopyResult = l0op::ViewCopy(equalOpOut, out, uniqueExecutor.get()); | ||
| 178 | + CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 179 | + | ||
| 180 | + *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | ||
| 181 | + uniqueExecutor.ReleaseTo(executor); | ||
| 182 | + return ACLNN_SUCCESS; | ||
| 183 | + } | ||
| 184 | + | ||
| 185 | + auto promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); | ||
| 186 | + if (promoteType == op::DataType::DT_BF16) { | ||
| 187 | + promoteType = op::DataType::DT_FLOAT; | ||
| 188 | + } | ||
| 189 | + // 固定写法,将输入self转换成连续的tensor | ||
| 190 | + auto selfContiguous = l0op::Contiguous(self, uniqueExecutor.get()); | ||
| 191 | + CHECK_RET(selfContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 192 | + | ||
| 193 | + // 将输入self的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 194 | + auto selfCasted = l0op::Cast(selfContiguous, promoteType, uniqueExecutor.get()); | ||
| 195 | + CHECK_RET(selfCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 196 | + | ||
| 197 | + // 固定写法,将输入other转换成连续的tensor | ||
| 198 | + auto otherContiguous = l0op::Contiguous(other, uniqueExecutor.get()); | ||
| 199 | + CHECK_RET(otherContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 200 | + | ||
| 201 | + // 将输入other的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 202 | + auto otherCasted = l0op::Cast(otherContiguous, promoteType, uniqueExecutor.get()); | ||
| 203 | + CHECK_RET(otherCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 204 | + | ||
| 205 | + // 调用TensorEqual算子kernel | ||
| 206 | + auto equalOpOut = l0op::TensorEqual(selfCasted, otherCasted, uniqueExecutor.get()); | ||
| 165 | CHECK_RET(equalOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | 207 | CHECK_RET(equalOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 166 | - | 208 | + |
| 167 | - // 固定写法,将计算结果拷贝到输出out上,out可能是非连续的tensor | ||
| 168 | auto viewCopyResult = l0op::ViewCopy(equalOpOut, out, uniqueExecutor.get()); | 209 | auto viewCopyResult = l0op::ViewCopy(equalOpOut, out, uniqueExecutor.get()); |
| 169 | CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | 210 | CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); |
| 170 | 211 | ||
| 212 | + // 固定写法,获取计算过程中需要使用的workspace大小 | ||
| 171 | *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | 213 | *workspaceSize = uniqueExecutor->GetWorkspaceSize(); |
| 172 | uniqueExecutor.ReleaseTo(executor); | 214 | uniqueExecutor.ReleaseTo(executor); |
| 173 | return ACLNN_SUCCESS; | 215 | return ACLNN_SUCCESS; |
| 174 | - } | ||
| 175 | - | ||
| 176 | - auto promoteType = op::PromoteType(self->GetDataType(), other->GetDataType()); | ||
| 177 | - if (promoteType == op::DataType::DT_BF16) { | ||
| 178 | - promoteType = op::DataType::DT_FLOAT; | ||
| 179 | - } | ||
| 180 | - // 固定写法,将输入self转换成连续的tensor | ||
| 181 | - auto selfContiguous = l0op::Contiguous(self, uniqueExecutor.get()); | ||
| 182 | - CHECK_RET(selfContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 183 | - | ||
| 184 | - // 将输入self的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 185 | - auto selfCasted = l0op::Cast(selfContiguous, promoteType, uniqueExecutor.get()); | ||
| 186 | - CHECK_RET(selfCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 187 | - | ||
| 188 | - // 固定写法,将输入other转换成连续的tensor | ||
| 189 | - auto otherContiguous = l0op::Contiguous(other, uniqueExecutor.get()); | ||
| 190 | - CHECK_RET(otherContiguous != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 191 | - | ||
| 192 | - // 将输入other的数据类型转换成隐式数据类型,根据具体算子语义按需调用 | ||
| 193 | - auto otherCasted = l0op::Cast(otherContiguous, promoteType, uniqueExecutor.get()); | ||
| 194 | - CHECK_RET(otherCasted != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 195 | - | ||
| 196 | - // 调用TensorEqual算子kernel | ||
| 197 | - auto equalOpOut = l0op::TensorEqual(selfCasted, otherCasted, uniqueExecutor.get()); | ||
| 198 | - CHECK_RET(equalOpOut != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 199 | - | ||
| 200 | - auto viewCopyResult = l0op::ViewCopy(equalOpOut, out, uniqueExecutor.get()); | ||
| 201 | - CHECK_RET(viewCopyResult != nullptr, ACLNN_ERR_INNER_NULLPTR); | ||
| 202 | - | ||
| 203 | - // 固定写法,获取计算过程中需要使用的workspace大小 | ||
| 204 | - *workspaceSize = uniqueExecutor->GetWorkspaceSize(); | ||
| 205 | - uniqueExecutor.ReleaseTo(executor); | ||
| 206 | - return ACLNN_SUCCESS; | ||
| 207 | } | 216 | } |
| 208 | 217 | ||
| 209 | -aclnnStatus aclnnEqual(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream) { | 218 | +aclnnStatus aclnnEqual(void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream) |
| 210 | - L2_DFX_PHASE_2(aclnnEqual); | 219 | +{ |
| 211 | - // 固定写法,调用框架能力,完成计算 | 220 | + L2_DFX_PHASE_2(aclnnEqual); |
| 212 | - return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | 221 | + // 固定写法,调用框架能力,完成计算 |
| 222 | + return CommonOpExecutorRun(workspace, workspaceSize, executor, stream); | ||
| 213 | } | 223 | } |
| 214 | 224 | ||
| 215 | 225 | ||


本 PR 改了 13 个算子的 dtype 列表选择逻辑,但没有配套 UT。仓里已有 SetPlatformNpuArch 模拟芯片架构的用法(#5519 给 remainder 加的用例就是这种写法),建议至少给列表选择有变化的算子补上 DAV_3510 下的用例,不然后续改动没有回归保护。