已合并
feat: 算子910B~910E区间判断追加IsRegBase()并统一DAV_3510硬编码判断以兼容后续Regbase芯片 #5467
hahaha22创建于 28 天前
feat: 算子910B~910E区间判断追加IsRegBase()并统一DAV_3510硬编码判断以兼容后续Regbase芯片 #5467
已合并
hahaha22创建于 28 天前
共 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
Hhahaha2227 天前

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

likedislike
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#include "opdev/op_log.h"17#include "opdev/op_log.h"
18#include "opdev/platform.h"18#include "opdev/platform.h"
19#include "opdev/shape_utils.h"19#include "opdev/shape_utils.h"
20+#include "op_api/aclnn_check.h"
20 21 
21using namespace op;22using namespace op;
22 23 
@@ -35,9 +36,11 @@ static const std::initializer_list<op::DataType> ASCEND910B_AICORE_DTYPE_SUPPORT
35static inline const std::initializer_list<op::DataType>& GetAiCoreDtypeSupportListBySocVersion()36static 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算子kernel60// 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算子kernel73// 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)
198static const std::initializer_list<op::DataType> GetInputDtypeSupportList()198static 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;
102static const std::initializer_list<op::DataType> GetInputDtypeSupportList()102static 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#include "opdev/op_executor.h"16#include "opdev/op_executor.h"
17#include "opdev/op_log.h"17#include "opdev/op_log.h"
18#include "opdev/shape_utils.h"18#include "opdev/shape_utils.h"
19+#include "op_api/aclnn_check.h"
19 20 
20using namespace op;21using namespace op;
21 22 
@@ -34,8 +35,10 @@ static const std::initializer_list<op::DataType> REGBASE_DTYPE_SUPPORT_LIST = {
34static inline const std::initializer_list<op::DataType>& GetAiCoreDtypeSupportListBySocVersion()35static 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算子kernel57// 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算子kernel68// 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 
46static const std::initializer_list<DataType>& GetDtypeSupportList(NpuArch npuArch, SocVersion socVersion)46static 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
143static inline const std::initializer_list<op::DataType>& GetDtypeSupportList()143static 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()
153static inline const std::initializer_list<op::DataType>& GetOutDtypeSupportList()153static 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 // 固定写法,创建OpExecutor304 // 固定写法,创建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 scalar444// 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
85static inline const std::initializer_list<op::DataType>& GetDtypeSupportList()85static 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()
95static inline const std::initializer_list<op::DataType>& GetOutDtypeSupportList()95static 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// InplaceLt251// 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 of3 * 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()
147static inline const std::initializer_list<op::DataType>& GetOutputDtypeSupportList()147static 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
Hhahaha2228 天前

这里加的 || IsRegBase() 对当前唯一的调用点来说是恒命中的:第163行调用处写的是 IsRegBase(npuArch) ? GetOutputDtypeSupportList() : supportList,只有 RegBase 平台才会进这个函数,所以函数里 910B~910E 的区间判断和最后 return OUT_DTYPE_SUPPORT_910_LIST; 的兜底分支对现有调用链都是走不到的死代码。既然语义已经变成 RegBase 专用,建议把函数体直接简化成 return REGBASE_OUT_DTYPE_SUPPORT_LIST;,或者改个能体现用途的名字,不然读代码的人会以为 910B 平台也会走到这个分支。

likedislike
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 
72inline static bool CheckSocVersionIsSupportBf16(void)72inline 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 
78inline static bool CheckDtypeValid(const aclTensor* self, const aclTensor* other, const aclTensor* out)79inline 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 
59static inline const std::initializer_list<op::DataType>& GetDtypeSupportList()59static 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 
55static inline const std::initializer_list<op::DataType>& GetDtypeSupportList()55static 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 
122static const std::initializer_list<DataType>& GetDtypeSupportList()122static 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;
47static inline const std::initializer_list<DataType>& GetAiCoreDtypeSupportListBySocVersion()47static 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#include "opdev/tensor_view_utils.h"24#include "opdev/tensor_view_utils.h"
25#include "opdev/op_dfx.h"25#include "opdev/op_dfx.h"
26#include "opdev/platform.h"26#include "opdev/platform.h"
27+#include "op_api/aclnn_check.h"
27 28 
28using namespace op;29using namespace op;
29#ifdef __cplusplus30#ifdef __cplusplus
@@ -66,8 +67,9 @@ inline static bool CheckNotNull(const aclTensor* self, const aclTensor* other, c
66 67 
67inline static bool CheckSocVersionIsSupportBf16(void)68inline 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 
133static inline bool CheckSocVersionIsSupportBf16(void)133static 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#include "opdev/shape_utils.h"22#include "opdev/shape_utils.h"
23#include "opdev/tensor_view_utils.h"23#include "opdev/tensor_view_utils.h"
24#include "opdev/platform.h"24#include "opdev/platform.h"
25+#include "op_api/aclnn_check.h"
25 26 
26using namespace op;27using namespace op;
27#ifdef __cplusplus28#ifdef __cplusplus
@@ -49,174 +50,181 @@ constexpr size_t MAX_DIM_LEN = 8;
49 50 
50// 根据API定义,需要列出所能支持的所有dtype51// 根据API定义,需要列出所能支持的所有dtype
51static const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST = {52static 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();
H
Hhahaha2227 天前

这个文件除了追加 IsRegBase 判断,还顺带做了全文件重排(2 空格改 4 空格、大括号换行风格),±130 行里绝大部分是格式噪音,真正的功能改动只有两处,review 和后续 blame 都不好定位。建议把纯格式化单独拆一个 commit 或 PR。

likedislike
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. 检查双输入是否能broadcast139+ // 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#ifdef __cplusplus228#ifdef __cplusplus
221}229}
222-#endif230+#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 
49static const std::initializer_list<op::DataType> SIGNBIT_DTYPE_SUPPORT_LIST = {49static 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 
54static bool CanUseSignbit(const aclTensor* self)54static bool CanUseSignbit(const aclTensor* self)
55{55{
@@ -65,8 +65,9 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* out)
65 65 
66static inline const std::initializer_list<op::DataType>& GetDtypeSupportList()66static 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的tensor176 // 创建数据为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
114static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion()114static 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
140static inline const std::initializer_list<op::DataType>& GetDtypeSupportListBySocVersion()140static 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#include "opdev/op_log.h"19#include "opdev/op_log.h"
20#include "opdev/shape_utils.h"20#include "opdev/shape_utils.h"
21#include "opdev/platform.h"21#include "opdev/platform.h"
22+#include "op_api/aclnn_check.h"
22 23 
23using namespace op;24using namespace op;
24 25 
@@ -42,8 +43,9 @@ static const std::initializer_list<op::DataType> ASCEND610LITE_DTYPE_SUPPORT_LIS
42static bool IsAiCoreSupport(const aclTensor* self)43static bool IsAiCoreSupport(const aclTensor* self)
43{44{
44 // 获取芯片类型,判断是1971还是198045 // 获取芯片类型,判断是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算子kernel59// 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算子kernel70// 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推导算子输出shape86 // 通过输入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 l0op103+} // namespace l0op
@@ -40,131 +40,139 @@ static const std::initializer_list<op::DataType> ASCEND910_DTYPE_SUPPORT_LIST =
40static const std::initializer_list<op::DataType> ASCEND910B_DTYPE_SUPPORT_LIST = {40static 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#ifdef __cplusplus178#ifdef __cplusplus
@@ -24,6 +24,7 @@
24#include "opdev/tensor_view_utils.h"24#include "opdev/tensor_view_utils.h"
25#include "opdev/platform.h"25#include "opdev/platform.h"
26#include "aclnn_kernels/common/op_error_check.h"26#include "aclnn_kernels/common/op_error_check.h"
27+#include "op_api/aclnn_check.h"
27 28 
28using namespace op;29using namespace op;
29#ifdef __cplusplus30#ifdef __cplusplus
@@ -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所能支持的所有dtype58// 列出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// 算子支持的最大维度
62static const size_t DIM_SUPPORT_MAX = 8;62static 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].",
H
Hhahaha2227 天前

这一行本次重排过,但格式串还是 [%s],op::ToString 对列表的输出本身带中括号,会打出双层括号,和 #5512 修的是同一类问题。既然这行已经动了,建议直接把 [%s] 改成 %s,免得再等一轮。

likedislike
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#ifdef __cplusplus225#ifdef __cplusplus