已合并
【feat】: 补充split_pass UT #4399
WangYanMale创建于 11 天前
【feat】: 补充split_pass UT #4399
已合并
共 2 个文件变更+264-0
| @@ -0,0 +1,263 @@ | |||
| 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 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | +using namespace std; | ||
| 30 | +using namespace testing; | ||
| 31 | +namespace ge { | ||
| 32 | + | ||
| 33 | +REG_OP(SplitD) | ||
| 34 | + .INPUT(x, TensorType::BasicType()) | ||
| 35 | + .DYNAMIC_OUTPUT(y, TensorType::BasicType()) | ||
| 36 | + .REQUIRED_ATTR(split_dim, Int) | ||
| 37 | + .ATTR(num_split, Int, 1) | ||
| 38 | + .OP_END_FACTORY_REG(SplitD) | ||
| 39 | + | ||
| 40 | + REG_OP(Const) | ||
| 41 | + .OUTPUT(y, TensorType({DT_FLOAT, DT_FLOAT16, DT_INT8, DT_INT16, DT_UINT16, DT_UINT8, DT_INT32, DT_INT64, DT_UINT32, | ||
| 42 | + DT_UINT64, DT_BOOL, DT_DOUBLE})) | ||
| 43 | + .ATTR(value, Tensor, Tensor()) | ||
| 44 | + .OP_END_FACTORY_REG(Const); | ||
| 45 | + | ||
| 46 | +namespace { | ||
| 47 | +GeTensorDesc MakeFp16Desc(const vector<int64_t> &dims) { | ||
| 48 | + GeTensorDesc desc{GeShape(dims)}; | ||
| 49 | + desc.SetFormat(FORMAT_ND); | ||
| 50 | + desc.SetOriginFormat(FORMAT_ND); | ||
| 51 | + desc.SetDataType(DT_FLOAT16); | ||
| 52 | + desc.SetOriginDataType(DT_FLOAT16); | ||
| 53 | + return desc; | ||
| 54 | +} | ||
| 55 | + | ||
| 56 | +NodePtr MakeReluNode(const ComputeGraphPtr &graph, const string &name, const GeTensorDesc &desc) { | ||
| 57 | + auto op_desc = std::make_shared<OpDesc>(name, "Relu"); | ||
| 58 | + op_desc->AddInputDesc(desc); | ||
| 59 | + op_desc->AddOutputDesc(desc); | ||
| 60 | + return graph->AddNode(op_desc); | ||
| 61 | +} | ||
| 62 | + | ||
| 63 | +NodePtr MakeSwitchReluChain(const ComputeGraphPtr &graph, const string &sw_name, const string &relu_name, | ||
| 64 | + const GeTensorDesc &desc) { | ||
| 65 | + auto sw_op_desc = std::make_shared<OpDesc>(sw_name, "switch"); | ||
| 66 | + sw_op_desc->AddOutputDesc(desc); | ||
| 67 | + auto sw = graph->AddNode(sw_op_desc); | ||
| 68 | + auto relu = MakeReluNode(graph, relu_name, desc); | ||
| 69 | + GraphUtils::AddEdge(sw->GetOutDataAnchor(0), relu->GetInDataAnchor(0)); | ||
| 70 | + return relu; | ||
| 71 | +} | ||
| 72 | + | ||
| 73 | +NodePtr MakeSplitDNode(const ComputeGraphPtr &graph, const string &name, uint32_t num_split, int64_t split_dim) { | ||
| 74 | + op::SplitD op(name.c_str()); | ||
| 75 | + op.BreakConnect(); | ||
| 76 | + op.create_dynamic_output_y(num_split); | ||
| 77 | + op.set_attr_split_dim(split_dim); | ||
| 78 | + op.set_attr_num_split(static_cast<int>(num_split)); | ||
| 79 | + return graph->AddNode(ge::OpDescUtils::GetOpDescFromOperator(op)); | ||
| 80 | +} | ||
| 81 | + | ||
| 82 | +void SetupSplitOutputs(const NodePtr &split_node, const GeTensorDesc &out_desc) { | ||
| 83 | + for (uint32_t i = 0U; i < split_node->GetAllOutDataAnchorsSize(); i++) { | ||
| 84 | + split_node->GetOpDesc()->UpdateOutputDesc(i, out_desc); | ||
| 85 | + split_node->GetOpDesc() | ||
| 86 | + ->MutableOutputDesc(i) | ||
| 87 | + ->GetOrCreateAttrsGroup<ge::SymbolicDescAttr>() | ||
| 88 | + ->symbolic_tensor.MutableOriginSymbolShape() = {Symbol("s0"), Symbol("s1")}; | ||
| 89 | + } | ||
| 90 | +} | ||
| 91 | + | ||
| 92 | +NodePtr BuildSplitCascade(const ComputeGraphPtr &graph, const GeTensorDesc &desc_a, const GeTensorDesc &desc_b, | ||
| 93 | + const GeTensorDesc &desc_c, int64_t split_dim) { | ||
| 94 | + auto relu = MakeSwitchReluChain(graph, "sw1", "relu1", desc_a); | ||
| 95 | + auto sp1 = MakeSplitDNode(graph, "S1", 2U, split_dim); | ||
| 96 | + auto sp2 = MakeSplitDNode(graph, "S2", 4U, split_dim); | ||
| 97 | + auto sp3 = MakeSplitDNode(graph, "S3", 4U, split_dim); | ||
| 98 | + GraphUtils::AddEdge(relu->GetOutDataAnchor(0), sp1->GetInDataAnchor(0)); | ||
| 99 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(0), sp2->GetInDataAnchor(0)); | ||
| 100 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(1), sp3->GetInDataAnchor(0)); | ||
| 101 | + sp1->GetOpDesc()->UpdateInputDesc(0, desc_a); | ||
| 102 | + SetupSplitOutputs(sp1, desc_b); | ||
| 103 | + sp2->GetOpDesc()->UpdateInputDesc(0, desc_b); | ||
| 104 | + SetupSplitOutputs(sp2, desc_c); | ||
| 105 | + sp3->GetOpDesc()->UpdateInputDesc(0, desc_b); | ||
| 106 | + SetupSplitOutputs(sp3, desc_c); | ||
| 107 | + return sp1; | ||
| 108 | +} | ||
| 109 | +} // namespace | ||
| 110 | + | ||
| 111 | +class FlattenSplitPassUT : public testing::Test { | ||
| 112 | + protected: | ||
| 113 | + void SetUp() override { | ||
| 114 | + dlog_setlevel(0, 3, 0); | ||
| 115 | + ge::autofuse::AutoFuseConfig::MutableLoweringConfig().experimental_lowering_split = true; | ||
| 116 | + } | ||
| 117 | + void TearDown() override { | ||
| 118 | + ge::autofuse::AutoFuseConfig::MutableLoweringConfig().experimental_lowering_split = false; | ||
| 119 | + dlog_setlevel(0, 3, 0); | ||
| 120 | + } | ||
| 121 | +}; | ||
| 122 | + | ||
| 123 | +TEST_F(FlattenSplitPassUT, RunSplitDisabled) { | ||
| 124 | + ge::autofuse::AutoFuseConfig::MutableLoweringConfig().experimental_lowering_split = false; | ||
| 125 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_disabled"); | ||
| 126 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 127 | +} | ||
| 128 | + | ||
| 129 | +TEST_F(FlattenSplitPassUT, RunNullGraph) { | ||
| 130 | + ComputeGraphPtr null_graph = nullptr; | ||
| 131 | + EXPECT_EQ(FlattenSplitPass::Run(null_graph), ge::GRAPH_SUCCESS); | ||
| 132 | +} | ||
| 133 | + | ||
| 134 | +TEST_F(FlattenSplitPassUT, RunNoSplitNodes) { | ||
| 135 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_no_split"); | ||
| 136 | + auto desc = MakeFp16Desc({128, 32}); | ||
| 137 | + MakeSwitchReluChain(graph, "sw1", "relu1", desc); | ||
| 138 | + graph->TopologicalSorting(); | ||
| 139 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 140 | +} | ||
| 141 | + | ||
| 142 | +TEST_F(FlattenSplitPassUT, RunSingleSplitDNoFusion) { | ||
| 143 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_single_splitd"); | ||
| 144 | + auto desc_a = MakeFp16Desc({128, 32}); | ||
| 145 | + auto desc_b = MakeFp16Desc({128, 16}); | ||
| 146 | + auto relu = MakeSwitchReluChain(graph, "sw1", "relu1", desc_a); | ||
| 147 | + auto split_node = MakeSplitDNode(graph, "SplitNode1", 2U, 1); | ||
| 148 | + auto c1 = MakeReluNode(graph, "c1", desc_b); | ||
| 149 | + auto c2 = MakeReluNode(graph, "c2", desc_b); | ||
| 150 | + GraphUtils::AddEdge(relu->GetOutDataAnchor(0), split_node->GetInDataAnchor(0)); | ||
| 151 | + GraphUtils::AddEdge(split_node->GetOutDataAnchor(0), c1->GetInDataAnchor(0)); | ||
| 152 | + GraphUtils::AddEdge(split_node->GetOutDataAnchor(1), c2->GetInDataAnchor(0)); | ||
| 153 | + split_node->GetOpDesc()->UpdateInputDesc(0, desc_a); | ||
| 154 | + SetupSplitOutputs(split_node, desc_b); | ||
| 155 | + graph->TopologicalSorting(); | ||
| 156 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 157 | +} | ||
| 158 | + | ||
| 159 | +TEST_F(FlattenSplitPassUT, RunSplitDDifferentDimNoFusion) { | ||
| 160 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_diff_dim"); | ||
| 161 | + BuildSplitCascade(graph, MakeFp16Desc({128, 32}), MakeFp16Desc({128, 16}), MakeFp16Desc({128, 2}), 1); | ||
| 162 | + graph->TopologicalSorting(); | ||
| 163 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 164 | +} | ||
| 165 | + | ||
| 166 | +TEST_F(FlattenSplitPassUT, RunSplitDNonSplitPeerNoFusion) { | ||
| 167 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_non_split_peer"); | ||
| 168 | + auto desc_a = MakeFp16Desc({128, 32}); | ||
| 169 | + auto desc_b = MakeFp16Desc({128, 16}); | ||
| 170 | + auto relu = MakeSwitchReluChain(graph, "sw1", "relu1", desc_a); | ||
| 171 | + auto sp1 = MakeSplitDNode(graph, "S1", 2U, 1); | ||
| 172 | + auto c1 = MakeReluNode(graph, "c1", desc_b); | ||
| 173 | + auto c2 = MakeReluNode(graph, "c2", desc_b); | ||
| 174 | + GraphUtils::AddEdge(relu->GetOutDataAnchor(0), sp1->GetInDataAnchor(0)); | ||
| 175 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(0), c1->GetInDataAnchor(0)); | ||
| 176 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(1), c2->GetInDataAnchor(0)); | ||
| 177 | + sp1->GetOpDesc()->UpdateInputDesc(0, desc_a); | ||
| 178 | + SetupSplitOutputs(sp1, desc_b); | ||
| 179 | + graph->TopologicalSorting(); | ||
| 180 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 181 | +} | ||
| 182 | + | ||
| 183 | +TEST_F(FlattenSplitPassUT, RunSplitNegativeDim) { | ||
| 184 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_neg_dim"); | ||
| 185 | + BuildSplitCascade(graph, MakeFp16Desc({128, 32}), MakeFp16Desc({128, 16}), MakeFp16Desc({128, 2}), -1); | ||
| 186 | + graph->TopologicalSorting(); | ||
| 187 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 188 | +} | ||
| 189 | + | ||
| 190 | +TEST_F(FlattenSplitPassUT, RunSplitDMultiConsumerNoFusion) { | ||
| 191 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_multi_consumer"); | ||
| 192 | + auto desc_a = MakeFp16Desc({128, 32}); | ||
| 193 | + auto desc_b = MakeFp16Desc({128, 16}); | ||
| 194 | + auto relu = MakeSwitchReluChain(graph, "sw1", "relu1", desc_a); | ||
| 195 | + auto sp1 = MakeSplitDNode(graph, "S1", 2U, 1); | ||
| 196 | + auto c1 = MakeReluNode(graph, "c1", desc_b); | ||
| 197 | + auto c2 = MakeReluNode(graph, "c2", desc_b); | ||
| 198 | + GraphUtils::AddEdge(relu->GetOutDataAnchor(0), sp1->GetInDataAnchor(0)); | ||
| 199 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(0), c1->GetInDataAnchor(0)); | ||
| 200 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(0), c2->GetInDataAnchor(0)); | ||
| 201 | + sp1->GetOpDesc()->UpdateInputDesc(0, desc_a); | ||
| 202 | + SetupSplitOutputs(sp1, desc_b); | ||
| 203 | + graph->TopologicalSorting(); | ||
| 204 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 205 | +} | ||
| 206 | + | ||
| 207 | +TEST_F(FlattenSplitPassUT, RunSplitDUnknownDimNum) { | ||
| 208 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_unknown_dim"); | ||
| 209 | + auto desc_a = MakeFp16Desc({128, 32}); | ||
| 210 | + auto desc_b = MakeFp16Desc({128, 16}); | ||
| 211 | + auto desc_c = MakeFp16Desc({128, 2}); | ||
| 212 | + auto sp1 = BuildSplitCascade(graph, desc_a, desc_b, desc_c, 1); | ||
| 213 | + sp1->GetOpDesc()->MutableInputDesc(0)->MutableShape().SetIsUnknownDimNum(); | ||
| 214 | + graph->TopologicalSorting(); | ||
| 215 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 216 | +} | ||
| 217 | + | ||
| 218 | +TEST_F(FlattenSplitPassUT, CanFlattenSingleConsumer) { | ||
| 219 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_can_flatten"); | ||
| 220 | + auto sp1 = BuildSplitCascade(graph, MakeFp16Desc({128, 32}), MakeFp16Desc({128, 16}), MakeFp16Desc({128, 16}), 1); | ||
| 221 | + graph->TopologicalSorting(); | ||
| 222 | + EXPECT_EQ(FlattenSplitPass::CanFlatten(sp1, 1, 2), ge::GRAPH_SUCCESS); | ||
| 223 | +} | ||
| 224 | + | ||
| 225 | +TEST_F(FlattenSplitPassUT, CanFlattenMultiConsumerFail) { | ||
| 226 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_can_flatten_fail"); | ||
| 227 | + auto desc_a = MakeFp16Desc({128, 32}); | ||
| 228 | + auto desc_b = MakeFp16Desc({128, 16}); | ||
| 229 | + auto relu = MakeSwitchReluChain(graph, "sw1", "relu1", desc_a); | ||
| 230 | + auto sp1 = MakeSplitDNode(graph, "S1", 2U, 1); | ||
| 231 | + auto c1 = MakeReluNode(graph, "c1", desc_b); | ||
| 232 | + auto c2 = MakeReluNode(graph, "c2", desc_b); | ||
| 233 | + GraphUtils::AddEdge(relu->GetOutDataAnchor(0), sp1->GetInDataAnchor(0)); | ||
| 234 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(0), c1->GetInDataAnchor(0)); | ||
| 235 | + GraphUtils::AddEdge(sp1->GetOutDataAnchor(0), c2->GetInDataAnchor(0)); | ||
| 236 | + sp1->GetOpDesc()->UpdateInputDesc(0, desc_a); | ||
| 237 | + SetupSplitOutputs(sp1, desc_b); | ||
| 238 | + graph->TopologicalSorting(); | ||
| 239 | + EXPECT_EQ(FlattenSplitPass::CanFlatten(sp1, 1, 2), ge::GRAPH_FAILED); | ||
| 240 | +} | ||
| 241 | + | ||
| 242 | +TEST_F(FlattenSplitPassUT, RunMultipleSplitNodesInGraph) { | ||
| 243 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_multi_split_nodes"); | ||
| 244 | + BuildSplitCascade(graph, MakeFp16Desc({128, 32}), MakeFp16Desc({128, 16}), MakeFp16Desc({128, 2}), 1); | ||
| 245 | + auto desc_a = MakeFp16Desc({128, 32}); | ||
| 246 | + auto desc_b = MakeFp16Desc({128, 16}); | ||
| 247 | + auto relu2 = MakeSwitchReluChain(graph, "sw2", "relu2", desc_a); | ||
| 248 | + auto sp4 = MakeSplitDNode(graph, "S4", 2U, 1); | ||
| 249 | + GraphUtils::AddEdge(relu2->GetOutDataAnchor(0), sp4->GetInDataAnchor(0)); | ||
| 250 | + sp4->GetOpDesc()->UpdateInputDesc(0, desc_a); | ||
| 251 | + SetupSplitOutputs(sp4, desc_b); | ||
| 252 | + graph->TopologicalSorting(); | ||
| 253 | + EXPECT_EQ(FlattenSplitPass::Run(graph), ge::GRAPH_SUCCESS); | ||
| 254 | +} | ||
| 255 | + | ||
| 256 | +TEST_F(FlattenSplitPassUT, RunPatternFusionIntegration) { | ||
| 257 | + ComputeGraphPtr graph = std::make_shared<ComputeGraph>("test_pf_integration"); | ||
| 258 | + BuildSplitCascade(graph, MakeFp16Desc({128, 32}), MakeFp16Desc({128, 16}), MakeFp16Desc({128, 2}), 1); | ||
| 259 | + graph->TopologicalSorting(); | ||
| 260 | + PatternFusion pattern_fusion; | ||
| 261 | + EXPECT_EQ(pattern_fusion.RunAllPatternFusion(graph), ge::GRAPH_SUCCESS); | ||
| 262 | +} | ||
| 263 | +} // namespace ge | ||
| @@ -19,6 +19,7 @@ std::unique_ptr<AutofuseBackendSpec> GetAutofuseBackendSpec() { | |||
| 19 | auto instance = std::make_unique<AutofuseBackendSpec>(); | 19 | auto instance = std::make_unique<AutofuseBackendSpec>(); |
| 20 | instance->concat_max_input_num = 0; | 20 | instance->concat_max_input_num = 0; |
| 21 | instance->concat_alg = 0; | 21 | instance->concat_alg = 0; |
| 22 | + instance->slice_split_spec.enable_split_flatten = true; | ||
| 22 | return instance; | 23 | return instance; |
| 23 | } | 24 | } |
| 24 | } // namespace ge | 25 | } // namespace ge |