已合并
fix: validate RemainderScalarTensor compute dtype on regbase #5519
hahaha22创建于 18 天前
fix: validate RemainderScalarTensor compute dtype on regbase #5519
已合并
共 2 个文件变更+99-0
| @@ -232,6 +232,7 @@ static bool CheckPromoteTypeTensorScalar(const op::DataType selfDtype, const op: | |||
| 232 | } | 232 | } |
| 233 | 233 | ||
| 234 | // 1. self和other没有complex 2. other能cast成outDtype 3. outDtype为算子支持的数据类型 | 234 | // 1. self和other没有complex 2. other能cast成outDtype 3. outDtype为算子支持的数据类型 |
| 235 | +// 4. RegBase场景计算类型为castDtype,castDtype为算子支持的数据类型 | ||
| 235 | static bool CheckPromoteTypeScalarTensor(const op::DataType selfDtype, const op::DataType otherDtype, | 236 | static bool CheckPromoteTypeScalarTensor(const op::DataType selfDtype, const op::DataType otherDtype, |
| 236 | const op::DataType outDtype) | 237 | const op::DataType outDtype) |
| 237 | { | 238 | { |
| @@ -257,6 +258,17 @@ static bool CheckPromoteTypeScalarTensor(const op::DataType selfDtype, const op: | |||
| 257 | return false; | 258 | return false; |
| 258 | } | 259 | } |
| 259 | 260 | ||
| 261 | + // RegBase场景以castDtype作为计算类型,计算后再转换为outDtype | ||
| 262 | + // 其他平台先将两个输入都转换为outDtype再计算 | ||
| 263 | + if (IsRegBase(npuArch) && !CheckType(castDtype, DTYPE_SUPPORT_LIST)) { | ||
| 264 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, | ||
| 265 | + "self dtype %s and other dtype %s promote to unsupported computation dtype %s, " | ||
| 266 | + "should be in dtype support list %s.", | ||
| 267 | + op::ToString(selfDtype).GetString(), op::ToString(otherDtype).GetString(), | ||
| 268 | + op::ToString(castDtype).GetString(), op::ToString(DTYPE_SUPPORT_LIST).GetString()); | ||
| 269 | + return false; | ||
| 270 | + } | ||
| 271 | + | ||
| 260 | return true; | 272 | return true; |
| 261 | } | 273 | } |
| 262 | 274 | ||
| @@ -37,3 +37,90 @@ TEST_F(l2_remainder_scalar_tensor_ascend950_test, double_scalar_int32_tensor_to_ | |||
| 37 | uint64_t workspaceSize = 0; | 37 | uint64_t workspaceSize = 0; |
| 38 | EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACL_SUCCESS); | 38 | EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACL_SUCCESS); |
| 39 | } | 39 | } |
| 40 | + | ||
| 41 | +TEST_F(l2_remainder_scalar_tensor_ascend950_test, integer_scalar_unsupported_compute_dtype) | ||
| 42 | +{ | ||
| 43 | + SetPlatformNpuArch(NpuArch::DAV_3510); | ||
| 44 | + for (auto otherDtype : {ACL_INT8, ACL_UINT8, ACL_INT16}) { | ||
| 45 | + for (auto outDtype : {ACL_INT32, ACL_FLOAT}) { | ||
| 46 | + SCOPED_TRACE(static_cast<int>(otherDtype)); | ||
| 47 | + SCOPED_TRACE(static_cast<int>(outDtype)); | ||
| 48 | + auto self = ScalarDesc(int64_t{7}); | ||
| 49 | + auto other = TensorDesc({4}, otherDtype, ACL_FORMAT_ND).ValueRange(2, 5); | ||
| 50 | + auto out = TensorDesc({4}, outDtype, ACL_FORMAT_ND); | ||
| 51 | + auto ut = OP_API_UT(aclnnRemainderScalarTensor, INPUT(self, other), OUTPUT(out)); | ||
| 52 | + uint64_t workspaceSize = 0; | ||
| 53 | + EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACLNN_ERR_PARAM_INVALID); | ||
| 54 | + } | ||
| 55 | + } | ||
| 56 | +} | ||
| 57 | + | ||
| 58 | +TEST_F(l2_remainder_scalar_tensor_ascend950_test, float_scalar_promotes_int8_tensor) | ||
| 59 | +{ | ||
| 60 | + SetPlatformNpuArch(NpuArch::DAV_3510); | ||
| 61 | + auto self = ScalarDesc(7.0f); | ||
| 62 | + auto other = TensorDesc({4}, ACL_INT8, ACL_FORMAT_ND).ValueRange(2, 5); | ||
| 63 | + auto out = TensorDesc({4}, ACL_FLOAT, ACL_FORMAT_ND); | ||
| 64 | + auto ut = OP_API_UT(aclnnRemainderScalarTensor, INPUT(self, other), OUTPUT(out)); | ||
| 65 | + uint64_t workspaceSize = 0; | ||
| 66 | + EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACL_SUCCESS); | ||
| 67 | +} | ||
| 68 | + | ||
| 69 | +TEST_F(l2_remainder_scalar_tensor_ascend950_test, non_regbase_widens_int8_tensor_to_output_dtype) | ||
| 70 | +{ | ||
| 71 | + SetPlatformNpuArch(NpuArch::DAV_2201); | ||
| 72 | + auto self = ScalarDesc(int64_t{7}); | ||
| 73 | + auto other = TensorDesc({4}, ACL_INT8, ACL_FORMAT_ND).ValueRange(2, 5); | ||
| 74 | + auto out = TensorDesc({4}, ACL_INT64, ACL_FORMAT_ND); | ||
| 75 | + auto ut = OP_API_UT(aclnnRemainderScalarTensor, INPUT(self, other), OUTPUT(out)); | ||
| 76 | + uint64_t workspaceSize = 0; | ||
| 77 | + EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACL_SUCCESS); | ||
| 78 | +} | ||
| 79 | + | ||
| 80 | +TEST_F(l2_remainder_scalar_tensor_ascend950_test, empty_tensor_rejects_unsupported_compute_dtype) | ||
| 81 | +{ | ||
| 82 | + SetPlatformNpuArch(NpuArch::DAV_3510); | ||
| 83 | + auto self = ScalarDesc(int64_t{7}); | ||
| 84 | + auto other = TensorDesc({0}, ACL_INT8, ACL_FORMAT_ND); | ||
| 85 | + auto out = TensorDesc({0}, ACL_INT32, ACL_FORMAT_ND); | ||
| 86 | + auto ut = OP_API_UT(aclnnRemainderScalarTensor, INPUT(self, other), OUTPUT(out)); | ||
| 87 | + uint64_t workspaceSize = 0; | ||
| 88 | + EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACLNN_ERR_PARAM_INVALID); | ||
| 89 | +} | ||
| 90 | + | ||
| 91 | +TEST_F(l2_remainder_scalar_tensor_ascend950_test, supported_compute_dtypes) | ||
| 92 | +{ | ||
| 93 | + SetPlatformNpuArch(NpuArch::DAV_3510); | ||
| 94 | + for (auto dtype : {ACL_INT32, ACL_INT64, ACL_FLOAT16, ACL_FLOAT, ACL_DOUBLE, ACL_BF16}) { | ||
| 95 | + SCOPED_TRACE(static_cast<int>(dtype)); | ||
| 96 | + auto self = ScalarDesc(int64_t{7}); | ||
| 97 | + auto other = TensorDesc({4}, dtype, ACL_FORMAT_ND).ValueRange(2, 5); | ||
| 98 | + auto out = TensorDesc({4}, dtype, ACL_FORMAT_ND); | ||
| 99 | + auto ut = OP_API_UT(aclnnRemainderScalarTensor, INPUT(self, other), OUTPUT(out)); | ||
| 100 | + uint64_t workspaceSize = 0; | ||
| 101 | + EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACL_SUCCESS); | ||
| 102 | + } | ||
| 103 | +} | ||
| 104 | + | ||
| 105 | +TEST_F(l2_remainder_scalar_tensor_ascend950_test, empty_tensor_supported_compute_dtype) | ||
| 106 | +{ | ||
| 107 | + SetPlatformNpuArch(NpuArch::DAV_3510); | ||
| 108 | + auto self = ScalarDesc(int64_t{7}); | ||
| 109 | + auto other = TensorDesc({0}, ACL_INT32, ACL_FORMAT_ND); | ||
| 110 | + auto out = TensorDesc({0}, ACL_INT32, ACL_FORMAT_ND); | ||
| 111 | + auto ut = OP_API_UT(aclnnRemainderScalarTensor, INPUT(self, other), OUTPUT(out)); | ||
| 112 | + uint64_t workspaceSize = 1; | ||
| 113 | + EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACL_SUCCESS); | ||
| 114 | + EXPECT_EQ(workspaceSize, 0U); | ||
| 115 | +} | ||
| 116 | + | ||
| 117 | +TEST_F(l2_remainder_scalar_tensor_ascend950_test, unsupported_output_dtype) | ||
| 118 | +{ | ||
| 119 | + SetPlatformNpuArch(NpuArch::DAV_3510); | ||
| 120 | + auto self = ScalarDesc(int64_t{7}); | ||
| 121 | + auto other = TensorDesc({4}, ACL_INT32, ACL_FORMAT_ND).ValueRange(2, 5); | ||
| 122 | + auto out = TensorDesc({4}, ACL_INT8, ACL_FORMAT_ND); | ||
| 123 | + auto ut = OP_API_UT(aclnnRemainderScalarTensor, INPUT(self, other), OUTPUT(out)); | ||
| 124 | + uint64_t workspaceSize = 0; | ||
| 125 | + EXPECT_EQ(ut.TestGetWorkspaceSize(&workspaceSize), ACLNN_ERR_PARAM_INVALID); | ||
| 126 | +} | ||