已合并
【feat】: 补充split_pass UT #4399
WangYanMale创建于 11 天前
【feat】: 补充split_pass UT #4399
已合并
WangYanMale创建于 11 天前
2 个文件变更+264-0
Atests/autofuse/ut/autofuse/flatten_split_pass_unittest.cpp+263-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+#include "graph/ascendc_ir/ascendc_ir_core/ascendc_ir.h"
12+#include "graph/attribute_group/attr_group_symbolic_desc.h"
13+#include "graph/debug/ge_op_types.h"
14+#include "graph/utils/graph_utils_ex.h"
15+#include "graph/debug/ge_attr_define.h"
16+#include "../../eager_style_graph_builder/all_ops_cpp.h"
17+#include "../../eager_style_graph_builder/esb_graph.h"
18+#include "lowering/asc_lowerer/loop_api.h"
19+#include "lowering/asc_lowerer/asc_overrides.h"
20+#include "lowering/lowerings.h"
21+#include "fusion/autofuse_attrs.h"
22+#include "utils/auto_fuse_config.h"
23+#include "../../eager_style_graph_builder/compliant_op_desc_builder.h"
24+#include "pattern_fusion/flatten_split_pass.h"
25+#include "pattern_fusion/pattern_fusion.h"
26+#include "graph_metadef/graph/debug/ge_util.h"
27+#include <gtest/gtest.h>
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
Mtests/depends/aihacb_autofusion/src/autofuse_new_stub.cc+1-0
@@ -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 ge25} // namespace ge