已合并
test: 补充ascir_common与ascir_att_impl单元测试(#210) #1985
test: 补充ascir_common与ascir_att_impl单元测试(#210) #1985
已合并
Leechi666创建于 28 天前
共 3 个文件变更+336-0
@@ -37,8 +37,11 @@ add_library(test_ascir_ut OBJECT
37 code_dumper_unittest.cc37 code_dumper_unittest.cc
38 ascir_utils_unittest.cc38 ascir_utils_unittest.cc
39 test_asc_graph_utils.cpp39 test_asc_graph_utils.cpp
40+ test_ascir_att_impl.cpp
41+ test_ascir_common.cpp
40)42)
41target_include_directories(test_ascir_ut PRIVATE43target_include_directories(test_ascir_ut PRIVATE
44+ ${CODE_ROOT_DIR}/ascir/generator
42 ${ASCEND_ROOT}/x86_64-linux/include45 ${ASCEND_ROOT}/x86_64-linux/include
43 ${ASCEND_ROOT}/opp/built-in/op_proto/inc)46 ${ASCEND_ROOT}/opp/built-in/op_proto/inc)
44target_link_libraries(test_ascir_ut47target_link_libraries(test_ascir_ut
@@ -0,0 +1,67 @@
1+/**
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 of
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.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+#include "gtest/gtest.h"
11+ 
12+#include "graph/ascendc_ir/ascir_registry.h"
13+#include "v1_ascir_att_impl.h"
14+ 
15+namespace af {
16+namespace ascir {
17+ 
18+class AscIrAttImplTest : public ::testing::Test {
19+ protected:
20+ void SetUp() override {}
21+ void TearDown() override {}
22+};
23+ 
24+// Each AscIrAtt impl exposes the IR name through both perf-table entry points.
25+#define EXPECT_ATT_IMPL_NAMED(ir_name) \
26+ { \
27+ JOIN(ir_name, AscIrAttImpl) impl; \
28+ EXPECT_STREQ(static_cast<const char *>(impl.GetApiPerf()), #ir_name); \
29+ EXPECT_STREQ(static_cast<const char *>(impl.GetAscendCApiPerfTable()), #ir_name); \
30+ }
31+ 
32+TEST_F(AscIrAttImplTest, GetApiPerfReturnsIrName_ElementWise) {
33+ EXPECT_ATT_IMPL_NAMED(Add);
34+ EXPECT_ATT_IMPL_NAMED(Gather);
35+ EXPECT_ATT_IMPL_NAMED(Abs);
36+ EXPECT_ATT_IMPL_NAMED(Broadcast);
37+ EXPECT_ATT_IMPL_NAMED(Cast);
38+ EXPECT_ATT_IMPL_NAMED(Div);
39+ EXPECT_ATT_IMPL_NAMED(Erf);
40+ EXPECT_ATT_IMPL_NAMED(Exp);
41+ EXPECT_ATT_IMPL_NAMED(LogicalAnd);
42+ EXPECT_ATT_IMPL_NAMED(LogicalOr);
43+ EXPECT_ATT_IMPL_NAMED(LogicalNot);
44+}
45+ 
46+TEST_F(AscIrAttImplTest, GetApiPerfReturnsIrName_ReduceArgMax) {
47+ EXPECT_ATT_IMPL_NAMED(ReduceArgMax);
48+ EXPECT_ATT_IMPL_NAMED(ReduceArgMaxMultiRPhase1);
49+ EXPECT_ATT_IMPL_NAMED(ReduceArgMaxMultiRPhase2);
50+}
51+ 
52+TEST_F(AscIrAttImplTest, GetApiPerfReturnsIrName_NoModeling) {
53+ EXPECT_ATT_IMPL_NAMED(Data);
54+ EXPECT_ATT_IMPL_NAMED(Scalar);
55+ EXPECT_ATT_IMPL_NAMED(IndexExpr);
56+ EXPECT_ATT_IMPL_NAMED(Output);
57+ EXPECT_ATT_IMPL_NAMED(Workspace);
58+ EXPECT_ATT_IMPL_NAMED(MatMul);
59+ EXPECT_ATT_IMPL_NAMED(Conv2D);
60+ EXPECT_ATT_IMPL_NAMED(Pad);
61+ EXPECT_ATT_IMPL_NAMED(Nop);
62+ EXPECT_ATT_IMPL_NAMED(Ln);
63+ EXPECT_ATT_IMPL_NAMED(Isnan);
64+}
65+ 
66+} // namespace ascir
67+} // namespace af
@@ -0,0 +1,266 @@
1+/**
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 of
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.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+#include "gtest/gtest.h"
11+ 
12+#include "ascir_common.h"
13+#include "ascir.h"
14+#include "ascir_ops.h"
15+ 
16+namespace af {
17+namespace ascir {
18+ 
19+using namespace af::ascir_op;
20+ 
21+class AscirCommonTest : public ::testing::Test {
22+ protected:
23+ void SetUp() override {}
24+ void TearDown() override {}
25+};
26+ 
27+// Test: dtypes present in the conversion map are replaced on both the input and the output side
28+TEST_F(AscirCommonTest, GetConversionFromDtypeMap_ShouldConvertMappedDtypes) {
29+ af::AscGraph graph("test");
30+ auto s0 = graph.CreateSizeVar("s0");
31+ auto s1 = graph.CreateSizeVar("s1");
32+ 
33+ auto z0 = graph.CreateAxis("z0", s0);
34+ auto z1 = graph.CreateAxis("z1", s1);
35+ 
36+ af::ascir_op::Data x1("x1", graph);
37+ af::ascir_op::Load load1("load1");
38+ af::ascir_op::Cast cast("cast");
39+ af::ascir_op::Store store("store");
40+ af::ascir_op::Output y("y");
41+ 
42+ x1.attr.sched.axis = {z0.id, z1.id};
43+ x1.y.dtype = af::DT_FLOAT;
44+ *x1.y.axis = {z0.id, z1.id};
45+ *x1.y.repeats = {s0, s1};
46+ *x1.y.strides = {s1, Symbol(1)};
47+ 
48+ load1.x = x1.y;
49+ load1.attr.sched.axis = {z0.id, z1.id};
50+ load1.y.dtype = af::DT_FLOAT;
51+ *load1.y.axis = {z0.id, z1.id};
52+ *load1.y.repeats = {s0, s1};
53+ *load1.y.strides = {s1, Symbol(1)};
54+ *load1.y.vectorized_axis = {z0.id, z1.id};
55+ 
56+ cast.x = load1.y;
57+ cast.attr.sched.axis = {z0.id, z1.id};
58+ cast.y.dtype = af::DT_INT32;
59+ *cast.y.axis = {z0.id, z1.id};
60+ *cast.y.repeats = {s0, s1};
61+ *cast.y.strides = {s1, Symbol(1)};
62+ *cast.y.vectorized_axis = {z0.id, z1.id};
63+ 
64+ store.x = cast.y;
65+ store.attr.sched.axis = {z0.id, z1.id};
66+ store.y.dtype = af::DT_INT32;
67+ *store.y.axis = {z0.id, z1.id};
68+ *store.y.repeats = {s0, s1};
69+ *store.y.strides = {s1, Symbol(1)};
70+ 
71+ y.x = store.y;
72+ y.attr.sched.axis = {z0.id, z1.id};
73+ y.y.dtype = af::DT_INT32;
74+ *y.y.axis = {z0.id, z1.id};
75+ *y.y.repeats = {s0, s1};
76+ *y.y.strides = {s1, Symbol(1)};
77+ 
78+ std::shared_ptr<af::AscNode> node = graph.FindNode("cast");
79+ node->inputs[0].attr.vectorized_strides = {s1, Symbol(1)};
80+ node->outputs[0].attr.vectorized_strides = {s1, Symbol(1)};
81+ 
82+ std::map<ge::DataType, ge::DataType> dtype_conversion_map = {{af::DT_FLOAT, af::DT_BF16},
83+ {af::DT_INT32, af::DT_FLOAT16}};
84+ auto conversion = GetConversionFromDtypeMap(*node, dtype_conversion_map);
85+ ASSERT_EQ(conversion.first.size(), 1U);
86+ EXPECT_EQ(conversion.first[0], af::DT_BF16);
87+ ASSERT_EQ(conversion.second.size(), 1U);
88+ EXPECT_EQ(conversion.second[0], af::DT_FLOAT16);
89+}
90+ 
91+// Test: both input and output have two vectorized axes whose strides are continuous -> true
92+TEST_F(AscirCommonTest, IsAllVecAxisContinuous_ShouldReturnTrue_WhenStridesContinuous) {
93+ af::AscGraph graph("test");
94+ auto s0 = graph.CreateSizeVar("s0");
95+ auto s1 = graph.CreateSizeVar("s1");
96+ 
97+ auto z0 = graph.CreateAxis("z0", s0);
98+ auto z1 = graph.CreateAxis("z1", s1);
99+ 
100+ af::ascir_op::Data x1("x1", graph);
101+ af::ascir_op::Load load1("load1");
102+ af::ascir_op::Abs abs("abs");
103+ af::ascir_op::Store store("store");
104+ af::ascir_op::Output y("y");
105+ 
106+ x1.attr.sched.axis = {z0.id, z1.id};
107+ x1.y.dtype = af::DT_FLOAT;
108+ *x1.y.axis = {z0.id, z1.id};
109+ *x1.y.repeats = {s0, s1};
110+ *x1.y.strides = {s1, Symbol(1)};
111+ 
112+ load1.x = x1.y;
113+ load1.attr.sched.axis = {z0.id, z1.id};
114+ load1.y.dtype = af::DT_FLOAT;
115+ *load1.y.axis = {z0.id, z1.id};
116+ *load1.y.repeats = {s0, s1};
117+ *load1.y.strides = {s1, Symbol(1)};
118+ *load1.y.vectorized_axis = {z0.id, z1.id};
119+ 
120+ abs.x = load1.y;
121+ abs.attr.sched.axis = {z0.id, z1.id};
122+ abs.y.dtype = af::DT_FLOAT;
123+ *abs.y.axis = {z0.id, z1.id};
124+ *abs.y.repeats = {s0, s1};
125+ *abs.y.strides = {s1, Symbol(1)};
126+ *abs.y.vectorized_axis = {z0.id, z1.id};
127+ 
128+ store.x = abs.y;
129+ store.attr.sched.axis = {z0.id, z1.id};
130+ store.y.dtype = af::DT_FLOAT;
131+ *store.y.axis = {z0.id, z1.id};
132+ *store.y.repeats = {s0, s1};
133+ *store.y.strides = {s1, Symbol(1)};
134+ 
135+ y.x = store.y;
136+ y.attr.sched.axis = {z0.id, z1.id};
137+ y.y.dtype = af::DT_FLOAT;
138+ *y.y.axis = {z0.id, z1.id};
139+ *y.y.repeats = {s0, s1};
140+ *y.y.strides = {s1, Symbol(1)};
141+ 
142+ std::shared_ptr<af::AscNode> node = graph.FindNode("abs");
143+ // repeats[axis_id] * strides[j] == strides[j-1]: s1 * 1 == s1 on both sides
144+ node->inputs[0].attr.vectorized_strides = {s1, Symbol(1)};
145+ node->outputs[0].attr.vectorized_strides = {s1, Symbol(1)};
146+ EXPECT_TRUE(IsAllVecAxisContinuous(*node));
147+}
148+ 
149+// Test: input strides are continuous but output strides are not -> false (output loop returns early)
150+TEST_F(AscirCommonTest, IsAllVecAxisContinuous_ShouldReturnFalse_WhenOutputStridesNotContinuous) {
151+ af::AscGraph graph("test");
152+ auto s0 = graph.CreateSizeVar("s0");
153+ auto s1 = graph.CreateSizeVar("s1");
154+ 
155+ auto z0 = graph.CreateAxis("z0", s0);
156+ auto z1 = graph.CreateAxis("z1", s1);
157+ 
158+ af::ascir_op::Data x1("x1", graph);
159+ af::ascir_op::Load load1("load1");
160+ af::ascir_op::Abs abs("abs");
161+ af::ascir_op::Store store("store");
162+ af::ascir_op::Output y("y");
163+ 
164+ x1.attr.sched.axis = {z0.id, z1.id};
165+ x1.y.dtype = af::DT_FLOAT;
166+ *x1.y.axis = {z0.id, z1.id};
167+ *x1.y.repeats = {s0, s1};
168+ *x1.y.strides = {s1, Symbol(1)};
169+ 
170+ load1.x = x1.y;
171+ load1.attr.sched.axis = {z0.id, z1.id};
172+ load1.y.dtype = af::DT_FLOAT;
173+ *load1.y.axis = {z0.id, z1.id};
174+ *load1.y.repeats = {s0, s1};
175+ *load1.y.strides = {s1, Symbol(1)};
176+ *load1.y.vectorized_axis = {z0.id, z1.id};
177+ 
178+ abs.x = load1.y;
179+ abs.attr.sched.axis = {z0.id, z1.id};
180+ abs.y.dtype = af::DT_FLOAT;
181+ *abs.y.axis = {z0.id, z1.id};
182+ *abs.y.repeats = {s0, s1};
183+ *abs.y.strides = {s1, Symbol(1)};
184+ *abs.y.vectorized_axis = {z0.id, z1.id};
185+ 
186+ store.x = abs.y;
187+ store.attr.sched.axis = {z0.id, z1.id};
188+ store.y.dtype = af::DT_FLOAT;
189+ *store.y.axis = {z0.id, z1.id};
190+ *store.y.repeats = {s0, s1};
191+ *store.y.strides = {s1, Symbol(1)};
192+ 
193+ y.x = store.y;
194+ y.attr.sched.axis = {z0.id, z1.id};
195+ y.y.dtype = af::DT_FLOAT;
196+ *y.y.axis = {z0.id, z1.id};
197+ *y.y.repeats = {s0, s1};
198+ *y.y.strides = {s1, Symbol(1)};
199+ 
200+ std::shared_ptr<af::AscNode> node = graph.FindNode("abs");
201+ node->inputs[0].attr.vectorized_strides = {s1, Symbol(1)};
202+ // s1 * 1 != 2 * s1 -> the output loop detects discontinuity
203+ node->outputs[0].attr.vectorized_strides = {Symbol(2) * s1, Symbol(1)};
204+ EXPECT_FALSE(IsAllVecAxisContinuous(*node));
205+}
206+ 
207+// Test: input strides themselves are not continuous -> false (input loop returns early)
208+TEST_F(AscirCommonTest, IsAllVecAxisContinuous_ShouldReturnFalse_WhenInputStridesNotContinuous) {
209+ af::AscGraph graph("test");
210+ auto s0 = graph.CreateSizeVar("s0");
211+ auto s1 = graph.CreateSizeVar("s1");
212+ 
213+ auto z0 = graph.CreateAxis("z0", s0);
214+ auto z1 = graph.CreateAxis("z1", s1);
215+ 
216+ af::ascir_op::Data x1("x1", graph);
217+ af::ascir_op::Load load1("load1");
218+ af::ascir_op::Abs abs("abs");
219+ af::ascir_op::Store store("store");
220+ af::ascir_op::Output y("y");
221+ 
222+ x1.attr.sched.axis = {z0.id, z1.id};
223+ x1.y.dtype = af::DT_FLOAT;
224+ *x1.y.axis = {z0.id, z1.id};
225+ *x1.y.repeats = {s0, s1};
226+ *x1.y.strides = {s1, Symbol(1)};
227+ 
228+ load1.x = x1.y;
229+ load1.attr.sched.axis = {z0.id, z1.id};
230+ load1.y.dtype = af::DT_FLOAT;
231+ *load1.y.axis = {z0.id, z1.id};
232+ *load1.y.repeats = {s0, s1};
233+ *load1.y.strides = {s1, Symbol(1)};
234+ *load1.y.vectorized_axis = {z0.id, z1.id};
235+ 
236+ abs.x = load1.y;
237+ abs.attr.sched.axis = {z0.id, z1.id};
238+ abs.y.dtype = af::DT_FLOAT;
239+ *abs.y.axis = {z0.id, z1.id};
240+ *abs.y.repeats = {s0, s1};
241+ *abs.y.strides = {s1, Symbol(1)};
242+ *abs.y.vectorized_axis = {z0.id, z1.id};
243+ 
244+ store.x = abs.y;
245+ store.attr.sched.axis = {z0.id, z1.id};
246+ store.y.dtype = af::DT_FLOAT;
247+ *store.y.axis = {z0.id, z1.id};
248+ *store.y.repeats = {s0, s1};
249+ *store.y.strides = {s1, Symbol(1)};
250+ 
251+ y.x = store.y;
252+ y.attr.sched.axis = {z0.id, z1.id};
253+ y.y.dtype = af::DT_FLOAT;
254+ *y.y.axis = {z0.id, z1.id};
255+ *y.y.repeats = {s0, s1};
256+ *y.y.strides = {s1, Symbol(1)};
257+ 
258+ std::shared_ptr<af::AscNode> node = graph.FindNode("abs");
259+ // s1 * 1 != 2 * s1 -> the input loop detects discontinuity
260+ node->inputs[0].attr.vectorized_strides = {Symbol(2) * s1, Symbol(1)};
261+ node->outputs[0].attr.vectorized_strides = {s1, Symbol(1)};
262+ EXPECT_FALSE(IsAllVecAxisContinuous(*node));
263+}
264+ 
265+} // namespace ascir
266+} // namespace af