已合并
test: 补充ascir_common与ascir_att_impl单元测试(#210) #1985
Leechi666创建于 28 天前
test: 补充ascir_common与ascir_att_impl单元测试(#210) #1985
已合并
共 3 个文件变更+336-0
| @@ -37,8 +37,11 @@ add_library(test_ascir_ut OBJECT | |||
| 37 | code_dumper_unittest.cc | 37 | code_dumper_unittest.cc |
| 38 | ascir_utils_unittest.cc | 38 | ascir_utils_unittest.cc |
| 39 | test_asc_graph_utils.cpp | 39 | test_asc_graph_utils.cpp |
| 40 | + test_ascir_att_impl.cpp | ||
| 41 | + test_ascir_common.cpp | ||
| 40 | ) | 42 | ) |
| 41 | target_include_directories(test_ascir_ut PRIVATE | 43 | target_include_directories(test_ascir_ut PRIVATE |
| 44 | + ${CODE_ROOT_DIR}/ascir/generator | ||
| 42 | ${ASCEND_ROOT}/x86_64-linux/include | 45 | ${ASCEND_ROOT}/x86_64-linux/include |
| 43 | ${ASCEND_ROOT}/opp/built-in/op_proto/inc) | 46 | ${ASCEND_ROOT}/opp/built-in/op_proto/inc) |
| 44 | target_link_libraries(test_ascir_ut | 47 | target_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 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 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 | + | ||
| 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 | + | ||
| 11 | + | ||
| 12 | + | ||
| 13 | + | ||
| 14 | + | ||
| 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 | ||