已合并
feat: 算子910B~910E区间判断追加IsRegbase()并统一IsArch3510判断以兼容后续Regbase芯片 #10101
hahaha22创建于 21 天前
feat: 算子910B~910E区间判断追加IsRegbase()并统一IsArch3510判断以兼容后续Regbase芯片 #10101
已合并
hahaha22创建于 21 天前
共 5 个文件变更+13-8
@@ -45,8 +45,9 @@ static const size_t AXIS_LIMIT = 8; // 底层算子不支持超过8维
45 45 
46static inline bool CheckSocVersionIsSupportBf16(void)46static inline bool CheckSocVersionIsSupportBf16(void)
47{47{
48- return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&48+ return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&
49- GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E;49+ GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) ||
50+ Ops::NN::AclnnUtil::IsRegbase();
50}51}
51 52 
52static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output)53static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output)
@@ -22,6 +22,7 @@
22#include "opdev/data_type_utils.h"22#include "opdev/data_type_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_util.h"
25 26 
26using namespace op;27using namespace op;
27#ifdef __cplusplus28#ifdef __cplusplus
@@ -42,8 +43,9 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* out)
42 43 
43static inline bool CheckSocVersionIsSupportBf16(void)44static inline bool CheckSocVersionIsSupportBf16(void)
44{45{
45- return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&46+ return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&
46- GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E;47+ GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) ||
48+ Ops::NN::AclnnUtil::IsRegbase();
47}49}
48 50 
49static bool CheckDtypeValid(const aclTensor* self, const aclTensor* out)51static bool CheckDtypeValid(const aclTensor* self, const aclTensor* out)
@@ -45,8 +45,9 @@ static const int64_t AXIS_LIMIT = 8; // 底层算子不支持超过8维
45 45 
46static inline bool CheckSocVersionIsSupportBf16(void)46static inline bool CheckSocVersionIsSupportBf16(void)
47{47{
48- return GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&48+ return (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&
49- GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E;49+ GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) ||
50+ Ops::NN::AclnnUtil::IsRegbase();
50}51}
51 52 
52static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output)53static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output)
@@ -26,6 +26,7 @@
26#include "opdev/op_log.h"26#include "opdev/op_log.h"
27#include "opdev/tensor_view_utils.h"27#include "opdev/tensor_view_utils.h"
28#include "aclnn_batch_norm_elemt.h"28#include "aclnn_batch_norm_elemt.h"
29+#include "op_api/aclnn_util.h"
29 30 
30using namespace op;31using namespace op;
31 32 
@@ -50,7 +51,7 @@ static bool CheckDtypeValid(const T& t, const Ts&... args)
50 if constexpr (std::is_same_v<T, aclTensor*> || std::is_same_v<T, const aclTensor*>) {51 if constexpr (std::is_same_v<T, aclTensor*> || std::is_same_v<T, const aclTensor*>) {
51 bool isBf16SupportedSocVersion = (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&52 bool isBf16SupportedSocVersion = (GetCurrentPlatformInfo().GetSocVersion() >= SocVersion::ASCEND910B &&
52 GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) ||53 GetCurrentPlatformInfo().GetSocVersion() <= SocVersion::ASCEND910E) ||
53- GetCurrentPlatformInfo().GetSocVersion() == SocVersion::ASCEND950;54+ Ops::NN::AclnnUtil::IsRegbase();
54 const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST = isBf16SupportedSocVersion ?55 const std::initializer_list<op::DataType> DTYPE_SUPPORT_LIST = isBf16SupportedSocVersion ?
55 ASCEND910B_DTYPE_SUPPORT_LIST :56 ASCEND910B_DTYPE_SUPPORT_LIST :
56 ASCEND910_DTYPE_SUPPORT_LIST;57 ASCEND910_DTYPE_SUPPORT_LIST;
@@ -120,7 +120,7 @@ static bool CheckMeanRstdOutputShape(const aclTensor* input, const aclIntArray*
120 return true;120 return true;
121}121}
122 122 
123-static bool IsArch3510() { return GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510; }123+static bool IsArch3510() { return Ops::NN::AclnnUtil::IsRegbase(); }
H
Hhahaha2220 天前

IsArch3510 和 add_rms_norm 里的 IsSocVersion950 是同一个问题,名字说的是 DAV_3510 硬编码判断,实际已经是 IsRegbase,后续 Regbase 芯片都会命中,建议趁这次统一改名,不然"统一判断"只统一了实现没统一名字。

likedislike
124 124 
125static bool CheckInputDtype(const aclTensor* input, const aclTensor* weightOptional, const aclTensor* biasOptional)125static bool CheckInputDtype(const aclTensor* input, const aclTensor* weightOptional, const aclTensor* biasOptional)
126{126{