已合并
【PR】fix :Scalar 多下游不同 shape 时 Broadcast 属性复用导致图校验失败 #1165
【PR】fix :Scalar 多下游不同 shape 时 Broadcast 属性复用导致图校验失败 #1165
已合并
czways创建于 7月2日
共 2 个文件变更+158-11
@@ -16,6 +16,7 @@
16#include "ascir_ops.h"16#include "ascir_ops.h"
17#include "graph/utils/graph_utils.h"17#include "graph/utils/graph_utils.h"
18#include "graph/ascendc_ir/utils/asc_graph_utils.h"18#include "graph/ascendc_ir/utils/asc_graph_utils.h"
19+#include "graph/symbolizer/symbolic_utils.h"
19#include "ascgen_log.h"20#include "ascgen_log.h"
20 21 
21namespace af::pre_process {22namespace af::pre_process {
@@ -26,9 +27,36 @@ bool IsSkipDownstreamType(const af::NodePtr &node) {
26 type == af::ascir_op::Output::Type;27 type == af::ascir_op::Output::Type;
27}28}
28 29 
30+bool IsSameExpressions(const std::vector<af::Expression> &lhs, const std::vector<af::Expression> &rhs) {
31+ if (lhs.size() != rhs.size()) {
32+ return false;
33+ }
34+ for (size_t i = 0; i < lhs.size(); ++i) {
35+ if (af::SymbolicUtils::StaticCheckEq(lhs[i], rhs[i]) != af::TriBool::kTrue) {
36+ return false;
37+ }
38+ }
39+ return true;
40+}
41+ 
42+bool IsSameBroadcastTarget(const af::AscNodePtr &lhs, const af::AscNodePtr &rhs) {
43+ const auto &lhs_attr = lhs->outputs[0].attr;
44+ const auto &rhs_attr = rhs->outputs[0].attr;
45+ return lhs->attr.sched.axis == rhs->attr.sched.axis && lhs_attr.axis == rhs_attr.axis &&
46+ IsSameExpressions(lhs_attr.repeats, rhs_attr.repeats) && IsSameExpressions(lhs_attr.strides, rhs_attr.strides);
47+}
48+ 
49+struct BroadcastGroup {
50+ af::AscNodePtr ref_node;
51+ std::vector<af::InDataAnchorPtr> anchors;
52+};
53+ 
29af::AscNodePtr BuildBroadcastNode(af::AscGraph &asc_graph, const af::AscNodePtr &scalar_node,54af::AscNodePtr BuildBroadcastNode(af::AscGraph &asc_graph, const af::AscNodePtr &scalar_node,
30- const af::AscNodePtr &ref_node) {55+ const af::AscNodePtr &ref_node, size_t group_idx) {
31 std::string brc_name = scalar_node->GetName() + "_broadcast";56 std::string brc_name = scalar_node->GetName() + "_broadcast";
57+ if (group_idx != 0U) {
58+ brc_name += "_" + std::to_string(group_idx);
59+ }
32 af::ascir_op::Broadcast brc(brc_name.c_str());60 af::ascir_op::Broadcast brc(brc_name.c_str());
33 auto b_node = asc_graph.AddNode(brc);61 auto b_node = asc_graph.AddNode(brc);
34 b_node->attr.sched = ref_node->attr.sched;62 b_node->attr.sched = ref_node->attr.sched;
@@ -37,6 +65,19 @@ af::AscNodePtr BuildBroadcastNode(af::AscGraph &asc_graph, const af::AscNodePtr
37 return b_node;65 return b_node;
38}66}
39 67 
68+Status AddAnchorToBroadcastGroup(std::vector<BroadcastGroup> &groups, const af::InDataAnchorPtr &anchor) {
69+ auto ref_node = std::dynamic_pointer_cast<af::AscNode>(anchor->GetOwnerNode());
70+ GE_ASSERT_NOTNULL(ref_node);
71+ for (auto &group : groups) {
72+ if (IsSameBroadcastTarget(group.ref_node, ref_node)) {
73+ group.anchors.push_back(anchor);
74+ return af::SUCCESS;
75+ }
76+ }
77+ groups.push_back({ref_node, {anchor}});
78+ return af::SUCCESS;
79+}
80+ 
40Status InsertBroadcastAfterScalar(af::AscGraph &asc_graph, const af::AscNodePtr &scalar_node, bool &inserted) {81Status InsertBroadcastAfterScalar(af::AscGraph &asc_graph, const af::AscNodePtr &scalar_node, bool &inserted) {
41 auto out_anchor = scalar_node->GetOutDataAnchor(0);82 auto out_anchor = scalar_node->GetOutDataAnchor(0);
42 GE_ASSERT_NOTNULL(out_anchor);83 GE_ASSERT_NOTNULL(out_anchor);
@@ -57,18 +98,24 @@ Status InsertBroadcastAfterScalar(af::AscGraph &asc_graph, const af::AscNodePtr
57 return af::SUCCESS;98 return af::SUCCESS;
58 }99 }
59 100 
60- auto ref_node = std::dynamic_pointer_cast<af::AscNode>(compute_anchors[0]->GetOwnerNode());101+ std::vector<BroadcastGroup> groups;
61- GE_ASSERT_NOTNULL(ref_node);102+ for (const auto &anchor : compute_anchors) {
62- auto b_node = BuildBroadcastNode(asc_graph, scalar_node, ref_node);103+ GE_ASSERT_SUCCESS(AddAnchorToBroadcastGroup(groups, anchor));
63- auto b_out = b_node->GetOutDataAnchor(0);104+ }
64- for (const auto &dst : compute_anchors) {105+ 
65- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(out_anchor, dst));106+ for (size_t i = 0; i < groups.size(); ++i) {
66- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(b_out, dst));107+ auto b_node = BuildBroadcastNode(asc_graph, scalar_node, groups[i].ref_node, i);
108+ auto b_out = b_node->GetOutDataAnchor(0);
109+ for (const auto &dst : groups[i].anchors) {
110+ GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(out_anchor, dst));
111+ GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(b_out, dst));
112+ }
113+ GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(out_anchor, b_node->GetInDataAnchor(0)));
114+ 
115+ GELOGD("insert broadcast %s after scalar %s in graph %s.", b_node->GetName().c_str(),
116+ scalar_node->GetName().c_str(), asc_graph.GetName().c_str());
67 }117 }
68- GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(out_anchor, b_node->GetInDataAnchor(0)));
69 118 
70- GELOGD("insert broadcast %s after scalar %s in graph %s.", b_node->GetName().c_str(), scalar_node->GetName().c_str(),
71- asc_graph.GetName().c_str());
72 inserted = true;119 inserted = true;
73 return af::SUCCESS;120 return af::SUCCESS;
74}121}
@@ -212,6 +212,106 @@ TEST_F(TestScalarBroadcastInsert, ScalarMultipleDownstreams_InsertsOneSharedBroa
212 }212 }
213}213}
214 214 
215+TEST_F(TestScalarBroadcastInsert, ScalarMultipleDownstreamsWithDifferentView_InsertsSeparateBroadcasts) {
216+ auto graph = AscGraphBuilder("test_scalar_multi_downstream_different_view")
217+ .Loops({Sym("s0"), Sym("s1")})
218+ .Data("data0", 0)
219+ .Data("data1", 1)
220+ .Load("load0", "data0")
221+ .Load("load1", "data1")
222+ .Scalar("scalar0", "1.0")
223+ .Add("add0", "load0", "scalar0")
224+ .Mul("mul0", "load1", "scalar0")
225+ .Store("store0", "add0")
226+ .Store("store1", "mul0")
227+ .Output("output0", "store0")
228+ .Output("output1", "store1", 1)
229+ .Build();
230+ 
231+ optimize::AscGraphInfoComplete::CompleteApiInfo(graph);
232+ auto mul_node = graph.FindNode("mul0");
233+ ASSERT_NE(mul_node, nullptr);
234+ ASSERT_GE(mul_node->outputs[0].attr.repeats.size(), 1U);
235+ mul_node->outputs[0].attr.repeats[0] = af::sym::kSymbolOne;
236+ 
237+ auto ret = InsertBroadcastAfterScalarForAscGraph(graph);
238+ ASSERT_EQ(ret, af::SUCCESS);
239+ 
240+ EXPECT_EQ(CountNodesByType(graph, ascir_op::Broadcast::Type), 2U);
241+ EXPECT_FALSE(IsConnected(graph, "scalar0", "add0"));
242+ EXPECT_FALSE(IsConnected(graph, "scalar0", "mul0"));
243+ EXPECT_EQ(CountOutDataEdges(graph, "scalar0"), 2U);
244+}
245+ 
246+TEST_F(TestScalarBroadcastInsert, ScalarMultipleDownstreamsWithDifferentShapes_GroupsCompatibleViews) {
247+ auto graph = AscGraphBuilder("test_scalar_multi_downstream_different_shapes")
248+ .Loops({Sym("s0"), Sym("s1")})
249+ .Data("data0", 0)
250+ .Load("load0", "data0")
251+ .Scalar("scalar0", "2.0")
252+ .Add("add0", "load0", "scalar0")
253+ .Mul("mul0", "load0", "scalar0")
254+ .Add("add1", "load0", "scalar0")
255+ .Sub("sub0", "load0", "scalar0")
256+ .Build();
257+ 
258+ optimize::AscGraphInfoComplete::CompleteApiInfo(graph);
259+ auto mul_node = std::dynamic_pointer_cast<AscNode>(graph.FindNode("mul0"));
260+ auto add1_node = std::dynamic_pointer_cast<AscNode>(graph.FindNode("add1"));
261+ ASSERT_NE(mul_node, nullptr);
262+ ASSERT_NE(add1_node, nullptr);
263+ ASSERT_FALSE(mul_node->outputs[0].attr.repeats.empty());
264+ ASSERT_FALSE(add1_node->outputs[0].attr.repeats.empty());
265+ mul_node->outputs[0].attr.repeats[0] = Sym(14);
266+ mul_node->outputs[0].attr.strides = {Sym(322)};
267+ add1_node->outputs[0].attr.repeats[0] = Sym(32);
268+ add1_node->outputs[0].attr.strides = {Sym(736)};
269+ 
270+ ASSERT_EQ(InsertBroadcastAfterScalarForAscGraph(graph), af::SUCCESS);
271+ EXPECT_EQ(CountNodesByType(graph, ascir_op::Broadcast::Type), 3U);
272+ 
273+ const auto add_brc = graph.FindNode("add0")->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode();
274+ const auto mul_brc = graph.FindNode("mul0")->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode();
275+ const auto add1_brc = graph.FindNode("add1")->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode();
276+ const auto sub_brc = graph.FindNode("sub0")->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode();
277+ ASSERT_NE(add_brc, nullptr);
278+ ASSERT_NE(mul_brc, nullptr);
279+ ASSERT_NE(add1_brc, nullptr);
280+ ASSERT_NE(sub_brc, nullptr);
281+ EXPECT_EQ(add_brc, sub_brc);
282+ EXPECT_NE(add_brc, mul_brc);
283+ EXPECT_NE(add_brc, add1_brc);
284+ EXPECT_NE(mul_brc, add1_brc);
285+ EXPECT_EQ(CountOutDataEdges(graph, add_brc->GetName()), 2U);
286+}
287+ 
288+TEST_F(TestScalarBroadcastInsert, ScalarDownstreamWithoutTensorAttr_DoesNotSynthesizeBroadcastAttr) {
289+ auto graph = AscGraphBuilder("test_scalar_downstream_without_tensor_attr")
290+ .Loops({Sym("s0")})
291+ .Data("data0", 0)
292+ .Load("load0", "data0")
293+ .Scalar("scalar0", "1.0")
294+ .Add("add0", "load0", "scalar0")
295+ .Build();
296+ 
297+ optimize::AscGraphInfoComplete::CompleteApiInfo(graph);
298+ auto add_node = std::dynamic_pointer_cast<AscNode>(graph.FindNode("add0"));
299+ ASSERT_NE(add_node, nullptr);
300+ ASSERT_FALSE(add_node->attr.sched.axis.empty());
301+ add_node->outputs[0].attr.axis.clear();
302+ add_node->outputs[0].attr.repeats.clear();
303+ add_node->outputs[0].attr.strides.clear();
304+ 
305+ ASSERT_EQ(InsertBroadcastAfterScalarForAscGraph(graph), af::SUCCESS);
306+ 
307+ auto brc_node = std::dynamic_pointer_cast<AscNode>(
308+ graph.FindNode("add0")->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode());
309+ ASSERT_NE(brc_node, nullptr);
310+ EXPECT_TRUE(brc_node->outputs[0].attr.axis.empty());
311+ EXPECT_TRUE(brc_node->outputs[0].attr.repeats.empty());
312+ EXPECT_TRUE(brc_node->outputs[0].attr.strides.empty());
313+}
314+ 
215// ==================== Scalar 直连 Output → 不插入 ====================315// ==================== Scalar 直连 Output → 不插入 ====================
216 316 
217TEST_F(TestScalarBroadcastInsert, ScalarDirectToOutput_NoInsert) {317TEST_F(TestScalarBroadcastInsert, ScalarDirectToOutput_NoInsert) {