已合并
fix: validate RemainderScalarTensor compute dtype on regbase #5519
fix: validate RemainderScalarTensor compute dtype on regbase #5519
已合并
hahaha22创建于 18 天前
共 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为算子支持的数据类型
235static bool CheckPromoteTypeScalarTensor(const op::DataType selfDtype, const op::DataType otherDtype,236static 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+}