已合并
【PR】fix :Scalar 多下游不同 shape 时 Broadcast 属性复用导致图校验失败 #1165
czways创建于 7月2日
【PR】fix :Scalar 多下游不同 shape 时 Broadcast 属性复用导致图校验失败 #1165
已合并
共 2 个文件变更+158-11
| @@ -16,6 +16,7 @@ | |||
| 16 | 16 | ||
| 17 | 17 | ||
| 18 | 18 | ||
| 19 | + | ||
| 19 | 20 | ||
| 20 | 21 | ||
| 21 | namespace af::pre_process { | 22 | namespace 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 | + | ||
| 29 | af::AscNodePtr BuildBroadcastNode(af::AscGraph &asc_graph, const af::AscNodePtr &scalar_node, | 54 | af::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 | + | ||
| 40 | Status InsertBroadcastAfterScalar(af::AscGraph &asc_graph, const af::AscNodePtr &scalar_node, bool &inserted) { | 81 | Status 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 | ||
| 217 | TEST_F(TestScalarBroadcastInsert, ScalarDirectToOutput_NoInsert) { | 317 | TEST_F(TestScalarBroadcastInsert, ScalarDirectToOutput_NoInsert) { |