已合并
【PR】: broadcast backward revert #1912
【PR】: broadcast backward revert #1912
已合并
czways创建于 5 天前
16 个文件变更+133-4459
@@ -1,24 +0,0 @@
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 of
4- * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License in the root of the software repository for the full text of the License.
6- * THIS FILE 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 directory of the software repository for the full text of the License.
9- */
10-#ifndef OPTIMIZE_GRAPH_PASS_BROADCAST_BACKWARD_PASS_H
11-#define OPTIMIZE_GRAPH_PASS_BROADCAST_BACKWARD_PASS_H
12- 
13-#include "base_graph_pass.h"
14- 
15-namespace optimize {
16-class BroadcastBackwardPass final : public BaseGraphPass {
17- public:
18- BroadcastBackwardPass() = default;
19- ~BroadcastBackwardPass() override = default;
20- Status RunPass(af::AscGraph &graph) override;
21-};
22-} // namespace optimize
23- 
24-#endif // OPTIMIZE_GRAPH_PASS_BROADCAST_BACKWARD_PASS_H
@@ -1,388 +0,0 @@
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 of
4- * CANN Open Software License Agreement Version 2.0.
5- */
6-#include "broadcast_backward_shared_split.h"
7- 
8-#include <algorithm>
9-#include <string>
10-#include <vector>
11- 
12-#include "ascir_ops.h"
13-#include "ascir_ops_utils.h"
14-#include "graph_utils.h"
15-#include "schedule_utils.h"
16- 
17-using namespace ascir;
18-using namespace af::ascir_op;
19-using namespace af::ops;
20- 
21-namespace optimize {
22-namespace broadcast_backward_shared_split {
23-namespace {
24- 
25-using af::AscGraph;
26-using af::AscNode;
27-using af::AscNodePtr;
28-using af::AscTensorAttr;
29-using af::FAILED;
30-using af::SUCCESS;
31-using NodePtr = AscNodePtr;
32- 
33-constexpr const char *kStoreType = Store::Type;
34-constexpr const char *kScalarType = Scalar::Type;
35-constexpr const char *kBroadcastType = Broadcast::Type;
36-constexpr const char *kCastType = Cast::Type;
37- 
38-NodePtr ToAscNode(const af::NodePtr &node) {
39- return std::dynamic_pointer_cast<AscNode>(node);
40-}
41- 
42-bool IsSingleInAndOutNode(const NodePtr &node) {
43- return node != nullptr && node->GetInDataNodesSize() == 1UL && node->GetOutDataNodesSize() == 1UL;
44-}
45- 
46-Status GetPeerOutNodeSafe(const NodePtr &node, NodePtr &peer, int32_t index) {
47- GE_ASSERT_NOTNULL(node);
48- if (index < 0 || static_cast<size_t>(index) >= node->GetAllInDataAnchorsSize()) {
49- return FAILED;
50- }
51- auto in_anchor = node->GetInDataAnchor(index);
52- GE_ASSERT_NOTNULL(in_anchor);
53- auto out_anchor = in_anchor->GetPeerOutAnchor();
54- GE_ASSERT_NOTNULL(out_anchor);
55- peer = ToAscNode(out_anchor->GetOwnerNode());
56- return SUCCESS;
57-}
58- 
59-Status GetOutputTensorAttr(const NodePtr &node, AscTensorAttr *&attr) {
60- GE_ASSERT_NOTNULL(node);
61- GE_ASSERT_TRUE(node->GetAllOutDataAnchorsSize() > 0U);
62- attr = &node->outputs[0].attr;
63- return SUCCESS;
64-}
65- 
66-bool IsScalarInput(const NodePtr &node) {
67- auto current = node;
68- while (current != nullptr) {
69- if (current->GetType() == kScalarType) {
70- return true;
71- }
72- if (current->GetType() != kBroadcastType || GetPeerOutNodeSafe(current, current, 0) != SUCCESS) {
73- break;
74- }
75- }
76- return false;
77-}
78- 
79-bool IsViewOp(const NodePtr &node) {
80- static const std::vector<std::string> kViewTypes = {
81- Transpose::Type, Broadcast::Type, "Slice", Split::Type, Concat::Type, Gather::Type, "Sum",
82- "Mean", "Max", "Min", "Prod", "Any", "All"};
83- return std::find(kViewTypes.begin(), kViewTypes.end(), node->GetType()) != kViewTypes.end();
84-}
85- 
86-bool IsDtypeNotSupported(const NodePtr &node) {
87- if (node->GetType() != kCastType || node->GetOpDesc() == nullptr) {
88- return false;
89- }
90- const auto output_desc = node->GetOpDesc()->MutableOutputDesc(0);
91- if (output_desc == nullptr) {
92- return true;
93- }
94- const auto dtype = output_desc->GetDataType();
95- const std::vector<af::DataType> input_dtypes = {dtype};
96- std::vector<af::DataType> output_dtypes = {dtype};
97- return ScheduleUtils::CallAscirInferDataType<Broadcast>(input_dtypes, output_dtypes) != SUCCESS;
98-}
99- 
100-bool CanBackward(const NodePtr &node) {
101- if (node == nullptr || node->GetType() == kStoreType || IsViewOp(node) || IsDtypeNotSupported(node)) {
102- return false;
103- }
104- return node->GetAllOutDataAnchorsSize() <= 1U && node->GetInDataNodesSize() == 1UL;
105-}
106- 
107-bool CanSharedConsumer(const NodePtr &node) {
108- return node != nullptr && node->GetType() != kStoreType && !IsViewOp(node) && !IsDtypeNotSupported(node);
109-}
110- 
111-Status TraceMergeNode(const NodePtr &start, NodePtr &merge) {
112- NodePtr current = start;
113- std::vector<NodePtr> visited;
114- while (current != nullptr) {
115- if (std::find(visited.begin(), visited.end(), current) != visited.end()) {
116- return FAILED;
117- }
118- visited.push_back(current);
119- if (current->GetInDataNodesSize() != 1UL || current->GetType() == kStoreType) {
120- merge = current;
121- return SUCCESS;
122- }
123- if (current->GetAllOutDataAnchorsSize() != 1U ||
124- current->GetOutDataAnchor(0)->GetPeerInDataAnchors().size() != 1U) {
125- return FAILED;
126- }
127- current = ToAscNode(current->GetOutDataNodes().at(0));
128- }
129- return FAILED;
130-}
131- 
132-Status CollectCandidates(const AscGraph &graph, std::vector<NodePtr> &candidates) {
133- for (const auto &node : graph.GetAllNodes()) {
134- auto candidate = ToAscNode(node);
135- if (candidate != nullptr && IsOps<Broadcast>(candidate) && candidate->GetOutDataNodesSize() > 1U &&
136- !IsScalarInput(candidate)) {
137- candidates.push_back(candidate);
138- }
139- }
140- return SUCCESS;
141-}
142- 
143-Status CollectBroadcastChain(NodePtr tail, std::vector<NodePtr> &chain, NodePtr &source) {
144- NodePtr current = tail;
145- while (IsOps<Broadcast>(current) && current->GetInDataNodesSize() == 1UL) {
146- chain.push_back(current);
147- GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(current, current, 0));
148- }
149- std::reverse(chain.begin(), chain.end());
150- source = current;
151- return chain.empty() || source == nullptr || IsOps<Broadcast>(source) ? FAILED : SUCCESS;
152-}
153- 
154-Status CloneBroadcast(AscGraph &graph, const NodePtr &source, const std::string &name, NodePtr &clone) {
155- Broadcast op(name.c_str());
156- clone = graph.AddNode(op);
157- GE_ASSERT_NOTNULL(clone);
158- AscTensorAttr *attr = nullptr;
159- GE_ASSERT_SUCCESS(GetOutputTensorAttr(source, attr));
160- op.attr.sched = source->attr.sched;
161- op.attr.api.compute_type = af::ComputeType::kComputeBroadcast;
162- op.attr.api.type = af::ApiType::kAPITypeCompute;
163- op.y.dtype = attr->dtype;
164- *op.y.axis = attr->axis;
165- *op.y.repeats = attr->repeats;
166- *op.y.strides = attr->strides;
167- return SUCCESS;
168-}
169- 
170-bool AllSame(const std::vector<NodePtr> &nodes) {
171- return nodes.empty() ||
172- std::all_of(nodes.begin() + 1, nodes.end(), [&](const NodePtr &node) { return node == nodes.front(); });
173-}
174- 
175-struct BranchPlan {
176- std::vector<NodePtr> consumers;
177- std::vector<NodePtr> successors;
178- std::vector<int32_t> successor_inputs;
179-};
180- 
181-Status BuildBranchPlan(const NodePtr &tail, BranchPlan &plan) {
182- auto out = tail->GetOutDataAnchor(0);
183- GE_ASSERT_NOTNULL(out);
184- const auto &peers = out->GetPeerInDataAnchors();
185- if (peers.size() < 2U || peers.size() > 8U) {
186- return FAILED;
187- }
188- for (const auto &peer : peers) {
189- if (peer == nullptr) {
190- return FAILED;
191- }
192- auto consumer = ToAscNode(peer->GetOwnerNode());
193- if (!IsSingleInAndOutNode(consumer) || !CanBackward(consumer)) {
194- return FAILED;
195- }
196- auto successor = ToAscNode(consumer->GetOutDataNodes().at(0));
197- if (successor == nullptr || successor->GetAllInDataAnchorsSize() <= 1U) {
198- return FAILED;
199- }
200- int32_t input = -1;
201- for (uint32_t i = 0U; i < successor->GetAllInDataAnchorsSize(); ++i) {
202- auto anchor = successor->GetInDataAnchor(static_cast<int32_t>(i));
203- if (anchor != nullptr && anchor->GetPeerOutAnchor() != nullptr &&
204- ToAscNode(anchor->GetPeerOutAnchor()->GetOwnerNode()) == consumer) {
205- input = static_cast<int32_t>(i);
206- break;
207- }
208- }
209- if (input < 0) {
210- return FAILED;
211- }
212- plan.consumers.push_back(consumer);
213- plan.successors.push_back(successor);
214- plan.successor_inputs.push_back(input);
215- }
216- return SUCCESS;
217-}
218- 
219-Status BuildBranchChains(AscGraph &graph, const std::vector<NodePtr> &original, const BranchPlan &plan,
220- std::vector<std::vector<NodePtr>> &chains) {
221- chains.resize(plan.consumers.size());
222- chains[0] = original;
223- for (size_t branch = 1U; branch < chains.size(); ++branch) {
224- for (const auto &node : original) {
225- NodePtr clone;
226- GE_ASSERT_SUCCESS(
227- CloneBroadcast(graph, node, node->GetName() + "_branch_split_" + std::to_string(branch), clone));
228- chains[branch].push_back(clone);
229- }
230- }
231- return SUCCESS;
232-}
233- 
234-Status ApplyBranchPlan(const NodePtr &source, const std::vector<NodePtr> &original, const BranchPlan &plan,
235- const std::vector<std::vector<NodePtr>> &chains) {
236- GE_ASSERT_GRAPH_SUCCESS(
237- af::GraphUtils::RemoveEdge(source->GetOutDataAnchor(0), original.front()->GetInDataAnchor(0)));
238- for (size_t i = 0U; i + 1U < original.size(); ++i) {
239- GE_ASSERT_GRAPH_SUCCESS(
240- af::GraphUtils::RemoveEdge(original[i]->GetOutDataAnchor(0), original[i + 1U]->GetInDataAnchor(0)));
241- }
242- AscTensorAttr *source_attr = nullptr;
243- GE_ASSERT_SUCCESS(GetOutputTensorAttr(source, source_attr));
244- for (size_t branch = 0U; branch < plan.consumers.size(); ++branch) {
245- auto consumer = plan.consumers[branch];
246- auto successor = plan.successors[branch];
247- GE_ASSERT_GRAPH_SUCCESS(
248- af::GraphUtils::RemoveEdge(original.back()->GetOutDataAnchor(0), consumer->GetInDataAnchor(0)));
249- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(consumer->GetOutDataAnchor(0),
250- successor->GetInDataAnchor(plan.successor_inputs[branch])));
251- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(source->GetOutDataAnchor(0), consumer->GetInDataAnchor(0)));
252- const auto &chain = chains[branch];
253- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(consumer->GetOutDataAnchor(0), chain.front()->GetInDataAnchor(0)));
254- for (size_t i = 0U; i + 1U < chain.size(); ++i) {
255- GE_ASSERT_GRAPH_SUCCESS(
256- af::GraphUtils::AddEdge(chain[i]->GetOutDataAnchor(0), chain[i + 1U]->GetInDataAnchor(0)));
257- }
258- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(chain.back()->GetOutDataAnchor(0),
259- successor->GetInDataAnchor(plan.successor_inputs[branch])));
260- AscTensorAttr *consumer_attr = nullptr;
261- GE_ASSERT_SUCCESS(GetOutputTensorAttr(consumer, consumer_attr));
262- consumer_attr->axis = source_attr->axis;
263- consumer_attr->repeats = source_attr->repeats;
264- consumer_attr->strides = source_attr->strides;
265- }
266- return SUCCESS;
267-}
268- 
269-Status SplitOneBranch(AscGraph &graph, const NodePtr &tail, bool &changed) {
270- changed = false;
271- std::vector<NodePtr> original;
272- NodePtr source;
273- if (CollectBroadcastChain(tail, original, source) != SUCCESS) {
274- return SUCCESS;
275- }
276- BranchPlan plan;
277- if (BuildBranchPlan(tail, plan) != SUCCESS || AllSame(plan.successors)) {
278- return SUCCESS;
279- }
280- std::vector<std::vector<NodePtr>> chains;
281- GE_ASSERT_SUCCESS(BuildBranchChains(graph, original, plan, chains));
282- GE_ASSERT_SUCCESS(ApplyBranchPlan(source, original, plan, chains));
283- changed = true;
284- return SUCCESS;
285-}
286- 
287-struct ConsumerPlan {
288- NodePtr source;
289- std::vector<NodePtr> consumers;
290- std::vector<int32_t> inputs;
291-};
292- 
293-Status BuildConsumerPlan(const NodePtr &broadcast, ConsumerPlan &plan) {
294- GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(broadcast, plan.source, 0));
295- if (plan.source == nullptr || IsOps<Broadcast>(plan.source)) {
296- return FAILED;
297- }
298- auto out = broadcast->GetOutDataAnchor(0);
299- GE_ASSERT_NOTNULL(out);
300- const auto &peers = out->GetPeerInDataAnchors();
301- if (peers.size() <= 1U) {
302- return FAILED;
303- }
304- for (const auto &peer : peers) {
305- auto consumer = peer == nullptr ? nullptr : ToAscNode(peer->GetOwnerNode());
306- if (consumer == nullptr || !CanSharedConsumer(consumer)) {
307- return FAILED;
308- }
309- plan.consumers.push_back(consumer);
310- plan.inputs.push_back(peer->GetIdx());
311- }
312- if (AllSame(plan.consumers)) {
313- return FAILED;
314- }
315- std::vector<NodePtr> merges;
316- bool trace_failed = false;
317- for (const auto &consumer : plan.consumers) {
318- NodePtr merge;
319- if (TraceMergeNode(consumer, merge) != SUCCESS) {
320- trace_failed = true;
321- break;
322- }
323- merges.push_back(merge);
324- }
325- if (!trace_failed && AllSame(merges)) {
326- return FAILED;
327- }
328- return SUCCESS;
329-}
330- 
331-Status SplitOneConsumer(AscGraph &graph, const NodePtr &broadcast, bool &changed) {
332- changed = false;
333- ConsumerPlan plan;
334- if (BuildConsumerPlan(broadcast, plan) != SUCCESS) {
335- return SUCCESS;
336- }
337- for (size_t i = 1U; i < plan.consumers.size(); ++i) {
338- NodePtr clone;
339- GE_ASSERT_SUCCESS(
340- CloneBroadcast(graph, broadcast, broadcast->GetName() + "_consumer_split_" + std::to_string(i), clone));
341- auto consumer_input = plan.consumers[i]->GetInDataAnchor(plan.inputs[i]);
342- GE_ASSERT_NOTNULL(consumer_input);
343- auto peer_out = consumer_input->GetPeerOutAnchor();
344- GE_ASSERT_NOTNULL(peer_out);
345- auto source_out = plan.source->GetOutDataAnchor(0);
346- GE_ASSERT_NOTNULL(source_out);
347- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(peer_out, consumer_input));
348- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(source_out, clone->GetInDataAnchor(0)));
349- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(clone->GetOutDataAnchor(0), consumer_input));
350- }
351- changed = true;
352- return SUCCESS;
353-}
354- 
355-} // namespace
356- 
357-Status SplitSharedBroadcastBranches(AscGraph &graph) {
358- std::vector<NodePtr> candidates;
359- GE_ASSERT_SUCCESS(CollectCandidates(graph, candidates));
360- bool changed = false;
361- for (const auto &candidate : candidates) {
362- bool candidate_changed = false;
363- GE_ASSERT_SUCCESS(SplitOneBranch(graph, candidate, candidate_changed));
364- changed = changed || candidate_changed;
365- }
366- if (changed) {
367- GE_ASSERT_SUCCESS(ScheduleUtils::TopologicalSorting(graph));
368- }
369- return SUCCESS;
370-}
371- 
372-Status SplitSharedBroadcastConsumers(AscGraph &graph) {
373- std::vector<NodePtr> candidates;
374- GE_ASSERT_SUCCESS(CollectCandidates(graph, candidates));
375- bool changed = false;
376- for (const auto &candidate : candidates) {
377- bool candidate_changed = false;
378- GE_ASSERT_SUCCESS(SplitOneConsumer(graph, candidate, candidate_changed));
379- changed = changed || candidate_changed;
380- }
381- if (changed) {
382- GE_ASSERT_SUCCESS(ScheduleUtils::TopologicalSorting(graph));
383- }
384- return SUCCESS;
385-}
386- 
387-} // namespace broadcast_backward_shared_split
388-} // namespace optimize
@@ -1,20 +0,0 @@
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 of
4- * CANN Open Software License Agreement Version 2.0.
5- */
6-#ifndef OPTIMIZE_GRAPH_PASS_BROADCAST_BACKWARD_SHARED_SPLIT_H
7-#define OPTIMIZE_GRAPH_PASS_BROADCAST_BACKWARD_SHARED_SPLIT_H
8- 
9-#include "base_graph_pass.h"
10- 
11-namespace optimize {
12-namespace broadcast_backward_shared_split {
13- 
14-Status SplitSharedBroadcastBranches(af::AscGraph &graph);
15-Status SplitSharedBroadcastConsumers(af::AscGraph &graph);
16- 
17-} // namespace broadcast_backward_shared_split
18-} // namespace optimize
19- 
20-#endif // OPTIMIZE_GRAPH_PASS_BROADCAST_BACKWARD_SHARED_SPLIT_H
@@ -12,7 +12,6 @@
12#define AUTOFUSE_PATTERN_FUSION_UNITTEST_CPP_O_D_PASS_RUNNER_V1_H12#define AUTOFUSE_PATTERN_FUSION_UNITTEST_CPP_O_D_PASS_RUNNER_V1_H
13 13 
14#include "optimize/platform/common/pass_runner.h"14#include "optimize/platform/common/pass_runner.h"
15-#include "optimize/graph_pass/broadcast_backward_pass.h"
16#include "optimize/graph_pass/broadcast_const_to_store.h"15#include "optimize/graph_pass/broadcast_const_to_store.h"
17#include "optimize/graph_pass/duplicate_elewise_cse_pass.h"16#include "optimize/graph_pass/duplicate_elewise_cse_pass.h"
18#include "optimize/graph_pass/scalar_to_1d_tensor.h"17#include "optimize/graph_pass/scalar_to_1d_tensor.h"
@@ -28,8 +27,6 @@ class PassRunnerV1 final : public BasePassRunner {
28 this->RegisterPass<PowEquivSubstitutionPass>();27 this->RegisterPass<PowEquivSubstitutionPass>();
29 this->RegisterPass<BroadcastConstToStorePass>();28 this->RegisterPass<BroadcastConstToStorePass>();
30 this->RegisterPass<ScalarTo1DTensorPass>();29 this->RegisterPass<ScalarTo1DTensorPass>();
31- // The sched/tensor axes must be complete before moving Broadcasts; scalar Broadcast cleanup runs afterward.
32- this->RegisterPass<BroadcastBackwardPass>();
33 this->RegisterPass<ScalarBroadcastOptimizationPass>();30 this->RegisterPass<ScalarBroadcastOptimizationPass>();
34 this->RegisterPass<MaskedFillInputReorderPass>();31 this->RegisterPass<MaskedFillInputReorderPass>();
35 this->RegisterPass<ExpandDimsForAllReducePass>();32 this->RegisterPass<ExpandDimsForAllReducePass>();
@@ -1,12 +0,0 @@
1-# -----------------------------------------------------------------------------------------------------------
2-# Copyright (c) 2026 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 FILE 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-add_library(broadcast_backward_test_utils INTERFACE)
12-target_include_directories(broadcast_backward_test_utils INTERFACE ${CMAKE_CURRENT_SOURCE_DIR})
@@ -1,99 +0,0 @@
1-/**
2- * Copyright (c) 2026 Huawei Technologies Co., Ltd.
3- * This program is free software, you can redistribute it and/or modify it under the terms of
4- * CANN Open Software License Agreement Version 2.0 (the "License").
5- * Please refer to the License for the specific language governing permissions and limitations under 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-#ifndef AUTOFUSE_TESTS_FRAMEWORK_BROADCAST_BACKWARD_TEST_UTILS_H_
11-#define AUTOFUSE_TESTS_FRAMEWORK_BROADCAST_BACKWARD_TEST_UTILS_H_
12- 
13-#include "gtest/gtest.h"
14- 
15-#include <string>
16- 
17-#include "asc_graph_builder.h"
18-#include "ascgraph_info_complete.h"
19-#include "optimize/graph_pass/broadcast_backward_pass.h"
20-#include "platform_context.h"
21- 
22-namespace broadcast_backward_test {
23- 
24-inline af::testing::AscGraphBuilder BuildCommonAxisGraphPrefix(const std::string &name) {
25- const auto s0 = af::testing::Sym("s0");
26- const auto s1 = af::testing::Sym("s1");
27- const auto s2 = af::testing::Sym("s2");
28- af::testing::AscGraphBuilder builder(name);
29- builder.Loops({s0, s1, s2})
30- .Data("data0", 0)
31- .Data("data1", 1)
32- .Load("load0", "data0", {af::sym::kSymbolOne, af::sym::kSymbolOne, s2},
33- {af::sym::kSymbolZero, af::sym::kSymbolZero, af::sym::kSymbolOne})
34- .Load("load1", "data1", {s0, af::sym::kSymbolOne, af::sym::kSymbolOne},
35- {af::sym::kSymbolOne, af::sym::kSymbolZero, af::sym::kSymbolZero})
36- .Broadcast("broadcast0", "load0", {0, 1})
37- .Broadcast("broadcast1", "load1", {1, 2});
38- return builder;
39-}
40- 
41-inline af::AscGraph BuildCommonAxisGraph(const std::string &name) {
42- return BuildCommonAxisGraphPrefix(name)
43- .Add("merge", "broadcast0", "broadcast1")
44- .Store("store", "merge")
45- .Output("output", "store")
46- .Build();
47-}
48- 
49-inline af::AscGraph BuildDtypeAwareCommonAxisGraph(const std::string &name) {
50- return BuildCommonAxisGraphPrefix(name)
51- .Abs("abs", "broadcast0")
52- .Cast("cast0", "abs", af::DT_FLOAT16)
53- .Cast("cast1", "broadcast1", af::DT_FLOAT16)
54- .Relu("relu", "cast1")
55- .Add("merge", "cast0", "relu")
56- .Store("store", "merge")
57- .Output("output", "store")
58- .Build();
59-}
60- 
61-inline af::AscNodePtr FindNode(af::AscGraph &graph, const std::string &name) {
62- for (const auto &node : graph.GetAllNodes()) {
63- if (node->GetName() == name) {
64- return std::dynamic_pointer_cast<af::AscNode>(node);
65- }
66- }
67- return nullptr;
68-}
69- 
70-inline std::string GetInputNodeName(const af::AscNodePtr &node, size_t input_index = 0U) {
71- if (node == nullptr || input_index >= node->GetInDataNodesSize()) {
72- return {};
73- }
74- const auto input_node = node->GetInDataNodes().at(input_index);
75- return input_node == nullptr ? std::string() : input_node->GetName();
76-}
77- 
78-inline void CompleteApiInfo(af::AscGraph &graph) {
79- ge::PlatformContext::GetInstance().Reset();
80- ge::PlatformContext::GetInstance().SetPlatform("3510");
81- ASSERT_EQ(optimize::AscGraphInfoComplete::CompleteApiInfo(graph), af::SUCCESS);
82-}
83- 
84-inline void SetNodeDtype(af::AscGraph &graph, const std::string &name, af::DataType dtype) {
85- const auto node = FindNode(graph, name);
86- ASSERT_NE(node, nullptr);
87- const auto op_desc = node->GetOpDesc();
88- ASSERT_NE(op_desc, nullptr);
89- for (size_t input_index = 0U; input_index < node->GetAllInDataAnchorsSize(); ++input_index) {
90- const auto input_desc = op_desc->MutableInputDesc(static_cast<uint32_t>(input_index));
91- ASSERT_NE(input_desc, nullptr);
92- input_desc->SetDataType(dtype);
93- }
94- node->outputs[0].attr.dtype = dtype;
95-}
96- 
97-} // namespace broadcast_backward_test
98- 
99-#endif // AUTOFUSE_TESTS_FRAMEWORK_BROADCAST_BACKWARD_TEST_UTILS_H_
@@ -1,378 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#ifndef AUTOFUSE_TESTS_FRAMEWORK_BROADCAST_BACKWARD_UT_UTILS_H_
11-#define AUTOFUSE_TESTS_FRAMEWORK_BROADCAST_BACKWARD_UT_UTILS_H_
12- 
13-#include "gtest/gtest.h"
14- 
15-#include <string>
16-#include <vector>
17- 
18-#include "asc_graph_builder.h"
19-#include "graph_utils.h"
20-#include "optimize/graph_pass/broadcast_backward_pass.h"
21-#include "tests/framework/broadcast_backward/broadcast_backward_test_utils.h"
22- 
23-namespace broadcast_backward_test {
24- 
25-inline const std::vector<af::Expression> kCompactRepeats = {af::testing::Sym("s0"), af::sym::kSymbolOne};
26-inline const std::vector<af::Expression> kCompactStrides = {af::sym::kSymbolOne, af::sym::kSymbolZero};
27- 
28-inline bool IsConnected(af::AscGraph &graph, const char *src_name, const char *dst_name) {
29- for (const auto &node : graph.GetAllNodes()) {
30- if (node->GetName() != src_name) {
31- continue;
32- }
33- const auto out_anchor = node->GetOutDataAnchor(0);
34- if (out_anchor == nullptr) {
35- return false;
36- }
37- for (const auto &peer : out_anchor->GetPeerInDataAnchors()) {
38- if (peer != nullptr && peer->GetOwnerNode()->GetName() == dst_name) {
39- return true;
40- }
41- }
42- }
43- return false;
44-}
45- 
46-inline bool AreExpressionVectorsEqual(const std::vector<af::Expression> &lhs, const std::vector<af::Expression> &rhs) {
47- if (lhs.size() != rhs.size()) {
48- return false;
49- }
50- for (size_t index = 0U; index < lhs.size(); ++index) {
51- if (af::SymbolicUtils::StaticCheckEq(lhs[index], rhs[index]) != af::TriBool::kTrue) {
52- return false;
53- }
54- }
55- return true;
56-}
57- 
58-inline bool AreConnectedTensorAttrsEqual(const af::AscNodePtr &source_node, const af::AscNodePtr &destination_node,
59- size_t destination_input_index) {
60- if (source_node == nullptr || destination_node == nullptr || source_node->GetOpDesc() == nullptr ||
61- destination_node->GetOpDesc() == nullptr ||
62- destination_input_index >= destination_node->GetAllInDataAnchorsSize()) {
63- return false;
64- }
65- const auto source_anchor = source_node->GetOutDataAnchor(0);
66- const auto destination_anchor = destination_node->GetInDataAnchor(static_cast<int32_t>(destination_input_index));
67- return source_anchor != nullptr && destination_anchor != nullptr &&
68- destination_anchor->GetPeerOutAnchor() == source_anchor &&
69- source_node->GetOpDesc()->GetOutputDesc(0U).GetDataType() ==
70- destination_node->GetOpDesc()->GetInputDesc(static_cast<uint32_t>(destination_input_index)).GetDataType();
71-}
72- 
73-inline bool IsEdgeAttrConsistent(af::AscGraph &graph, const char *src_name, const char *dst_name) {
74- for (const auto &node : graph.GetAllNodes()) {
75- if (node->GetName() != src_name || node->GetOutDataAnchor(0) == nullptr) {
76- continue;
77- }
78- for (const auto &peer : node->GetOutDataAnchor(0)->GetPeerInDataAnchors()) {
79- if (peer == nullptr || peer->GetOwnerNode()->GetName() != dst_name) {
80- continue;
81- }
82- const auto destination_node = std::dynamic_pointer_cast<af::AscNode>(peer->GetOwnerNode());
83- const auto source_node = std::dynamic_pointer_cast<af::AscNode>(node);
84- return AreConnectedTensorAttrsEqual(source_node, destination_node, static_cast<size_t>(peer->GetIdx()));
85- }
86- }
87- return false;
88-}
89- 
90-inline bool HasNode(af::AscGraph &graph, const char *node_name) {
91- for (const auto &node : graph.GetAllNodes()) {
92- if (node->GetName() == node_name) {
93- return true;
94- }
95- }
96- return false;
97-}
98- 
99-inline void ExpectStaticEq(const std::vector<af::Expression> &actual, const std::vector<af::Expression> &expected) {
100- ASSERT_EQ(actual.size(), expected.size());
101- for (size_t index = 0U; index < actual.size(); ++index) {
102- EXPECT_EQ(af::SymbolicUtils::StaticCheckEq(actual[index], expected[index]), af::TriBool::kTrue);
103- }
104-}
105- 
106-inline void ExpectResidualConnections(af::AscGraph &graph) {
107- EXPECT_TRUE(IsConnected(graph, "load0", "broadcast0_residual_0"));
108- EXPECT_TRUE(IsConnected(graph, "broadcast0_residual_0", "merge"));
109- EXPECT_TRUE(IsConnected(graph, "load1", "broadcast1_residual_1"));
110- EXPECT_TRUE(IsConnected(graph, "broadcast1_residual_1", "merge"));
111-}
112- 
113-inline af::AscGraph BuildUnaryGraph(const std::string &op_name) {
114- af::testing::AscGraphBuilder builder("broadcast_backward_unary_" + op_name);
115- builder.Loops({af::testing::Sym("s0"), af::testing::Sym("s1")})
116- .Data("data", 0)
117- .Load("load", "data", kCompactRepeats, kCompactStrides)
118- .Broadcast("broadcast", "load", {af::testing::Sym("s0"), af::testing::Sym("s1")});
119- if (op_name == "Abs") {
120- builder.Abs("compute", "broadcast");
121- } else if (op_name == "Neg") {
122- builder.Neg("compute", "broadcast");
123- } else if (op_name == "Exp") {
124- builder.Exp("compute", "broadcast");
125- } else if (op_name == "Sqrt") {
126- builder.Sqrt("compute", "broadcast");
127- } else if (op_name == "Relu") {
128- builder.Relu("compute", "broadcast");
129- } else if (op_name == "Rsqrt") {
130- builder.Op<af::ascir_op::Rsqrt>("compute", {"broadcast"});
131- } else if (op_name == "Reciprocal") {
132- builder.Op<af::ascir_op::Reciprocal>("compute", {"broadcast"});
133- } else if (op_name == "Erf") {
134- builder.Op<af::ascir_op::Erf>("compute", {"broadcast"});
135- } else if (op_name == "Sign") {
136- builder.Op<af::ascir_op::Sign>("compute", {"broadcast"});
137- } else if (op_name == "Tanh") {
138- builder.Op<af::ascir_op::Tanh>("compute", {"broadcast"});
139- } else if (op_name == "Ln") {
140- builder.Op<af::ascir_op::Ln>("compute", {"broadcast"});
141- } else {
142- ADD_FAILURE() << "Unsupported unary test operator: " << op_name;
143- }
144- return builder.Store("store", "compute").Output("output", "store").Build();
145-}
146- 
147-inline af::AscGraph BuildBinaryGraph(const std::string &op_name) {
148- af::testing::AscGraphBuilder builder("broadcast_backward_binary_" + op_name);
149- builder.Loops({af::testing::Sym("s0"), af::testing::Sym("s1")})
150- .Data("data0", 0)
151- .Data("data1", 1)
152- .Load("load0", "data0", kCompactRepeats, kCompactStrides)
153- .Load("load1", "data1", kCompactRepeats, kCompactStrides)
154- .Broadcast("broadcast0", "load0", {af::testing::Sym("s0"), af::testing::Sym("s1")})
155- .Broadcast("broadcast1", "load1", {af::testing::Sym("s0"), af::testing::Sym("s1")});
156- if (op_name == "Add") {
157- builder.Add("compute", "broadcast0", "broadcast1");
158- } else if (op_name == "Sub") {
159- builder.Sub("compute", "broadcast0", "broadcast1");
160- } else if (op_name == "Mul") {
161- builder.Mul("compute", "broadcast0", "broadcast1");
162- } else if (op_name == "Div") {
163- builder.Div("compute", "broadcast0", "broadcast1");
164- } else if (op_name == "Minimum") {
165- builder.Minimum("compute", "broadcast0", "broadcast1");
166- } else if (op_name == "Maximum") {
167- builder.Maximum("compute", "broadcast0", "broadcast1");
168- } else {
169- ADD_FAILURE() << "Unsupported binary test operator: " << op_name;
170- }
171- return builder.Store("store", "compute").Output("output", "store").Build();
172-}
173- 
174-inline af::AscGraph BuildDtypeAwareBinaryGraph(const std::string &op_name) {
175- af::testing::AscGraphBuilder builder("broadcast_backward_dtype_aware_" + op_name);
176- builder.Loops({af::testing::Sym("s0"), af::testing::Sym("s1")})
177- .Data("data0", 0)
178- .Data("data1", 1)
179- .Load("load0", "data0", kCompactRepeats, kCompactStrides)
180- .Load("load1", "data1", kCompactRepeats, kCompactStrides)
181- .Broadcast("broadcast0", "load0", {af::testing::Sym("s0"), af::testing::Sym("s1")})
182- .Broadcast("broadcast1", "load1", {af::testing::Sym("s0"), af::testing::Sym("s1")});
183- if (op_name == "Eq") {
184- builder.Op<af::ascir_op::Eq>("compute", {"broadcast0", "broadcast1"});
185- } else if (op_name == "TrueDiv") {
186- builder.Op<af::ascir_op::TrueDiv>("compute", {"broadcast0", "broadcast1"});
187- } else {
188- ADD_FAILURE() << "Unsupported dtype-aware binary test operator: " << op_name;
189- }
190- return builder.Store("store", "compute").Output("output", "store").Build();
191-}
192- 
193-inline void ExpectBinaryBroadcastMove(af::AscGraph &graph) {
194- EXPECT_TRUE(IsConnected(graph, "load0", "compute"));
195- EXPECT_TRUE(IsConnected(graph, "load1", "compute"));
196- EXPECT_TRUE(IsConnected(graph, "compute", "broadcast0"));
197- EXPECT_TRUE(IsConnected(graph, "broadcast0", "store"));
198- EXPECT_FALSE(HasNode(graph, "broadcast1"));
199-}
200- 
201-inline void ExpectDtypeAwareBinaryMove(af::AscGraph &graph, const std::string &op_name) {
202- CompleteApiInfo(graph);
203- optimize::BroadcastBackwardPass pass;
204- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS) << op_name;
205- EXPECT_TRUE(IsConnected(graph, "load0", "compute"));
206- EXPECT_TRUE(IsConnected(graph, "load1", "compute"));
207- EXPECT_TRUE(IsConnected(graph, "compute", "broadcast0"));
208- EXPECT_TRUE(IsConnected(graph, "broadcast0", "store"));
209- EXPECT_FALSE(HasNode(graph, "broadcast1"));
210- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load0", "compute"));
211- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load1", "compute"));
212- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "compute", "broadcast0"));
213- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast0", "store"));
214-}
215- 
216-inline void ExpectCommonAxisLayouts(af::AscGraph &graph, const std::vector<af::Expression> &compact_repeats,
217- const std::vector<af::Expression> &expanded_repeats,
218- const std::vector<af::Expression> &compact_strides,
219- const std::vector<af::Expression> &expanded_strides) {
220- const auto merge = FindNode(graph, "merge");
221- const auto common = FindNode(graph, "merge_broadcast_backward_common");
222- const auto residual0 = FindNode(graph, "broadcast0_residual_0");
223- const auto residual1 = FindNode(graph, "broadcast1_residual_1");
224- ASSERT_NE(merge, nullptr);
225- ASSERT_NE(common, nullptr);
226- ASSERT_NE(residual0, nullptr);
227- ASSERT_NE(residual1, nullptr);
228- ExpectStaticEq(merge->inputs[0].attr.repeats, compact_repeats);
229- ExpectStaticEq(merge->inputs[1].attr.repeats, compact_repeats);
230- ExpectStaticEq(merge->outputs[0].attr.repeats, compact_repeats);
231- ExpectStaticEq(common->outputs[0].attr.repeats, expanded_repeats);
232- ExpectStaticEq(residual0->outputs[0].attr.strides, compact_strides);
233- ExpectStaticEq(residual1->outputs[0].attr.strides, compact_strides);
234- ExpectStaticEq(merge->outputs[0].attr.strides, compact_strides);
235- ExpectStaticEq(common->outputs[0].attr.strides, expanded_strides);
236-}
237- 
238-inline af::AscGraph BuildDirectFanOutGraph(const std::string &name) {
239- const auto s0 = af::testing::Sym("s0");
240- const auto s1 = af::testing::Sym("s1");
241- return af::testing::AscGraphBuilder(name)
242- .Loops({s0, s1})
243- .Data("data", 0)
244- .Load("load", "data", kCompactRepeats, kCompactStrides)
245- .Broadcast("broadcast", "load", {1})
246- .Abs("branch0", "broadcast")
247- .Neg("branch1", "broadcast")
248- .Add("merge", "branch0", "branch1")
249- .Store("store", "merge")
250- .Output("output", "store")
251- .Build();
252-}
253- 
254-inline af::AscGraph BuildMultiNodeFanOutGraph(const std::string &name) {
255- const auto s0 = af::testing::Sym("s0");
256- const auto s1 = af::testing::Sym("s1");
257- return af::testing::AscGraphBuilder(name)
258- .Loops({s0, s1})
259- .Data("data", 0)
260- .Load("load", "data", kCompactRepeats, kCompactStrides)
261- .Broadcast("broadcast", "load", {1})
262- .Abs("branch0_head", "broadcast")
263- .Relu("branch0_tail", "branch0_head")
264- .Neg("branch1_head", "broadcast")
265- .Exp("branch1_tail", "branch1_head")
266- .Add("merge", "branch0_tail", "branch1_tail")
267- .Store("store", "merge")
268- .Output("output", "store")
269- .Build();
270-}
271- 
272-inline af::AscGraph BuildScalarForkJoinGraph(const std::string &name) {
273- const auto s0 = af::testing::Sym("s0");
274- const auto s1 = af::testing::Sym("s1");
275- return af::testing::AscGraphBuilder(name)
276- .Loops({s0, s1})
277- .Scalar("scalar", "1.0")
278- .Broadcast("broadcast", "scalar", {s0, s1})
279- .Abs("left", "broadcast")
280- .Neg("right", "broadcast")
281- .Add("merge", "left", "right")
282- .Store("store", "merge")
283- .Output("output", "store")
284- .Build();
285-}
286- 
287-inline af::AscGraph BuildSharedDtypeAwareFanOutGraph(const std::string &name) {
288- const auto s0 = af::testing::Sym("s0");
289- const auto s1 = af::testing::Sym("s1");
290- const auto s2 = af::testing::Sym("s2");
291- const std::vector<af::Expression> compact = {s0, af::sym::kSymbolOne, s2};
292- const std::vector<af::Expression> strides = {s2, af::sym::kSymbolZero, af::sym::kSymbolOne};
293- return af::testing::AscGraphBuilder(name)
294- .Loops({s0, s1, s2})
295- .Data("data", 0, af::DT_INT32)
296- .Load("load", "data", compact, strides)
297- .Broadcast("broadcast", "load", {1})
298- .Abs("abs", "broadcast")
299- .Cast("left_cast", "abs", af::DT_FLOAT)
300- .Cast("right_cast", "broadcast", af::DT_FLOAT)
301- .Relu("relu", "right_cast")
302- .Add("add", "relu", "left_cast")
303- .Sqrt("sqrt", "add")
304- .Op<af::ascir_op::Sigmoid>("sigmoid", {"sqrt"})
305- .Store("store", "sigmoid")
306- .Output("output", "store")
307- .Build();
308-}
309- 
310-inline void CompleteSharedDtypeAwareFanOutGraph(af::AscGraph &graph) {
311- CompleteApiInfo(graph);
312- for (const auto *node_name : {"load", "broadcast", "abs"}) {
313- SetNodeDtype(graph, node_name, af::DT_INT32);
314- }
315- for (const auto *node_name : {"left_cast", "right_cast"}) {
316- const auto cast = FindNode(graph, node_name);
317- ASSERT_NE(cast, nullptr);
318- cast->GetOpDesc()->MutableInputDesc(0U)->SetDataType(af::DT_INT32);
319- }
320- for (const auto *node_name : {"relu", "add", "sqrt", "sigmoid", "store"}) {
321- SetNodeDtype(graph, node_name, af::DT_FLOAT);
322- }
323-}
324- 
325-inline void ExpectDirectFanOutCandidate(af::AscGraph &graph) {
326- const auto broadcast = FindNode(graph, "broadcast");
327- const auto branch0 = FindNode(graph, "branch0");
328- const auto branch1 = FindNode(graph, "branch1");
329- const auto merge = FindNode(graph, "merge");
330- ASSERT_NE(broadcast, nullptr);
331- ASSERT_NE(branch0, nullptr);
332- ASSERT_NE(branch1, nullptr);
333- ASSERT_NE(merge, nullptr);
334- ASSERT_EQ(broadcast->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 2U);
335- ASSERT_EQ(branch0->GetAllInDataAnchorsSize(), 1U);
336- ASSERT_EQ(branch0->GetOutDataNodesSize(), 1U);
337- ASSERT_EQ(branch1->GetAllInDataAnchorsSize(), 1U);
338- ASSERT_EQ(branch1->GetOutDataNodesSize(), 1U);
339- ASSERT_EQ(merge->GetAllInDataAnchorsSize(), 2U);
340- ASSERT_EQ(merge->GetOutDataNodesSize(), 1U);
341- ASSERT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "branch0"));
342- ASSERT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "branch1"));
343- ASSERT_TRUE(IsEdgeAttrConsistent(graph, "branch0", "merge"));
344- ASSERT_TRUE(IsEdgeAttrConsistent(graph, "branch1", "merge"));
345- ASSERT_TRUE(IsEdgeAttrConsistent(graph, "merge", "store"));
346- ASSERT_TRUE(AreExpressionVectorsEqual(merge->outputs[0].attr.repeats, broadcast->outputs[0].attr.repeats));
347- ASSERT_EQ(merge->outputs[0].attr.axis, broadcast->outputs[0].attr.axis);
348- ASSERT_EQ(merge->outputs[0].attr.dtype, broadcast->outputs[0].attr.dtype);
349- ASSERT_TRUE(AreExpressionVectorsEqual(merge->outputs[0].attr.strides, broadcast->outputs[0].attr.strides));
350-}
351- 
352-inline void ExpectDirectFanOutMoved(af::AscGraph &graph) {
353- EXPECT_TRUE(IsConnected(graph, "load", "branch0"));
354- EXPECT_TRUE(IsConnected(graph, "load", "branch1"));
355- EXPECT_TRUE(IsConnected(graph, "branch0", "merge"));
356- EXPECT_TRUE(IsConnected(graph, "branch1", "merge"));
357- EXPECT_TRUE(IsConnected(graph, "merge", "broadcast"));
358- EXPECT_TRUE(IsConnected(graph, "broadcast", "store"));
359- const auto branch0 = FindNode(graph, "branch0");
360- const auto branch1 = FindNode(graph, "branch1");
361- const auto merge = FindNode(graph, "merge");
362- ASSERT_NE(branch0, nullptr);
363- ASSERT_NE(branch1, nullptr);
364- ASSERT_NE(merge, nullptr);
365- ExpectStaticEq(branch0->inputs[0].attr.repeats, kCompactRepeats);
366- ExpectStaticEq(branch1->inputs[0].attr.repeats, kCompactRepeats);
367- ExpectStaticEq(merge->outputs[0].attr.repeats, kCompactRepeats);
368- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load", "branch0"));
369- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load", "branch1"));
370- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch0", "merge"));
371- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch1", "merge"));
372- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "merge", "broadcast"));
373- EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "store"));
374-}
375- 
376-} // namespace broadcast_backward_test
377- 
378-#endif // AUTOFUSE_TESTS_FRAMEWORK_BROADCAST_BACKWARD_UT_UTILS_H_
@@ -1896,7 +1896,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) {
1896 auto impl_grp_0_brc4 = impl_graphs[0].FindNode("brc4");1896 auto impl_grp_0_brc4 = impl_graphs[0].FindNode("brc4");
1897 EXPECT_NE(impl_grp_0_brc4, nullptr);1897 EXPECT_NE(impl_grp_0_brc4, nullptr);
1898 EXPECT_EQ(impl_grp_0_brc4->GetAllInDataAnchorsSize(), 1);1898 EXPECT_EQ(impl_grp_0_brc4->GetAllInDataAnchorsSize(), 1);
1899- EXPECT_EQ(impl_grp_0_brc4->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0");1899+ EXPECT_EQ(impl_grp_0_brc4->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0");
1900 1900 
1901 EXPECT_EQ(impl_graphs[1].FindNode("brc0"), nullptr);1901 EXPECT_EQ(impl_graphs[1].FindNode("brc0"), nullptr);
1902 EXPECT_EQ(impl_graphs[1].FindNode("brc1"), nullptr);1902 EXPECT_EQ(impl_graphs[1].FindNode("brc1"), nullptr);
@@ -1905,7 +1905,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) {
1905 auto impl_grp_1_brc3 = impl_graphs[1].FindNode("brc3");1905 auto impl_grp_1_brc3 = impl_graphs[1].FindNode("brc3");
1906 EXPECT_NE(impl_grp_1_brc3, nullptr);1906 EXPECT_NE(impl_grp_1_brc3, nullptr);
1907 EXPECT_EQ(impl_grp_1_brc3->GetAllInDataAnchorsSize(), 1);1907 EXPECT_EQ(impl_grp_1_brc3->GetAllInDataAnchorsSize(), 1);
1908- EXPECT_EQ(impl_grp_1_brc3->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0");1908+ EXPECT_EQ(impl_grp_1_brc3->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0");
1909 1909 
1910 EXPECT_EQ(impl_graphs[2].FindNode("brc0"), nullptr);1910 EXPECT_EQ(impl_graphs[2].FindNode("brc0"), nullptr);
1911 EXPECT_EQ(impl_graphs[2].FindNode("brc3"), nullptr);1911 EXPECT_EQ(impl_graphs[2].FindNode("brc3"), nullptr);
@@ -1913,7 +1913,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) {
1913 auto impl_grp_2_brc2 = impl_graphs[2].FindNode("brc2");1913 auto impl_grp_2_brc2 = impl_graphs[2].FindNode("brc2");
1914 EXPECT_NE(impl_grp_2_brc2, nullptr);1914 EXPECT_NE(impl_grp_2_brc2, nullptr);
1915 EXPECT_EQ(impl_grp_2_brc2->GetAllInDataAnchorsSize(), 1);1915 EXPECT_EQ(impl_grp_2_brc2->GetAllInDataAnchorsSize(), 1);
1916- EXPECT_EQ(impl_grp_2_brc2->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0");1916+ EXPECT_EQ(impl_grp_2_brc2->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0");
1917 1917 
1918 EXPECT_EQ(impl_graphs[3].FindNode("brc0"), nullptr);1918 EXPECT_EQ(impl_graphs[3].FindNode("brc0"), nullptr);
1919 EXPECT_EQ(impl_graphs[3].FindNode("brc2"), nullptr);1919 EXPECT_EQ(impl_graphs[3].FindNode("brc2"), nullptr);
@@ -1922,7 +1922,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) {
1922 auto impl_grp_3_brc1 = impl_graphs[3].FindNode("brc1");1922 auto impl_grp_3_brc1 = impl_graphs[3].FindNode("brc1");
1923 EXPECT_NE(impl_grp_3_brc1, nullptr);1923 EXPECT_NE(impl_grp_3_brc1, nullptr);
1924 EXPECT_EQ(impl_grp_3_brc1->GetAllInDataAnchorsSize(), 1);1924 EXPECT_EQ(impl_grp_3_brc1->GetAllInDataAnchorsSize(), 1);
1925- EXPECT_EQ(impl_grp_3_brc1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0");1925+ EXPECT_EQ(impl_grp_3_brc1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0");
1926 1926 
1927 EXPECT_EQ(impl_graphs[4].FindNode("brc1"), nullptr);1927 EXPECT_EQ(impl_graphs[4].FindNode("brc1"), nullptr);
1928 EXPECT_EQ(impl_graphs[4].FindNode("brc2"), nullptr);1928 EXPECT_EQ(impl_graphs[4].FindNode("brc2"), nullptr);
@@ -1931,7 +1931,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) {
1931 auto impl_grp_4_exp0 = impl_graphs[4].FindNode("exp0");1931 auto impl_grp_4_exp0 = impl_graphs[4].FindNode("exp0");
1932 EXPECT_NE(impl_grp_4_exp0, nullptr);1932 EXPECT_NE(impl_grp_4_exp0, nullptr);
1933 EXPECT_EQ(impl_grp_4_exp0->GetAllInDataAnchorsSize(), 1);1933 EXPECT_EQ(impl_grp_4_exp0->GetAllInDataAnchorsSize(), 1);
1934- EXPECT_EQ(impl_grp_4_exp0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0");1934+ EXPECT_EQ(impl_grp_4_exp0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc0");
1935}1935}
1936 1936 
1937TEST_F(OptimizerSt, ScalarBroadcastOptimization_Two_Scalar) {1937TEST_F(OptimizerSt, ScalarBroadcastOptimization_Two_Scalar) {
@@ -1967,10 +1967,10 @@ TEST_F(OptimizerSt, ScalarBroadcastOptimization_Two_Scalar) {
1967 EXPECT_EQ(impl_graphs.size(), 3);1967 EXPECT_EQ(impl_graphs.size(), 3);
1968 auto impl_graph0 = af::AscGraphUtils::GetComputeGraph(impl_graphs[0]);1968 auto impl_graph0 = af::AscGraphUtils::GetComputeGraph(impl_graphs[0]);
1969 EXPECT_EQ(impl_graph0->GetAllNodesSize(), 8);1969 EXPECT_EQ(impl_graph0->GetAllNodesSize(), 8);
1970- EXPECT_NE(impl_graph0->FindNode("brc1"), nullptr);1970+ EXPECT_EQ(impl_graph0->FindNode("brc1"), nullptr);
1971 EXPECT_EQ(impl_graph0->FindNode("brc2"), nullptr);1971 EXPECT_EQ(impl_graph0->FindNode("brc2"), nullptr);
1972 EXPECT_EQ(impl_graph0->FindNode("brc3"), nullptr);1972 EXPECT_EQ(impl_graph0->FindNode("brc3"), nullptr);
1973- EXPECT_EQ(impl_graph0->FindNode("brc4"), nullptr);1973+ EXPECT_NE(impl_graph0->FindNode("brc4"), nullptr);
1974 EXPECT_EQ(impl_graph0->FindNode("brc5"), nullptr);1974 EXPECT_EQ(impl_graph0->FindNode("brc5"), nullptr);
1975 EXPECT_EQ(impl_graph0->FindNode("brc6"), nullptr);1975 EXPECT_EQ(impl_graph0->FindNode("brc6"), nullptr);
1976}1976}
@@ -2001,23 +2001,23 @@ TEST_F(OptimizerSt, RemoveRedundantBroadcast) {
2001 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0].size(), 1UL);2001 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0].size(), 1UL);
2002 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL);2002 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL);
2003 auto impl_graphs = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs;2003 auto impl_graphs = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs;
2004- EXPECT_EQ(impl_graphs.size(), 4);2004+ EXPECT_EQ(impl_graphs.size(), 3);
2005- // consumer split creates clone for exp1; common-axis backward removes brc0/brc1 from add0's chain2005+ // check don't remove brc
2006 auto impl_grp_0_exp1 = impl_graphs[0].FindNode("exp1");2006 auto impl_grp_0_exp1 = impl_graphs[0].FindNode("exp1");
2007 EXPECT_NE(impl_grp_0_exp1, nullptr);2007 EXPECT_NE(impl_grp_0_exp1, nullptr);
2008 EXPECT_EQ(impl_grp_0_exp1->GetAllInDataAnchorsSize(), 1);2008 EXPECT_EQ(impl_grp_0_exp1->GetAllInDataAnchorsSize(), 1);
2009- EXPECT_EQ(impl_grp_0_exp1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(),2009+ EXPECT_EQ(impl_grp_0_exp1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc0");
2010- "brc0_consumer_split_1");
2011 2010 
2012 auto impol_grp_0_add0 = impl_graphs[0].FindNode("add0");2011 auto impol_grp_0_add0 = impl_graphs[0].FindNode("add0");
2013 EXPECT_NE(impol_grp_0_add0, nullptr);2012 EXPECT_NE(impol_grp_0_add0, nullptr);
2014 EXPECT_EQ(impol_grp_0_add0->GetAllInDataAnchorsSize(), 2);2013 EXPECT_EQ(impol_grp_0_add0->GetAllInDataAnchorsSize(), 2);
2015- EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0");2014+ EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc0");
2016- EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "exp0");2015+ EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc1");
2017 2016 
2018- EXPECT_NE(impl_graphs[0].FindNode("brc0_consumer_split_1"), nullptr);2017+ EXPECT_NE(impl_graphs[0].FindNode("brc0"), nullptr);
2018+ EXPECT_NE(impl_graphs[0].FindNode("brc1"), nullptr);
2019 2019 
2020- // check remove brc in unaligned template2020+ // check remove brc
2021 auto impl_grp_1_exp1 = impl_graphs[1].FindNode("exp1");2021 auto impl_grp_1_exp1 = impl_graphs[1].FindNode("exp1");
2022 EXPECT_NE(impl_grp_1_exp1, nullptr);2022 EXPECT_NE(impl_grp_1_exp1, nullptr);
2023 EXPECT_EQ(impl_grp_1_exp1->GetAllInDataAnchorsSize(), 1);2023 EXPECT_EQ(impl_grp_1_exp1->GetAllInDataAnchorsSize(), 1);
@@ -2262,9 +2262,20 @@ TEST_F(OptimizerSt, BufQueAllocator_RemovePad_MemUnique) {
2262 broadcast1.y.dtype = af::DataType::DT_FLOAT;2262 broadcast1.y.dtype = af::DataType::DT_FLOAT;
2263 broadcast1.attr.api.unit = ComputeUnit::kUnitVector;2263 broadcast1.attr.api.unit = ComputeUnit::kUnitVector;
2264 2264 
2265+ af::ascir_op::Abs abs0("abs0");
2266+ abs0.x = broadcast1.y;
2267+ abs0.attr.api.compute_type = ComputeType::kComputeElewise;
2268+ abs0.attr.api.type = af::ApiType::kAPITypeCompute;
2269+ abs0.attr.sched.axis = {z0.id, z1.id};
2270+ *abs0.y.axis = {z0.id, z1.id};
2271+ *abs0.y.repeats = {s0, s1};
2272+ *abs0.y.strides = {s1, One};
2273+ abs0.y.dtype = af::DataType::DT_FLOAT;
2274+ abs0.attr.api.unit = ComputeUnit::kUnitVector;
2275+ 
2265 af::ascir_op::Add add0("add0");2276 af::ascir_op::Add add0("add0");
2266 add0.x1 = load0.y;2277 add0.x1 = load0.y;
2267- add0.x2 = broadcast1.y;2278+ add0.x2 = abs0.y;
2268 add0.attr.api.compute_type = ComputeType::kComputeElewise;2279 add0.attr.api.compute_type = ComputeType::kComputeElewise;
2269 add0.attr.api.type = af::ApiType::kAPITypeCompute;2280 add0.attr.api.type = af::ApiType::kAPITypeCompute;
2270 add0.attr.sched.axis = {z0.id, z1.id};2281 add0.attr.sched.axis = {z0.id, z1.id};
@@ -2328,20 +2339,23 @@ TEST_F(OptimizerSt, BufQueAllocator_RemovePad_MemUnique) {
2328 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL);2339 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL);
2329 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs.size(), 3UL);2340 EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs.size(), 3UL);
2330 2341 
2331- auto impl_graph1 = af::AscGraphUtils::GetComputeGraph(2342+ auto impl_graph2 = af::AscGraphUtils::GetComputeGraph(
2332- fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[1]);2343+ fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[2]);
2333- EXPECT_EQ(impl_graph1->GetAllNodesSize(), 12);2344+ EXPECT_EQ(impl_graph2->GetAllNodesSize(), 13);
2334- EXPECT_NE(impl_graph1->FindNode("broadcast1"), nullptr);2345+ EXPECT_NE(impl_graph2->FindNode("broadcast1"), nullptr);
2335- EXPECT_NE(impl_graph1->FindNode("broadcast1_remove_pad_0"), nullptr);2346+ EXPECT_NE(impl_graph2->FindNode("broadcast1_remove_pad_0"), nullptr);
2336- EXPECT_NE(impl_graph1->FindNode("add0"), nullptr);2347+ EXPECT_NE(impl_graph2->FindNode("add0"), nullptr);
2337- const auto &impl_graph1_brc1 = std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("broadcast1"));2348+ EXPECT_NE(impl_graph2->FindNode("abs0"), nullptr);
2338- const auto &impl_graph1_rpd =2349+ const auto &impl_graph2_brc1 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("broadcast1"));
2339- std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("broadcast1_remove_pad_0"));2350+ const auto &impl_graph2_rpd =
2340- const auto &impl_graph1_add0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("add0"));2351+ std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("broadcast1_remove_pad_0"));
2341- const auto &impl_graph1_mul0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("mul0"));2352+ const auto &impl_graph2_add0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("add0"));
2342- EXPECT_EQ(impl_graph1_brc1->outputs[0].attr.buf.id, 1);2353+ const auto &impl_graph2_abs0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("abs0"));
2343- EXPECT_EQ(impl_graph1_rpd->outputs[0].attr.buf.id, 2);2354+ const auto &impl_graph2_mul0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("mul0"));
2344- EXPECT_EQ(impl_graph1_add0->outputs[0].attr.que.id, impl_graph1_mul0->outputs[0].attr.que.id);2355+ EXPECT_EQ(impl_graph2_brc1->outputs[0].attr.buf.id, 1);
2356+ EXPECT_EQ(impl_graph2_rpd->outputs[0].attr.buf.id, 2);
2357+ EXPECT_EQ(impl_graph2_abs0->outputs[0].attr.buf.id, 3);
2358+ EXPECT_EQ(impl_graph2_add0->outputs[0].attr.que.id, impl_graph2_mul0->outputs[0].attr.que.id);
2345}2359}
2346 2360 
2347TEST_F(OptimizerSt, BufQueAllocator_Inplace) {2361TEST_F(OptimizerSt, BufQueAllocator_Inplace) {
@@ -1,132 +0,0 @@
1-/**
2- * Copyright (c) 2026 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 "gtest/gtest.h"
12- 
13-#include "asc_graph_builder.h"
14-#include "tests/framework/broadcast_backward/broadcast_backward_test_utils.h"
15-#include "tests/framework/broadcast_backward/broadcast_backward_ut_utils.h"
16-#include "optimize/graph_pass/broadcast_backward_pass.h"
17- 
18-namespace {
19-using af::AscGraph;
20-using af::testing::AscGraphBuilder;
21-using af::testing::Sym;
22-using namespace broadcast_backward_test;
23-} // namespace
24- 
25-TEST(BroadcastBackwardPassSt, ChecksScalarForkJoinBranches) {
26- auto graph = BuildScalarForkJoinGraph("broadcast_backward_scalar_fork_join_st");
27- CompleteApiInfo(graph);
28- optimize::BroadcastBackwardPass pass;
29- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
30- EXPECT_TRUE(HasNode(graph, "store"));
31-}
32- 
33-TEST(BroadcastBackwardPassSt, MovesSingleInputChain) {
34- auto graph = BuildUnaryGraph("Abs");
35- CompleteApiInfo(graph);
36- optimize::BroadcastBackwardPass pass;
37- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
38- EXPECT_EQ(GetInputNodeName(FindNode(graph, "compute")), "load");
39- EXPECT_EQ(GetInputNodeName(FindNode(graph, "broadcast")), "compute");
40- EXPECT_EQ(GetInputNodeName(FindNode(graph, "store")), "broadcast");
41-}
42- 
43-TEST(BroadcastBackwardPassSt, MovesIdenticalMultiInputChains) {
44- // 该场景受全量 ST 的全局平台状态影响,单独运行通过但不适合作为全量 ST 用例。
45- GTEST_SKIP();
46- auto graph = BuildBinaryGraph("Add");
47- CompleteApiInfo(graph);
48- optimize::BroadcastBackwardPass pass;
49- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
50- ExpectBinaryBroadcastMove(graph);
51-}
52- 
53-TEST(BroadcastBackwardPassSt, MovesDirectFanOutBranches) {
54- auto graph = BuildDirectFanOutGraph("broadcast_backward_direct_fan_out_st");
55- CompleteApiInfo(graph);
56- ExpectDirectFanOutCandidate(graph);
57- optimize::BroadcastBackwardPass pass;
58- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
59- ExpectDirectFanOutMoved(graph);
60-}
61- 
62-TEST(BroadcastBackwardPassSt, MovesThreeDimensionalCommonAxis) {
63- // 当前本仓 BRC 不会对该 common-axis 图触发改写。
64- GTEST_SKIP();
65- auto graph = BuildCommonAxisGraph("broadcast_backward_common_axis_st");
66- CompleteApiInfo(graph);
67- optimize::BroadcastBackwardPass pass;
68- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
69- EXPECT_EQ(GetInputNodeName(FindNode(graph, "merge"), 0U), "broadcast0_residual_0");
70- EXPECT_EQ(GetInputNodeName(FindNode(graph, "merge"), 1U), "broadcast1_residual_1");
71- EXPECT_EQ(GetInputNodeName(FindNode(graph, "store")), "merge_broadcast_backward_common");
72-}
73- 
74-TEST(BroadcastBackwardPassSt, MovesSameConsumerMultiReference) {
75- const auto s0 = Sym("s0");
76- const auto s1 = Sym("s1");
77- auto graph = AscGraphBuilder("broadcast_backward_multi_reference_st")
78- .Loops({s0, s1})
79- .Data("data", 0)
80- .Load("load", "data", {s0, af::sym::kSymbolOne}, {af::sym::kSymbolOne, af::sym::kSymbolZero})
81- .Broadcast("broadcast", "load", {s0, s1})
82- .Add("merge", "broadcast", "broadcast")
83- .Store("store", "merge")
84- .Output("output", "store")
85- .Build();
86- CompleteApiInfo(graph);
87- optimize::BroadcastBackwardPass pass;
88- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
89- EXPECT_EQ(GetInputNodeName(FindNode(graph, "merge"), 0U), "load");
90- EXPECT_EQ(GetInputNodeName(FindNode(graph, "merge"), 1U), "load");
91- EXPECT_EQ(GetInputNodeName(FindNode(graph, "broadcast")), "merge");
92- EXPECT_EQ(GetInputNodeName(FindNode(graph, "store")), "broadcast");
93-}
94- 
95-TEST(BroadcastBackwardPassSt, HandlesPartialCommonBroadcastAxis) {
96- const auto s0 = Sym("s0");
97- const auto s1 = Sym("s1");
98- const auto s2 = Sym("s2");
99- auto graph = AscGraphBuilder("broadcast_backward_partial_common_axis_st")
100- .Loops({s0, s1, s2})
101- .Data("data0", 0)
102- .Data("data1", 1)
103- .Load("load0", "data0", {af::sym::kSymbolOne, s1, s2},
104- {af::sym::kSymbolZero, af::sym::kSymbolOne, af::sym::kSymbolOne})
105- .Load("load1", "data1", {af::sym::kSymbolOne, af::sym::kSymbolOne, s2},
106- {af::sym::kSymbolZero, af::sym::kSymbolZero, af::sym::kSymbolOne})
107- .Broadcast("broadcast0", "load0", {0})
108- .Broadcast("broadcast1", "load1", {0, 1})
109- .Add("merge", "broadcast0", "broadcast1")
110- .Store("store", "merge")
111- .Output("output", "store")
112- .Build();
113- CompleteApiInfo(graph);
114- optimize::BroadcastBackwardPass pass;
115- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
116- EXPECT_TRUE(HasNode(graph, "store"));
117-}
118- 
119-TEST(BroadcastBackwardPassSt, DtypeAwareBackwardEnablesCommonAxis) {
120- // 当前本仓 BRC 不会对该 dtype-aware common-axis 图触发改写。
121- GTEST_SKIP();
122- auto graph = BuildDtypeAwareCommonAxisGraph("broadcast_backward_dtype_aware_common_axis_st");
123- CompleteApiInfo(graph);
124- SetNodeDtype(graph, "relu", af::DT_FLOAT16);
125- SetNodeDtype(graph, "merge", af::DT_FLOAT16);
126- SetNodeDtype(graph, "store", af::DT_FLOAT16);
127- optimize::BroadcastBackwardPass pass;
128- ASSERT_EQ(pass.RunPass(graph), af::SUCCESS);
129- EXPECT_EQ(GetInputNodeName(FindNode(graph, "broadcast0_residual_0")), "cast0");
130- EXPECT_EQ(GetInputNodeName(FindNode(graph, "broadcast1_residual_1")), "relu");
131- EXPECT_EQ(GetInputNodeName(FindNode(graph, "store")), "merge_broadcast_backward_common");
132-}
@@ -3395,19 +3395,12 @@ TEST_F(TestOptimizer, ScalarBroadcastOptimization_Two_Scalar) {
3395 EXPECT_EQ(res, af::SUCCESS);3395 EXPECT_EQ(res, af::SUCCESS);
3396 auto compute_graph = af::AscGraphUtils::GetComputeGraph(graph);3396 auto compute_graph = af::AscGraphUtils::GetComputeGraph(graph);
3397 EXPECT_EQ(compute_graph->GetAllNodesSize(), 10);3397 EXPECT_EQ(compute_graph->GetAllNodesSize(), 10);
3398- const auto retained_brc1 = compute_graph->FindNode("brc1");3398+ EXPECT_EQ(compute_graph->FindNode("brc1"), nullptr);
3399- const auto retained_brc2 = compute_graph->FindNode("brc2");3399+ EXPECT_EQ(compute_graph->FindNode("brc2"), nullptr);
3400- const auto retained_brc3 = compute_graph->FindNode("brc3");3400+ EXPECT_EQ(compute_graph->FindNode("brc3"), nullptr);
3401- ASSERT_NE(retained_brc1, nullptr);3401+ EXPECT_NE(compute_graph->FindNode("brc4"), nullptr);
3402- ASSERT_NE(retained_brc2, nullptr);3402+ EXPECT_NE(compute_graph->FindNode("brc5"), nullptr);
3403- ASSERT_NE(retained_brc3, nullptr);3403+ EXPECT_NE(compute_graph->FindNode("brc6"), nullptr);
3404- EXPECT_EQ(compute_graph->FindNode("brc4"), nullptr);
3405- EXPECT_EQ(compute_graph->FindNode("brc5"), nullptr);
3406- EXPECT_EQ(compute_graph->FindNode("brc6"), nullptr);
3407- EXPECT_EQ(retained_brc1->GetInDataNodes().at(0)->GetName(), "add");
3408- EXPECT_EQ(retained_brc2->GetInDataNodes().at(0)->GetName(), "brc1");
3409- EXPECT_EQ(retained_brc3->GetInDataNodes().at(0)->GetName(), "brc2");
3410- EXPECT_EQ(compute_graph->FindNode("store")->GetInDataNodes().at(0)->GetName(), "brc3");
3411}3404}
3412 3405 
3413TEST_F(TestOptimizer, ScalarBroadcastOptimization_Same_Input) {3406TEST_F(TestOptimizer, ScalarBroadcastOptimization_Same_Input) {
@@ -406,32 +406,25 @@ TEST_F(OptimizerStV2, NddmaCaseBrcOutputWithMultiRef) {
406 .Data("data0", 0, af::DT_FLOAT)406 .Data("data0", 0, af::DT_FLOAT)
407 .Load("load0", "data0", load_shape, load_strides)407 .Load("load0", "data0", load_shape, load_strides)
408 .Broadcast("broadcast", "load0", {0, 1}) // broadcast on both axes408 .Broadcast("broadcast", "load0", {0, 1}) // broadcast on both axes
409- .Scalar("scalar0", "0", af::DT_FLOAT)409+ .Exp("exp0", "broadcast")
410- .Add("exp0", "broadcast", "scalar0")410+ .Abs("abs0", "broadcast")
411- .Abs("abs0", "exp0")411+ .Mul("mul0", "exp0", "abs0")
412- .Store("store", "abs0")412+ .Store("store", "mul0")
413 .Output("output", "store", 8, af::DT_FLOAT)413 .Output("output", "store", 8, af::DT_FLOAT)
414 .Build();414 .Build();
415 415 
416 ::ascir::FusedScheduledResult fused_scheduled_result;416 ::ascir::FusedScheduledResult fused_scheduled_result;
417- ASSERT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);417+ EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);
418- ASSERT_FALSE(fused_scheduled_result.node_idx_to_scheduled_results.empty());418+ 
419- ASSERT_FALSE(fused_scheduled_result.node_idx_to_scheduled_results[0].empty());419+ for (const auto &node :
420- ASSERT_FALSE(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.empty());420+ fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[1].GetAllNodes()) {
421- const auto &impl_graphs = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs;421+ if (node->GetOpDesc()->GetId() == 1) {
422- ASSERT_GT(impl_graphs.size(), 1UL);422+ EXPECT_EQ(node->GetOpDesc()->GetType(), "Nddma");
423- bool has_nddma = false;
424- bool has_vector_func = false;
425- for (const auto &node : impl_graphs[1].GetAllNodes()) {
426- if (node->GetOpDesc()->GetType() == "Nddma") {
427- has_nddma = true;
428 }423 }
429- if (node->GetOpDesc()->GetType() == "VectorFunc") {424+ if (node->GetOpDesc()->GetId() == 2) {
430- has_vector_func = true;425+ EXPECT_EQ(node->GetOpDesc()->GetType(), "VectorFunc");
431 }426 }
432 }427 }
433- EXPECT_TRUE(has_nddma);
434- EXPECT_TRUE(has_vector_func);
435}428}
436 429 
437TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) {430TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) {
@@ -463,24 +456,16 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) {
463 *broadcast0.y.repeats = {s0, s1};456 *broadcast0.y.repeats = {s0, s1};
464 *broadcast0.y.strides = {s1, af::ops::One};457 *broadcast0.y.strides = {s1, af::ops::One};
465 458 
466- Scalar scalar0("scalar0", graph);459+ Exp exp0("exp0");
467- scalar0.y.dtype = dtype;
468- scalar0.attr.sched.axis = {z0.id, z1.id};
469- *scalar0.y.axis = {z0.id, z1.id};
470- *scalar0.y.repeats = {s0, s1};
471- *scalar0.y.strides = {s1, af::ops::One};
472- 
473- Add exp0("exp0");
474 exp0.attr.sched.axis = {z0.id, z1.id};460 exp0.attr.sched.axis = {z0.id, z1.id};
475- exp0.x1 = broadcast0.y;461+ exp0.x = broadcast0.y;
476- exp0.x2 = scalar0.y;
477 *exp0.y.axis = {z0.id, z1.id};462 *exp0.y.axis = {z0.id, z1.id};
478 exp0.y.dtype = dtype;463 exp0.y.dtype = dtype;
479 *exp0.y.repeats = {s0, s1};464 *exp0.y.repeats = {s0, s1};
480 *exp0.y.strides = {s1, af::ops::One};465 *exp0.y.strides = {s1, af::ops::One};
481 466 
482 Abs abs0("abs0");467 Abs abs0("abs0");
483- abs0.x = exp0.y;468+ abs0.x = broadcast0.y;
484 abs0.attr.sched.axis = {z0.id, z1.id};469 abs0.attr.sched.axis = {z0.id, z1.id};
485 abs0.y.dtype = dtype;470 abs0.y.dtype = dtype;
486 *abs0.y.axis = {z0.id, z1.id};471 *abs0.y.axis = {z0.id, z1.id};
@@ -488,9 +473,18 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) {
488 *abs0.y.strides = {s1, One};473 *abs0.y.strides = {s1, One};
489 abs0.attr.api.compute_type = ComputeType::kComputeElewise;474 abs0.attr.api.compute_type = ComputeType::kComputeElewise;
490 475 
476+ Mul mul0("mul0");
477+ mul0.attr.sched.axis = {z0.id, z1.id};
478+ mul0.x1 = exp0.y;
479+ mul0.x2 = abs0.y;
480+ mul0.y.dtype = dtype;
481+ *mul0.y.axis = {z0.id, z1.id};
482+ *mul0.y.repeats = {s0, s1};
483+ *mul0.y.strides = {s1, One};
484+ 
491 Store store_op("store");485 Store store_op("store");
492 store_op.attr.sched.axis = {z0.id, z1.id};486 store_op.attr.sched.axis = {z0.id, z1.id};
493- store_op.x = abs0.y;487+ store_op.x = mul0.y;
494 *store_op.y.axis = {z0.id, z1.id};488 *store_op.y.axis = {z0.id, z1.id};
495 store_op.y.dtype = dtype;489 store_op.y.dtype = dtype;
496 *store_op.y.strides = {s1, af::ops::One};490 *store_op.y.strides = {s1, af::ops::One};
@@ -505,8 +499,15 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) {
505 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);499 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);
506 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];500 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];
507 501 
508- ASSERT_FALSE(schedule_group.impl_graphs.empty());502+ ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2);
509- EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL);503+ 
504+ const auto score_func_iter = schedule_group.graph_name_to_score_funcs.find(schedule_group.impl_graphs[2].GetName());
505+ ASSERT_NE(score_func_iter, schedule_group.graph_name_to_score_funcs.end());
506+ const auto res =
507+ "int32_t CalcScore(const AutofuseTilingData &tiling_data) {\n"
508+ " return -1;\n"
509+ "}\n";
510+ EXPECT_EQ(score_func_iter->second, res);
510}511}
511 512 
512TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) {513TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) {
@@ -538,24 +539,16 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) {
538 *broadcast0.y.repeats = {s0, s1};539 *broadcast0.y.repeats = {s0, s1};
539 *broadcast0.y.strides = {s1, af::ops::One};540 *broadcast0.y.strides = {s1, af::ops::One};
540 541 
541- Scalar scalar0("scalar0", graph);542+ Exp exp0("exp0");
542- scalar0.y.dtype = dtype;
543- scalar0.attr.sched.axis = {z0.id, z1.id};
544- *scalar0.y.axis = {z0.id, z1.id};
545- *scalar0.y.repeats = {s0, s1};
546- *scalar0.y.strides = {s1, af::ops::One};
547- 
548- Add exp0("exp0");
549 exp0.attr.sched.axis = {z0.id, z1.id};543 exp0.attr.sched.axis = {z0.id, z1.id};
550- exp0.x1 = broadcast0.y;544+ exp0.x = broadcast0.y;
551- exp0.x2 = scalar0.y;
552 *exp0.y.axis = {z0.id, z1.id};545 *exp0.y.axis = {z0.id, z1.id};
553 exp0.y.dtype = dtype;546 exp0.y.dtype = dtype;
554 *exp0.y.repeats = {s0, s1};547 *exp0.y.repeats = {s0, s1};
555 *exp0.y.strides = {s1, af::ops::One};548 *exp0.y.strides = {s1, af::ops::One};
556 549 
557 Abs abs0("abs0");550 Abs abs0("abs0");
558- abs0.x = exp0.y;551+ abs0.x = broadcast0.y;
559 abs0.attr.sched.axis = {z0.id, z1.id};552 abs0.attr.sched.axis = {z0.id, z1.id};
560 abs0.y.dtype = dtype;553 abs0.y.dtype = dtype;
561 *abs0.y.axis = {z0.id, z1.id};554 *abs0.y.axis = {z0.id, z1.id};
@@ -563,9 +556,18 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) {
563 *abs0.y.strides = {s1, One};556 *abs0.y.strides = {s1, One};
564 abs0.attr.api.compute_type = ComputeType::kComputeElewise;557 abs0.attr.api.compute_type = ComputeType::kComputeElewise;
565 558 
559+ Mul mul0("mul0");
560+ mul0.attr.sched.axis = {z0.id, z1.id};
561+ mul0.x1 = exp0.y;
562+ mul0.x2 = abs0.y;
563+ mul0.y.dtype = dtype;
564+ *mul0.y.axis = {z0.id, z1.id};
565+ *mul0.y.repeats = {s0, s1};
566+ *mul0.y.strides = {s1, One};
567+ 
566 Store store_op("store");568 Store store_op("store");
567 store_op.attr.sched.axis = {z0.id, z1.id};569 store_op.attr.sched.axis = {z0.id, z1.id};
568- store_op.x = abs0.y;570+ store_op.x = mul0.y;
569 *store_op.y.axis = {z0.id, z1.id};571 *store_op.y.axis = {z0.id, z1.id};
570 store_op.y.dtype = dtype;572 store_op.y.dtype = dtype;
571 *store_op.y.strides = {s1, af::ops::One};573 *store_op.y.strides = {s1, af::ops::One};
@@ -580,8 +582,17 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) {
580 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);582 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);
581 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];583 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];
582 584 
583- ASSERT_FALSE(schedule_group.impl_graphs.empty());585+ ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2);
584- EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL);586+ const auto score_func_iter = schedule_group.graph_name_to_score_funcs.find(schedule_group.impl_graphs[2].GetName());
587+ ASSERT_NE(score_func_iter, schedule_group.graph_name_to_score_funcs.end());
588+ const auto res =
589+ "int32_t CalcScore(const AutofuseTilingData &tiling_data) {\n"
590+ " const auto tail_size = static_cast<int64_t>((2 * tiling_data.s1));\n"
591+ " if (tail_size % 32 == 0) { return -1; }\n"
592+ " if (tail_size > 4096) { return -1; }\n"
593+ " return 0;\n"
594+ "}\n";
595+ EXPECT_EQ(score_func_iter->second, res);
585}596}
586 597 
587TEST_F(OptimizerStV2, NddmaCaseLargeTailBrcScoreFunc) {598TEST_F(OptimizerStV2, NddmaCaseLargeTailBrcScoreFunc) {
@@ -597,18 +608,25 @@ TEST_F(OptimizerStV2, NddmaCaseLargeTailBrcScoreFunc) {
597 .Data("data0", 0, af::DT_FLOAT)608 .Data("data0", 0, af::DT_FLOAT)
598 .Load("load0", "data0", load0_shape, load0_strides)609 .Load("load0", "data0", load0_shape, load0_strides)
599 .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1610 .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1
600- .Scalar("scalar0", "0", af::DT_FLOAT)611+ .Exp("exp0", "broadcast")
601- .Add("exp0", "broadcast", "scalar0")612+ .Abs("abs0", "broadcast")
602- .Abs("abs0", "exp0")613+ .Mul("mul0", "exp0", "abs0")
603- .Store("store", "abs0")614+ .Store("store", "mul0")
604 .Output("output", "store", 8, af::DT_FLOAT)615 .Output("output", "store", 8, af::DT_FLOAT)
605 .Build();616 .Build();
606 617 
607 ::ascir::FusedScheduledResult fused_scheduled_result;618 ::ascir::FusedScheduledResult fused_scheduled_result;
608 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);619 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);
609 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];620 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];
610- ASSERT_FALSE(schedule_group.impl_graphs.empty());621+ ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2);
611- EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL);622+ 
623+ const auto score_func_iter = schedule_group.graph_name_to_score_funcs.find(schedule_group.impl_graphs[2].GetName());
624+ ASSERT_NE(score_func_iter, schedule_group.graph_name_to_score_funcs.end());
625+ const auto res =
626+ "int32_t CalcScore(const AutofuseTilingData &tiling_data) {\n"
627+ " return -1;\n"
628+ "}\n";
629+ EXPECT_EQ(score_func_iter->second, res);
612}630}
613 631 
614TEST_F(OptimizerStV2, NddmaCaseLargeTailBrc_Dynamic) {632TEST_F(OptimizerStV2, NddmaCaseLargeTailBrc_Dynamic) {
@@ -624,18 +642,17 @@ TEST_F(OptimizerStV2, NddmaCaseLargeTailBrc_Dynamic) {
624 .Data("data0", 0, af::DT_FLOAT)642 .Data("data0", 0, af::DT_FLOAT)
625 .Load("load0", "data0", load0_shape, load0_strides)643 .Load("load0", "data0", load0_shape, load0_strides)
626 .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1644 .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1
627- .Scalar("scalar0", "0", af::DT_FLOAT)645+ .Exp("exp0", "broadcast")
628- .Add("exp0", "broadcast", "scalar0")646+ .Abs("abs0", "broadcast")
629- .Abs("abs0", "exp0")647+ .Mul("mul0", "exp0", "abs0")
630- .Store("store", "abs0")648+ .Store("store", "mul0")
631 .Output("output", "store", 8, af::DT_FLOAT)649 .Output("output", "store", 8, af::DT_FLOAT)
632 .Build();650 .Build();
633 651 
634 ::ascir::FusedScheduledResult fused_scheduled_result;652 ::ascir::FusedScheduledResult fused_scheduled_result;
635 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);653 EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0);
636 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];654 const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0];
637- ASSERT_FALSE(schedule_group.impl_graphs.empty());655+ ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2);
638- EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL);
639}656}
640 657 
641/**658/**
@@ -64,19 +64,14 @@ TEST_F(SameSourceBroadcastCseStTest, MergesEquivalentBroadcastsThroughGraphPassR
64 64 
65 ASSERT_EQ(optimizer.GraphPass(graph), af::SUCCESS);65 ASSERT_EQ(optimizer.GraphPass(graph), af::SUCCESS);
66 66 
67+ const auto canonical = graph.FindNode("broadcast0");
67 const auto add = graph.FindNode("add");68 const auto add = graph.FindNode("add");
69+ ASSERT_NE(canonical, nullptr);
68 ASSERT_NE(add, nullptr);70 ASSERT_NE(add, nullptr);
69 EXPECT_EQ(graph.FindNode("broadcast1"), nullptr);71 EXPECT_EQ(graph.FindNode("broadcast1"), nullptr);
70- 72+ EXPECT_EQ(add->GetInDataAnchor(0)->GetPeerOutAnchor(), canonical->GetOutDataAnchor(0));
71- const auto input0_peer = add->GetInDataAnchor(0)->GetPeerOutAnchor();73+ EXPECT_EQ(add->GetInDataAnchor(1)->GetPeerOutAnchor(), canonical->GetOutDataAnchor(0));
72- const auto input1_peer = add->GetInDataAnchor(1)->GetPeerOutAnchor();74+ EXPECT_EQ(reduce->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 1UL);
73- ASSERT_NE(input0_peer, nullptr);
74- ASSERT_NE(input1_peer, nullptr);
75- EXPECT_EQ(input0_peer, input1_peer);
76- ASSERT_NE(input0_peer->GetOwnerNode(), nullptr);
77- EXPECT_EQ(input0_peer->GetOwnerNode()->GetName(), "reduce");
78- EXPECT_EQ(input1_peer->GetOwnerNode()->GetName(), "reduce");
79- EXPECT_EQ(reduce->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 2UL);
80}75}
81 76 
82TEST_F(SameSourceBroadcastCseStTest, SkipsGraphWithoutNormStructureThroughGraphPassRunner) {77TEST_F(SameSourceBroadcastCseStTest, SkipsGraphWithoutNormStructureThroughGraphPassRunner) {
@@ -1497,7 +1497,7 @@ TEST_F(VectorFuncSt, CastNotFusion) {
1497 std::vector<af::AscGraph> asc_graphs;1497 std::vector<af::AscGraph> asc_graphs;
1498 fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[0].GetAllSubGraphs(1498 fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[0].GetAllSubGraphs(
1499 asc_graphs);1499 asc_graphs);
1500- EXPECT_EQ(asc_graphs.size(), 3UL);1500+ EXPECT_EQ(asc_graphs.size(), 2UL);
1501 auto graph1 = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[1];1501 auto graph1 = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[1];
1502 std::vector<af::AscGraph> asc_graphs1;1502 std::vector<af::AscGraph> asc_graphs1;
1503 graph1.GetAllSubGraphs(asc_graphs1);1503 graph1.GetAllSubGraphs(asc_graphs1);
@@ -12,7 +12,6 @@
12#define OPTIMIZE_PLATFORM_V2_PASS_RUNNER_V2_H12#define OPTIMIZE_PLATFORM_V2_PASS_RUNNER_V2_H
13 13 
14#include "optimize/platform/common/pass_runner.h"14#include "optimize/platform/common/pass_runner.h"
15-#include "optimize/graph_pass/broadcast_backward_pass.h"
16#include "optimize/graph_pass/broadcast_const_to_store.h"15#include "optimize/graph_pass/broadcast_const_to_store.h"
17#include "optimize/graph_pass/scalar_to_1d_tensor.h"16#include "optimize/graph_pass/scalar_to_1d_tensor.h"
18#include "optimize/graph_pass/scalar_broadcast_optimization.h"17#include "optimize/graph_pass/scalar_broadcast_optimization.h"
@@ -32,12 +31,11 @@ class PassRunnerV2 final : public BasePassRunner {
32 this->RegisterPass<PowEquivSubstitutionPass>();31 this->RegisterPass<PowEquivSubstitutionPass>();
33 this->RegisterPass<BroadcastConstToStorePass>();32 this->RegisterPass<BroadcastConstToStorePass>();
34 this->RegisterPass<ScalarTo1DTensorPass>();33 this->RegisterPass<ScalarTo1DTensorPass>();
35- this->RegisterPass<SameSourceBroadcastCsePass>();
36- this->RegisterPass<BroadcastBackwardPass>();
37 this->RegisterPass<ScalarBroadcastOptimizationPass>();34 this->RegisterPass<ScalarBroadcastOptimizationPass>();
38 this->RegisterPass<MaskedFillInputReorderPass>();35 this->RegisterPass<MaskedFillInputReorderPass>();
39 this->RegisterPass<ExpandDimsForAllReducePass>();36 this->RegisterPass<ExpandDimsForAllReducePass>();
40 this->RegisterPass<ContinuesBroadcastOptimizationPass>();37 this->RegisterPass<ContinuesBroadcastOptimizationPass>();
38+ this->RegisterPass<SameSourceBroadcastCsePass>();
41 this->RegisterPass<DuplicateElewiseCsePass>();39 this->RegisterPass<DuplicateElewiseCsePass>();
42 this->RegisterPass<GatherToLoadPass>();40 this->RegisterPass<GatherToLoadPass>();
43 this->RegisterPass<SplitConcatOptimizationPass>();41 this->RegisterPass<SplitConcatOptimizationPass>();