已合并
【PR】: broadcast backward revert #1912
czways创建于 5 天前
【PR】: broadcast backward revert #1912
已合并
共 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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 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 | - | ||
| @@ -12,7 +12,6 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | - | ||
| 16 | 15 | ||
| 17 | 16 | ||
| 18 | 17 | ||
| @@ -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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 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 | - | ||
| @@ -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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 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 | - | ||
| @@ -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 | ||
| 1937 | TEST_F(OptimizerSt, ScalarBroadcastOptimization_Two_Scalar) { | 1937 | TEST_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 chain | 2005 | + // 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 template | 2020 | + // 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 | ||
| 2347 | TEST_F(OptimizerSt, BufQueAllocator_Inplace) { | 2361 | TEST_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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 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 | ||
| 3413 | TEST_F(TestOptimizer, ScalarBroadcastOptimization_Same_Input) { | 3406 | TEST_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 axes | 408 | .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 | ||
| 437 | TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) { | 430 | TEST_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 | ||
| 512 | TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) { | 513 | TEST_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 | ||
| 587 | TEST_F(OptimizerStV2, NddmaCaseLargeTailBrcScoreFunc) { | 598 | TEST_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 1 | 610 | .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 | ||
| 614 | TEST_F(OptimizerStV2, NddmaCaseLargeTailBrc_Dynamic) { | 632 | TEST_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 1 | 644 | .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 | ||
| 82 | TEST_F(SameSourceBroadcastCseStTest, SkipsGraphWithoutNormStructureThroughGraphPassRunner) { | 77 | TEST_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 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | - | ||
| 16 | 15 | ||
| 17 | 16 | ||
| 18 | 17 | ||
| @@ -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>(); |