已合并
feat: 算子910B~910E区间判断追加IsRegbase()并统一IsArch3510判断以兼容后续Regbase芯片 #10101
hahaha22创建于 21 天前
feat: 算子910B~910E区间判断追加IsRegbase()并统一IsArch3510判断以兼容后续Regbase芯片 #10101
已合并
共 5 个文件变更+13-8
| @@ -45,8 +45,9 @@ static const size_t AXIS_LIMIT = 8; // 底层算子不支持超过8维 | |||
| 45 | 45 | ||
| 46 | static inline bool CheckSocVersionIsSupportBf16(void) | 46 | static 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 | ||
| 52 | static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output) | 53 | static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output) |
| @@ -22,6 +22,7 @@ | |||
| 22 | 22 | ||
| 23 | 23 | ||
| 24 | 24 | ||
| 25 | + | ||
| 25 | 26 | ||
| 26 | using namespace op; | 27 | using namespace op; |
| 27 | 28 | ||
| @@ -42,8 +43,9 @@ static bool CheckNotNull(const aclTensor* self, const aclTensor* out) | |||
| 42 | 43 | ||
| 43 | static inline bool CheckSocVersionIsSupportBf16(void) | 44 | static 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 | ||
| 49 | static bool CheckDtypeValid(const aclTensor* self, const aclTensor* out) | 51 | static bool CheckDtypeValid(const aclTensor* self, const aclTensor* out) |
| @@ -45,8 +45,9 @@ static const int64_t AXIS_LIMIT = 8; // 底层算子不支持超过8维 | |||
| 45 | 45 | ||
| 46 | static inline bool CheckSocVersionIsSupportBf16(void) | 46 | static 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 | ||
| 52 | static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output) | 53 | static bool CheckDtypeValid(const aclTensor* gradOutput, const aclTensor* output) |
| @@ -26,6 +26,7 @@ | |||
| 26 | 26 | ||
| 27 | 27 | ||
| 28 | 28 | ||
| 29 | + | ||
| 29 | 30 | ||
| 30 | using namespace op; | 31 | using 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 | |||
| 124 | 124 | ||
| 125 | static bool CheckInputDtype(const aclTensor* input, const aclTensor* weightOptional, const aclTensor* biasOptional) | 125 | static bool CheckInputDtype(const aclTensor* input, const aclTensor* weightOptional, const aclTensor* biasOptional) |
| 126 | { | 126 | { |
IsArch3510 和 add_rms_norm 里的 IsSocVersion950 是同一个问题,名字说的是 DAV_3510 硬编码判断,实际已经是 IsRegbase,后续 Regbase 芯片都会命中,建议趁这次统一改名,不然"统一判断"只统一了实现没统一名字。