已合并
【PR】: broadcast 后移 优化 #1782
czways创建于 14 天前
【PR】: broadcast 后移 优化 #1782
已合并
共 16 个文件变更+4380-133
| @@ -0,0 +1,1324 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2025 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS 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 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | + | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | + | ||
| 28 | + | ||
| 29 | +using namespace ascir; | ||
| 30 | +using namespace af::ascir_op; | ||
| 31 | +using namespace af::ops; | ||
| 32 | + | ||
| 33 | +namespace optimize { | ||
| 34 | +namespace { | ||
| 35 | + | ||
| 36 | +using af::AscGraph; | ||
| 37 | +using af::AscNode; | ||
| 38 | +using af::AscNodePtr; | ||
| 39 | +using af::AscTensorAttr; | ||
| 40 | +using af::Expression; | ||
| 41 | +using af::FAILED; | ||
| 42 | +using af::SUCCESS; | ||
| 43 | +using NodePtr = af::AscNodePtr; | ||
| 44 | + | ||
| 45 | +constexpr const char *kStoreType = Store::Type; | ||
| 46 | +constexpr const char *kScalarType = Scalar::Type; | ||
| 47 | +constexpr const char *kBroadcastType = Broadcast::Type; | ||
| 48 | +constexpr const char *kCastType = Cast::Type; | ||
| 49 | + | ||
| 50 | +std::vector<std::string> view_op_type = {Transpose::Type, Broadcast::Type, "Slice", Split::Type, Concat::Type, | ||
| 51 | + Gather::Type, "Sum", "Mean", "Max", "Min", | ||
| 52 | + "Prod", "Any", "All"}; | ||
| 53 | + | ||
| 54 | +// -------------------- compat shim(替代 GE asc_adapt:: / BackendUtils:: / AutofuseUtils::) -------------------- | ||
| 55 | + | ||
| 56 | +AscNodePtr ToAscNode(const af::NodePtr &node) { | ||
| 57 | + return std::dynamic_pointer_cast<AscNode>(node); | ||
| 58 | +} | ||
| 59 | + | ||
| 60 | +bool IsEqOne(const Expression &expr) { | ||
| 61 | + return af::SymbolicUtils::StaticCheckEq(expr, af::sym::kSymbolOne) == af::TriBool::kTrue; | ||
| 62 | +} | ||
| 63 | + | ||
| 64 | +bool IsEqZero(const Expression &expr) { | ||
| 65 | + return af::SymbolicUtils::StaticCheckEq(expr, af::sym::kSymbolZero) == af::TriBool::kTrue; | ||
| 66 | +} | ||
| 67 | + | ||
| 68 | +std::string VectorToStr(const std::vector<bool> &vec) { | ||
| 69 | + std::string str = "["; | ||
| 70 | + for (size_t i = 0U; i < vec.size(); ++i) { | ||
| 71 | + if (i > 0U) { | ||
| 72 | + str += ", "; | ||
| 73 | + } | ||
| 74 | + str += vec[i] ? "true" : "false"; | ||
| 75 | + } | ||
| 76 | + str += "]"; | ||
| 77 | + return str; | ||
| 78 | +} | ||
| 79 | + | ||
| 80 | +Status GetPeerInNodes(const NodePtr &node, std::vector<NodePtr> &vec, int32_t idx) { | ||
| 81 | + auto out_anchor = node->GetOutDataAnchor(idx); | ||
| 82 | + GE_ASSERT_NOTNULL(out_anchor); | ||
| 83 | + for (const auto &peer_in : out_anchor->GetPeerInDataAnchors()) { | ||
| 84 | + GE_ASSERT_NOTNULL(peer_in); | ||
| 85 | + vec.push_back(ToAscNode(peer_in->GetOwnerNode())); | ||
| 86 | + } | ||
| 87 | + return SUCCESS; | ||
| 88 | +} | ||
| 89 | + | ||
| 90 | +Status GetPeerOutNode(const NodePtr &node, NodePtr &peer, int32_t idx) { | ||
| 91 | + auto in_anchor = node->GetInDataAnchor(idx); | ||
| 92 | + GE_ASSERT_NOTNULL(in_anchor); | ||
| 93 | + auto out_anchor = in_anchor->GetPeerOutAnchor(); | ||
| 94 | + GE_ASSERT_NOTNULL(out_anchor); | ||
| 95 | + peer = ToAscNode(out_anchor->GetOwnerNode()); | ||
| 96 | + return SUCCESS; | ||
| 97 | +} | ||
| 98 | + | ||
| 99 | +Status GetPeerOutNodes(const NodePtr &node, std::vector<NodePtr> &vec) { | ||
| 100 | + auto in_size = node->GetAllInDataAnchorsSize(); | ||
| 101 | + for (uint32_t i = 0U; i < in_size; ++i) { | ||
| 102 | + auto in_anchor = node->GetInDataAnchor(static_cast<int32_t>(i)); | ||
| 103 | + if (in_anchor == nullptr) { | ||
| 104 | + continue; | ||
| 105 | + } | ||
| 106 | + auto out_anchor = in_anchor->GetPeerOutAnchor(); | ||
| 107 | + if (out_anchor == nullptr) { | ||
| 108 | + continue; | ||
| 109 | + } | ||
| 110 | + vec.push_back(ToAscNode(out_anchor->GetOwnerNode())); | ||
| 111 | + } | ||
| 112 | + return SUCCESS; | ||
| 113 | +} | ||
| 114 | + | ||
| 115 | +Status GetOutputTensorAttr(const NodePtr &node, AscTensorAttr *&attr) { | ||
| 116 | + GE_ASSERT_NOTNULL(node); | ||
| 117 | + GE_ASSERT_TRUE(node->GetAllOutDataAnchorsSize() > 0U); | ||
| 118 | + attr = &node->outputs[0].attr; | ||
| 119 | + return SUCCESS; | ||
| 120 | +} | ||
| 121 | + | ||
| 122 | +bool IsSingleInAndOutNode(const NodePtr &node) { | ||
| 123 | + return node->GetInDataNodesSize() == 1UL && node->GetOutDataNodesSize() == 1UL; | ||
| 124 | +} | ||
| 125 | + | ||
| 126 | +bool IsSingleInNode(const NodePtr &node) { | ||
| 127 | + return node->GetInDataNodesSize() == 1UL; | ||
| 128 | +} | ||
| 129 | + | ||
| 130 | +bool IsSingleOutNode(const NodePtr &node) { | ||
| 131 | + if (node->GetAllOutDataAnchorsSize() != 1U) { | ||
| 132 | + return false; | ||
| 133 | + } | ||
| 134 | + auto out_anchor = node->GetOutDataAnchor(0); | ||
| 135 | + if (out_anchor == nullptr) { | ||
| 136 | + return false; | ||
| 137 | + } | ||
| 138 | + return out_anchor->GetPeerInDataAnchors().size() == 1U; | ||
| 139 | +} | ||
| 140 | + | ||
| 141 | +void RemoveDuplicates(std::vector<NodePtr> &vec) { | ||
| 142 | + std::sort(vec.begin(), vec.end()); | ||
| 143 | + vec.erase(std::unique(vec.begin(), vec.end()), vec.end()); | ||
| 144 | +} | ||
| 145 | + | ||
| 146 | +// -------------------- 算法逻辑(移植自 GE broadcast_backward_pass.cpp) -------------------- | ||
| 147 | + | ||
| 148 | +Status GetSingleNextNode(NodePtr &node, NodePtr &peer_in_node) { | ||
| 149 | + std::vector<NodePtr> peer_in_nodes; | ||
| 150 | + GE_ASSERT_SUCCESS(GetPeerInNodes(node, peer_in_nodes, 0)); | ||
| 151 | + | ||
| 152 | + if (peer_in_nodes.size() != 1U) { | ||
| 153 | + GELOGI("node:%s(%s) has %zu peer out nodes", node->GetName().c_str(), node->GetType().c_str(), | ||
| 154 | + peer_in_nodes.size()); | ||
| 155 | + return FAILED; | ||
| 156 | + } | ||
| 157 | + peer_in_node = peer_in_nodes.at(0); | ||
| 158 | + return SUCCESS; | ||
| 159 | +} | ||
| 160 | + | ||
| 161 | +Status GetPeerOutNodeSafe(const NodePtr &node, NodePtr &peer_out_node, int32_t idx) { | ||
| 162 | + GE_ASSERT_NOTNULL(node); | ||
| 163 | + | ||
| 164 | + if (node->GetAllInDataAnchorsSize() <= static_cast<size_t>(idx)) { | ||
| 165 | + return FAILED; | ||
| 166 | + } | ||
| 167 | + | ||
| 168 | + auto in_anchor = node->GetInDataAnchor(idx); | ||
| 169 | + GE_ASSERT_NOTNULL(in_anchor); | ||
| 170 | + | ||
| 171 | + auto out_anchor = in_anchor->GetPeerOutAnchor(); | ||
| 172 | + GE_ASSERT_NOTNULL(out_anchor); | ||
| 173 | + | ||
| 174 | + return GetPeerOutNode(node, peer_out_node, idx); | ||
| 175 | +} | ||
| 176 | + | ||
| 177 | +bool IsNextViewOp(const NodePtr &next_node) { | ||
| 178 | + std::string type = next_node->GetType(); | ||
| 179 | + return std::find(view_op_type.begin(), view_op_type.end(), type) != view_op_type.end(); | ||
| 180 | +} | ||
| 181 | + | ||
| 182 | +bool IsDtypeNotSupportOp(const NodePtr &next_node, af::DataType &output_dtype) { | ||
| 183 | + std::vector<af::DataType> input_dtypes; | ||
| 184 | + std::vector<af::DataType> expect_output_dtypes; | ||
| 185 | + const auto output_tensor_desc = next_node->GetOpDesc()->MutableOutputDesc(0); | ||
| 186 | + output_dtype = output_tensor_desc->GetDataType(); | ||
| 187 | + expect_output_dtypes.push_back(output_dtype); | ||
| 188 | + input_dtypes.push_back(output_dtype); | ||
| 189 | + return (next_node->GetType() == kCastType) && | ||
| 190 | + (ScheduleUtils::CallAscirInferDataType<Broadcast>(input_dtypes, expect_output_dtypes) != SUCCESS); | ||
| 191 | +} | ||
| 192 | + | ||
| 193 | +Status ReverseCollectBrcNodes(const NodePtr &node, std::vector<NodePtr> &bro_nodes) { | ||
| 194 | + NodePtr cur_node = node; | ||
| 195 | + while ((cur_node->GetType() == kBroadcastType) && IsSingleInAndOutNode(cur_node)) { | ||
| 196 | + bro_nodes.push_back(cur_node); | ||
| 197 | + GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(cur_node, cur_node, 0)); | ||
| 198 | + } | ||
| 199 | + return SUCCESS; | ||
| 200 | +} | ||
| 201 | + | ||
| 202 | +Status GetBroAxisFromNode(const NodePtr &bro_node, int64_t &bro_axis) { | ||
| 203 | + bro_axis = -1; | ||
| 204 | + NodePtr pre_bro_node; | ||
| 205 | + GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(bro_node, pre_bro_node, 0)); | ||
| 206 | + AscTensorAttr *pre_bro_output_attr = nullptr; | ||
| 207 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(pre_bro_node, pre_bro_output_attr)); | ||
| 208 | + auto pre_bro_repeats = pre_bro_output_attr->repeats; | ||
| 209 | + auto pre_bro_strides = pre_bro_output_attr->strides; | ||
| 210 | + auto pre_bro_axis = pre_bro_output_attr->axis; | ||
| 211 | + | ||
| 212 | + AscTensorAttr *bro_output_attr = nullptr; | ||
| 213 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(bro_node, bro_output_attr)); | ||
| 214 | + auto bro_repeats = bro_output_attr->repeats; | ||
| 215 | + auto bro_strides = bro_output_attr->strides; | ||
| 216 | + auto bro_attr_axis = bro_output_attr->axis; | ||
| 217 | + GE_ASSERT_TRUE(pre_bro_repeats.size() == pre_bro_strides.size()); | ||
| 218 | + GE_ASSERT_TRUE(bro_repeats.size() == bro_attr_axis.size()); | ||
| 219 | + GE_ASSERT_TRUE(bro_strides.size() == bro_attr_axis.size()); | ||
| 220 | + GE_ASSERT_TRUE(pre_bro_repeats.size() == pre_bro_axis.size()); | ||
| 221 | + for (size_t index = 0U; index < bro_attr_axis.size(); index++) { | ||
| 222 | + if (IsEqOne(bro_repeats[index])) { | ||
| 223 | + continue; | ||
| 224 | + } | ||
| 225 | + | ||
| 226 | + const auto pre_axis_iter = std::find(pre_bro_axis.begin(), pre_bro_axis.end(), bro_attr_axis[index]); | ||
| 227 | + if (pre_axis_iter == pre_bro_axis.end()) { | ||
| 228 | + // A missing input axis is a scalar/implicit broadcast dimension. | ||
| 229 | + bro_axis = bro_attr_axis[index]; | ||
| 230 | + return SUCCESS; | ||
| 231 | + } | ||
| 232 | + | ||
| 233 | + const size_t pre_index = static_cast<size_t>(std::distance(pre_bro_axis.begin(), pre_axis_iter)); | ||
| 234 | + if (IsEqOne(pre_bro_repeats[pre_index]) && IsEqZero(pre_bro_strides[pre_index])) { | ||
| 235 | + bro_axis = bro_attr_axis[index]; | ||
| 236 | + return SUCCESS; | ||
| 237 | + } | ||
| 238 | + } | ||
| 239 | + GELOGW("Cannot infer broadcast axis: broadcast[%s], source[%s].", bro_node->GetName().c_str(), | ||
| 240 | + pre_bro_node->GetName().c_str()); | ||
| 241 | + return FAILED; | ||
| 242 | +} | ||
| 243 | + | ||
| 244 | +Status GetBroAxises(const std::vector<NodePtr> &bro_nodes, std::vector<int64_t> &bro_axis_idx) { | ||
| 245 | + for (const auto &bro_node : bro_nodes) { | ||
| 246 | + int64_t bro_axis = -1; | ||
| 247 | + if (GetBroAxisFromNode(bro_node, bro_axis) != SUCCESS) { | ||
| 248 | + GELOGI("GetBroAxisFromNode failed for node %s(%s), skipping.", bro_node->GetName().c_str(), | ||
| 249 | + bro_node->GetType().c_str()); | ||
| 250 | + continue; | ||
| 251 | + } | ||
| 252 | + if (bro_axis >= 0) { | ||
| 253 | + bro_axis_idx.push_back(bro_axis); | ||
| 254 | + } | ||
| 255 | + } | ||
| 256 | + return SUCCESS; | ||
| 257 | +} | ||
| 258 | + | ||
| 259 | +Status GetBroAxisesIndex(std::vector<size_t> &bro_axis_idx, const std::vector<Expression> &pre_bro_repeats, | ||
| 260 | + const std::vector<Expression> &pre_bro_strides, | ||
| 261 | + const std::vector<Expression> &last_bro_repeats) { | ||
| 262 | + GE_ASSERT_TRUE(pre_bro_repeats.size() == pre_bro_strides.size()); | ||
| 263 | + GE_ASSERT_TRUE(pre_bro_repeats.size() == last_bro_repeats.size()); | ||
| 264 | + for (size_t index = 0U; index < pre_bro_repeats.size(); index++) { | ||
| 265 | + if (IsEqOne(pre_bro_repeats[index]) && IsEqZero(pre_bro_strides[index])) { | ||
| 266 | + if (IsEqOne(last_bro_repeats[index])) { | ||
| 267 | + continue; | ||
| 268 | + } | ||
| 269 | + bro_axis_idx.push_back(index); | ||
| 270 | + } | ||
| 271 | + } | ||
| 272 | + return SUCCESS; | ||
| 273 | +} | ||
| 274 | + | ||
| 275 | +bool IsSameBroNodes(const std::vector<NodePtr> &bro_nodes1, const std::vector<NodePtr> &bro_nodes2) { | ||
| 276 | + if (bro_nodes1.size() != bro_nodes2.size()) { | ||
| 277 | + return false; | ||
| 278 | + } | ||
| 279 | + std::vector<int64_t> bro_axis_idx1; | ||
| 280 | + std::vector<int64_t> bro_axis_idx2; | ||
| 281 | + GetBroAxises(bro_nodes1, bro_axis_idx1); | ||
| 282 | + GetBroAxises(bro_nodes2, bro_axis_idx2); | ||
| 283 | + return bro_axis_idx1 == bro_axis_idx2; | ||
| 284 | +} | ||
| 285 | + | ||
| 286 | +Status RemoveAndRelinkNodeEdge(af::InDataAnchorPtr &bro_in_anchor, af::OutDataAnchorPtr &bro_out_anchor) { | ||
| 287 | + GE_ASSERT_NOTNULL(bro_in_anchor); | ||
| 288 | + GE_ASSERT_NOTNULL(bro_out_anchor); | ||
| 289 | + GE_ASSERT_TRUE(!bro_out_anchor->GetPeerInDataAnchors().empty()); | ||
| 290 | + auto before_bro_out_anchor = bro_in_anchor->GetPeerOutAnchor(); | ||
| 291 | + auto after_bro_in_anchor = bro_out_anchor->GetPeerInDataAnchors().at(0); | ||
| 292 | + GE_ASSERT_NOTNULL(before_bro_out_anchor); | ||
| 293 | + GE_ASSERT_NOTNULL(after_bro_in_anchor); | ||
| 294 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(before_bro_out_anchor, bro_in_anchor)); | ||
| 295 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(bro_out_anchor, after_bro_in_anchor)); | ||
| 296 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(before_bro_out_anchor, after_bro_in_anchor)); | ||
| 297 | + return SUCCESS; | ||
| 298 | +} | ||
| 299 | + | ||
| 300 | +Status RemoveBroadcastOneByOne(std::vector<NodePtr> &bro_nodes, AscGraph &graph) { | ||
| 301 | + for (auto &node : bro_nodes) { | ||
| 302 | + auto bro_in_anchor = node->GetInDataAnchor(0); | ||
| 303 | + auto bro_out_anchor = node->GetOutDataAnchor(0); | ||
| 304 | + GE_ASSERT_SUCCESS(RemoveAndRelinkNodeEdge(bro_in_anchor, bro_out_anchor)); | ||
| 305 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveNodeWithoutRelink(af::AscGraphUtils::GetComputeGraph(graph), node)); | ||
| 306 | + af::NodeUtils::UnlinkAll(*node); | ||
| 307 | + } | ||
| 308 | + return SUCCESS; | ||
| 309 | +} | ||
| 310 | + | ||
| 311 | +Status RemoveBroadcasts(std::vector<NodePtr> &bro_nodes, AscGraph &graph) { | ||
| 312 | + auto bro_in_anchor = bro_nodes.front()->GetInDataAnchor(0); | ||
| 313 | + auto bro_out_anchor = bro_nodes.back()->GetOutDataAnchor(0); | ||
| 314 | + GE_ASSERT_SUCCESS(RemoveAndRelinkNodeEdge(bro_in_anchor, bro_out_anchor)); | ||
| 315 | + for (auto &node : bro_nodes) { | ||
| 316 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveNodeWithoutRelink(af::AscGraphUtils::GetComputeGraph(graph), node)); | ||
| 317 | + af::NodeUtils::UnlinkAll(*node); | ||
| 318 | + } | ||
| 319 | + return SUCCESS; | ||
| 320 | +} | ||
| 321 | + | ||
| 322 | +std::set<int64_t> FindSubSet(std::vector<int64_t> &bro_axis_idx1, std::vector<int64_t> &bro_axis_idx2) { | ||
| 323 | + std::set<int64_t> common_elements; | ||
| 324 | + if (bro_axis_idx1.empty() || bro_axis_idx2.empty()) { | ||
| 325 | + return common_elements; | ||
| 326 | + } | ||
| 327 | + std::sort(bro_axis_idx1.begin(), bro_axis_idx1.end(), std::greater<int64_t>()); | ||
| 328 | + std::sort(bro_axis_idx2.begin(), bro_axis_idx2.end(), std::greater<int64_t>()); | ||
| 329 | + auto it1 = bro_axis_idx1.begin(); | ||
| 330 | + auto it2 = bro_axis_idx2.begin(); | ||
| 331 | + while (it1 != bro_axis_idx1.end() && it2 != bro_axis_idx2.end()) { | ||
| 332 | + if (*it1 == *it2) { | ||
| 333 | + common_elements.insert(*it1); | ||
| 334 | + ++it1; | ||
| 335 | + ++it2; | ||
| 336 | + } else if (*it1 > *it2) { | ||
| 337 | + ++it1; | ||
| 338 | + } else { | ||
| 339 | + ++it2; | ||
| 340 | + } | ||
| 341 | + } | ||
| 342 | + return common_elements; | ||
| 343 | +} | ||
| 344 | + | ||
| 345 | +bool CollectSameBrcAxis(NodePtr &cur_node, NodePtr &next_node, std::vector<std::vector<NodePtr>> &bro_nodes_list, | ||
| 346 | + std::set<int64_t> &common_axes, std::vector<NodePtr> &origin_bro_nodes) { | ||
| 347 | + std::vector<NodePtr> peer_out_nodes; | ||
| 348 | + GE_ASSERT_SUCCESS(GetPeerOutNodes(next_node, peer_out_nodes)); | ||
| 349 | + for (const auto &node : peer_out_nodes) { | ||
| 350 | + if ((cur_node != nullptr) && (node == cur_node)) { | ||
| 351 | + continue; | ||
| 352 | + } | ||
| 353 | + if (node->GetType() != kBroadcastType) { | ||
| 354 | + bro_nodes_list.clear(); | ||
| 355 | + common_axes.clear(); | ||
| 356 | + return false; | ||
| 357 | + } | ||
| 358 | + if (origin_bro_nodes.empty()) { | ||
| 359 | + ReverseCollectBrcNodes(node, origin_bro_nodes); | ||
| 360 | + std::reverse(origin_bro_nodes.begin(), origin_bro_nodes.end()); | ||
| 361 | + bro_nodes_list.push_back(origin_bro_nodes); | ||
| 362 | + continue; | ||
| 363 | + } | ||
| 364 | + std::vector<NodePtr> temp_bro_nodes; | ||
| 365 | + ReverseCollectBrcNodes(node, temp_bro_nodes); | ||
| 366 | + std::reverse(temp_bro_nodes.begin(), temp_bro_nodes.end()); | ||
| 367 | + bro_nodes_list.push_back(temp_bro_nodes); | ||
| 368 | + | ||
| 369 | + std::vector<int64_t> bro_axis_idx1; | ||
| 370 | + std::vector<int64_t> bro_axis_idx2; | ||
| 371 | + if (common_axes.empty()) { | ||
| 372 | + GetBroAxises(origin_bro_nodes, bro_axis_idx1); | ||
| 373 | + } else { | ||
| 374 | + bro_axis_idx1.assign(common_axes.begin(), common_axes.end()); | ||
| 375 | + } | ||
| 376 | + GetBroAxises(temp_bro_nodes, bro_axis_idx2); | ||
| 377 | + std::set<int64_t> temp_common_axes = FindSubSet(bro_axis_idx1, bro_axis_idx2); | ||
| 378 | + | ||
| 379 | + if (temp_common_axes.empty()) { | ||
| 380 | + bro_nodes_list.clear(); | ||
| 381 | + common_axes.clear(); | ||
| 382 | + return false; | ||
| 383 | + } | ||
| 384 | + common_axes = temp_common_axes; | ||
| 385 | + origin_bro_nodes = origin_bro_nodes.size() < temp_bro_nodes.size() ? origin_bro_nodes : temp_bro_nodes; | ||
| 386 | + } | ||
| 387 | + return true; | ||
| 388 | +} | ||
| 389 | + | ||
| 390 | +Status GetNodeScalarInputList(const af::AscNodePtr &asc_node, std::vector<bool> &is_scalar_list) { | ||
| 391 | + is_scalar_list.resize(asc_node->GetInDataNodesSize(), false); | ||
| 392 | + for (size_t i = 0UL; i < is_scalar_list.size(); ++i) { | ||
| 393 | + const std::vector<Expression> repeats = asc_node->inputs[i].attr.repeats; | ||
| 394 | + is_scalar_list[i] = ascgen_utils::IsScalarInput(repeats); | ||
| 395 | + } | ||
| 396 | + return SUCCESS; | ||
| 397 | +} | ||
| 398 | + | ||
| 399 | +Status ProcessOtherInputBranches(const NodePtr &next_comp_op, size_t current_idx, const std::vector<int64_t> &bro_axes, | ||
| 400 | + std::vector<bool> &is_scalar_list); | ||
| 401 | + | ||
| 402 | +bool CheckBackwardCommon(const NodePtr &next_node) { | ||
| 403 | + if (next_node->GetType() == kStoreType) { | ||
| 404 | + return false; | ||
| 405 | + } | ||
| 406 | + // IndirectLoad 的输出轴属于独立的物理视图,Broadcast 后移不能跨过该边界。 | ||
| 407 | + if (IsOps<af::ascir_op::IndirectLoad>(next_node)) { | ||
| 408 | + return false; | ||
| 409 | + } | ||
| 410 | + if (ScheduleUtils::IsRemovePad(next_node)) { | ||
| 411 | + return false; | ||
| 412 | + } | ||
| 413 | + if (IsNextViewOp(next_node)) { | ||
| 414 | + return false; | ||
| 415 | + } | ||
| 416 | + af::DataType output_dtype; | ||
| 417 | + if (IsDtypeNotSupportOp(next_node, output_dtype)) { | ||
| 418 | + GELOGI("Node %s(%s) cannot backward with dtype(%s)", next_node->GetName().c_str(), next_node->GetType().c_str(), | ||
| 419 | + af::TypeUtils::DataTypeToSerialString(output_dtype).c_str()); | ||
| 420 | + return false; | ||
| 421 | + } | ||
| 422 | + return true; | ||
| 423 | +} | ||
| 424 | + | ||
| 425 | +bool CanBackwardSimplified(const NodePtr &next_node) { | ||
| 426 | + if (!CheckBackwardCommon(next_node)) { | ||
| 427 | + return false; | ||
| 428 | + } | ||
| 429 | + if (next_node->GetAllOutDataAnchorsSize() > 1U) { | ||
| 430 | + return false; | ||
| 431 | + } | ||
| 432 | + if (!IsSingleInNode(next_node)) { | ||
| 433 | + return false; | ||
| 434 | + } | ||
| 435 | + return true; | ||
| 436 | +} | ||
| 437 | + | ||
| 438 | +bool IsScalarInput(const NodePtr &input_node) { | ||
| 439 | + NodePtr temp_node = input_node; | ||
| 440 | + while (temp_node != nullptr) { | ||
| 441 | + if (temp_node->GetType() == kScalarType) { | ||
| 442 | + return true; | ||
| 443 | + } | ||
| 444 | + if (temp_node->GetType() != kBroadcastType) { | ||
| 445 | + break; | ||
| 446 | + } | ||
| 447 | + NodePtr pre_node; | ||
| 448 | + GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(temp_node, pre_node, 0)); | ||
| 449 | + temp_node = pre_node; | ||
| 450 | + } | ||
| 451 | + return false; | ||
| 452 | +} | ||
| 453 | + | ||
| 454 | +Status CheckNodeSupportsScalarInput(const NodePtr &compute_node, int32_t input_idx, | ||
| 455 | + const std::vector<int64_t> &bro_axes, bool &is_support) { | ||
| 456 | + is_support = false; | ||
| 457 | + GE_ASSERT_NOTNULL(std::dynamic_pointer_cast<AscNode>(compute_node)); | ||
| 458 | + const auto &asc_node = std::dynamic_pointer_cast<AscNode>(compute_node); | ||
| 459 | + | ||
| 460 | + std::vector<bool> is_scalar_list; | ||
| 461 | + GE_ASSERT_SUCCESS(GetNodeScalarInputList(asc_node, is_scalar_list)); | ||
| 462 | + | ||
| 463 | + if (input_idx >= 0 && static_cast<size_t>(input_idx) < is_scalar_list.size()) { | ||
| 464 | + is_scalar_list[input_idx] = true; | ||
| 465 | + } | ||
| 466 | + | ||
| 467 | + GE_ASSERT_SUCCESS(ProcessOtherInputBranches(compute_node, input_idx, bro_axes, is_scalar_list)); | ||
| 468 | + | ||
| 469 | + is_support = ascgen_utils::IsNodeSupportsScalarInput(asc_node, is_scalar_list); | ||
| 470 | + if (!is_support) { | ||
| 471 | + GELOGD("Compute node %s does not support scalar input, is_scalar_list: %s", compute_node->GetName().c_str(), | ||
| 472 | + VectorToStr(is_scalar_list).c_str()); | ||
| 473 | + } | ||
| 474 | + return SUCCESS; | ||
| 475 | +} | ||
| 476 | + | ||
| 477 | +bool CheckScalarInputSupport(const NodePtr &next_node, const std::vector<NodePtr> &bro_nodes) { | ||
| 478 | + GE_ASSERT_NOTNULL(std::dynamic_pointer_cast<AscNode>(next_node)); | ||
| 479 | + | ||
| 480 | + auto in_data_anchor_size = next_node->GetAllInDataAnchorsSize(); | ||
| 481 | + for (uint32_t i = 0U; i < in_data_anchor_size; ++i) { | ||
| 482 | + NodePtr input_node; | ||
| 483 | + GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(next_node, input_node, i)); | ||
| 484 | + | ||
| 485 | + if (IsScalarInput(input_node)) { | ||
| 486 | + std::vector<int64_t> bro_axes; | ||
| 487 | + if (!bro_nodes.empty()) { | ||
| 488 | + GE_ASSERT_SUCCESS(GetBroAxises(bro_nodes, bro_axes)); | ||
| 489 | + } | ||
| 490 | + bool is_support = false; | ||
| 491 | + GE_ASSERT_SUCCESS(CheckNodeSupportsScalarInput(next_node, i, bro_axes, is_support)); | ||
| 492 | + if (!is_support) { | ||
| 493 | + return false; | ||
| 494 | + } | ||
| 495 | + } | ||
| 496 | + } | ||
| 497 | + return true; | ||
| 498 | +} | ||
| 499 | + | ||
| 500 | +bool IsMulInputsCanBackward(NodePtr &cur_node, NodePtr &next_node, std::vector<NodePtr> &bro_nodes, AscGraph &graph, | ||
| 501 | + std::set<NodePtr> &mul_input_nodes) { | ||
| 502 | + auto in_data_anchor_size = next_node->GetAllInDataAnchorsSize(); | ||
| 503 | + if (in_data_anchor_size == 1U) { | ||
| 504 | + return false; | ||
| 505 | + } | ||
| 506 | + | ||
| 507 | + std::vector<NodePtr> peer_out_nodes; | ||
| 508 | + GE_ASSERT_SUCCESS(GetPeerOutNodes(next_node, peer_out_nodes)); | ||
| 509 | + std::vector<std::vector<NodePtr>> remove_bro_nodes_list; | ||
| 510 | + for (const auto &node : peer_out_nodes) { | ||
| 511 | + if (node == cur_node) { | ||
| 512 | + continue; | ||
| 513 | + } | ||
| 514 | + if (node->GetType() != kBroadcastType) { | ||
| 515 | + return false; | ||
| 516 | + } | ||
| 517 | + | ||
| 518 | + std::vector<NodePtr> temp_bro_nodes; | ||
| 519 | + GE_ASSERT_SUCCESS(ReverseCollectBrcNodes(node, temp_bro_nodes)); | ||
| 520 | + std::reverse(temp_bro_nodes.begin(), temp_bro_nodes.end()); | ||
| 521 | + if (!IsSameBroNodes(bro_nodes, temp_bro_nodes)) { | ||
| 522 | + std::vector<std::vector<NodePtr>> bro_nodes_list; | ||
| 523 | + std::set<int64_t> common_axes; | ||
| 524 | + std::vector<NodePtr> origin_bro_nodes; | ||
| 525 | + origin_bro_nodes.assign(bro_nodes.begin(), bro_nodes.end()); | ||
| 526 | + bro_nodes_list.push_back(origin_bro_nodes); | ||
| 527 | + if (CollectSameBrcAxis(cur_node, next_node, bro_nodes_list, common_axes, origin_bro_nodes)) { | ||
| 528 | + mul_input_nodes.insert(next_node); | ||
| 529 | + } | ||
| 530 | + return false; | ||
| 531 | + } | ||
| 532 | + remove_bro_nodes_list.push_back(temp_bro_nodes); | ||
| 533 | + } | ||
| 534 | + | ||
| 535 | + if (!CheckScalarInputSupport(next_node, bro_nodes)) { | ||
| 536 | + return false; | ||
| 537 | + } | ||
| 538 | + | ||
| 539 | + for (std::vector<NodePtr> &remove_nodes : remove_bro_nodes_list) { | ||
| 540 | + RemoveBroadcasts(remove_nodes, graph); | ||
| 541 | + } | ||
| 542 | + return true; | ||
| 543 | +} | ||
| 544 | + | ||
| 545 | +bool CanBackward(NodePtr &cur_node, NodePtr &next_node, std::vector<NodePtr> &bro_nodes, AscGraph &graph, | ||
| 546 | + std::set<NodePtr> &mul_input_nodes) { | ||
| 547 | + if (!CheckBackwardCommon(next_node)) { | ||
| 548 | + return false; | ||
| 549 | + } | ||
| 550 | + if (!IsSingleOutNode(next_node)) { | ||
| 551 | + return false; | ||
| 552 | + } | ||
| 553 | + if (!IsSingleInNode(next_node) && !IsMulInputsCanBackward(cur_node, next_node, bro_nodes, graph, mul_input_nodes)) { | ||
| 554 | + return false; | ||
| 555 | + } | ||
| 556 | + return true; | ||
| 557 | +} | ||
| 558 | + | ||
| 559 | +Status CollectBroNodes(NodePtr &cur_node, NodePtr &next_node, std::vector<NodePtr> &nodes) { | ||
| 560 | + if (IsSingleInAndOutNode(cur_node) && (cur_node->GetType() == kBroadcastType)) { | ||
| 561 | + nodes.push_back(cur_node); | ||
| 562 | + } | ||
| 563 | + while (IsSingleInAndOutNode(next_node) && (next_node->GetType() == kBroadcastType)) { | ||
| 564 | + nodes.push_back(next_node); | ||
| 565 | + cur_node = next_node; | ||
| 566 | + GE_ASSERT_SUCCESS(GetSingleNextNode(cur_node, next_node)); | ||
| 567 | + } | ||
| 568 | + return SUCCESS; | ||
| 569 | +} | ||
| 570 | + | ||
| 571 | +Status CollectCmpNodes(NodePtr &cur_node, NodePtr &next_node, std::vector<NodePtr> &nodes, | ||
| 572 | + std::vector<NodePtr> &bro_nodes, AscGraph &graph, std::set<NodePtr> &mul_input_nodes) { | ||
| 573 | + while (CanBackward(cur_node, next_node, bro_nodes, graph, mul_input_nodes)) { | ||
| 574 | + nodes.push_back(next_node); | ||
| 575 | + cur_node = next_node; | ||
| 576 | + GE_ASSERT_SUCCESS(GetSingleNextNode(cur_node, next_node)); | ||
| 577 | + } | ||
| 578 | + return SUCCESS; | ||
| 579 | +} | ||
| 580 | + | ||
| 581 | +Status ReorderBroadcasts(std::vector<NodePtr> &compute_nodes, std::vector<NodePtr> &bro_nodes) { | ||
| 582 | + auto bro_in_anchor = bro_nodes.front()->GetInDataAnchor(0); | ||
| 583 | + auto bro_out_anchor = bro_nodes.back()->GetOutDataAnchor(0); | ||
| 584 | + auto comp_out_anchor = compute_nodes.back()->GetOutDataAnchor(0); | ||
| 585 | + GE_ASSERT_NOTNULL(bro_in_anchor); | ||
| 586 | + GE_ASSERT_NOTNULL(bro_out_anchor); | ||
| 587 | + GE_ASSERT_NOTNULL(comp_out_anchor); | ||
| 588 | + GE_ASSERT_TRUE(!comp_out_anchor->GetPeerInDataAnchors().empty()); | ||
| 589 | + GE_ASSERT_TRUE(!bro_out_anchor->GetPeerInDataAnchors().empty()); | ||
| 590 | + auto before_bro_out_anchor = bro_in_anchor->GetPeerOutAnchor(); | ||
| 591 | + auto after_comp_in_anchor = comp_out_anchor->GetPeerInDataAnchors().at(0); | ||
| 592 | + auto comp_in_anchor = bro_out_anchor->GetPeerInDataAnchors().at(0); | ||
| 593 | + GE_ASSERT_NOTNULL(before_bro_out_anchor); | ||
| 594 | + GE_ASSERT_NOTNULL(after_comp_in_anchor); | ||
| 595 | + GE_ASSERT_NOTNULL(comp_in_anchor); | ||
| 596 | + | ||
| 597 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(before_bro_out_anchor, bro_in_anchor)); | ||
| 598 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(bro_out_anchor, comp_in_anchor)); | ||
| 599 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(comp_out_anchor, after_comp_in_anchor)); | ||
| 600 | + | ||
| 601 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(before_bro_out_anchor, comp_in_anchor)); | ||
| 602 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(comp_out_anchor, bro_in_anchor)); | ||
| 603 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(bro_out_anchor, after_comp_in_anchor)); | ||
| 604 | + return SUCCESS; | ||
| 605 | +} | ||
| 606 | + | ||
| 607 | +Status UpdateComputeNodesAscTensorAttr(std::vector<NodePtr> &bro_nodes, std::vector<NodePtr> &compute_nodes, | ||
| 608 | + const NodePtr &pre_bro_node) { | ||
| 609 | + AscTensorAttr *pre_bro_output_attr = nullptr; | ||
| 610 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(pre_bro_node, pre_bro_output_attr)); | ||
| 611 | + | ||
| 612 | + AscTensorAttr *last_bro_output_attr = nullptr; | ||
| 613 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(bro_nodes.back(), last_bro_output_attr)); | ||
| 614 | + | ||
| 615 | + std::vector<size_t> bro_axis_idx; | ||
| 616 | + auto pre_bro_axis = pre_bro_output_attr->axis; | ||
| 617 | + auto pre_bro_repeats = pre_bro_output_attr->repeats; | ||
| 618 | + auto pre_bro_strides = pre_bro_output_attr->strides; | ||
| 619 | + auto last_bro_repeats = last_bro_output_attr->repeats; | ||
| 620 | + GE_ASSERT_SUCCESS(GetBroAxisesIndex(bro_axis_idx, pre_bro_repeats, pre_bro_strides, last_bro_repeats)); | ||
| 621 | + GELOGD("Broadcast chain contains %zu broadcast axes.", bro_axis_idx.size()); | ||
| 622 | + for (const auto &compute_node : compute_nodes) { | ||
| 623 | + AscTensorAttr *compute_output_attr = nullptr; | ||
| 624 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(compute_node, compute_output_attr)); | ||
| 625 | + compute_output_attr->axis = pre_bro_axis; | ||
| 626 | + compute_output_attr->repeats = pre_bro_repeats; | ||
| 627 | + compute_output_attr->strides = pre_bro_strides; | ||
| 628 | + GE_ASSERT_TRUE(compute_output_attr->strides.size() > 0U); | ||
| 629 | + auto it = std::find(bro_axis_idx.begin(), bro_axis_idx.end(), pre_bro_strides.size() - 1U); | ||
| 630 | + if (it != bro_axis_idx.end()) { | ||
| 631 | + compute_output_attr->strides[compute_output_attr->strides.size() - 1U] = af::sym::kSymbolZero; | ||
| 632 | + } | ||
| 633 | + if (pre_bro_node->GetType() != kScalarType) { | ||
| 634 | + GE_ASSERT_SUCCESS( | ||
| 635 | + ScheduleUtils::RecalculateStridesFromRepeats(compute_output_attr->repeats, compute_output_attr->strides)); | ||
| 636 | + } | ||
| 637 | + } | ||
| 638 | + return SUCCESS; | ||
| 639 | +} | ||
| 640 | + | ||
| 641 | +Status UpdateBroadcastNodesDataType(std::vector<NodePtr> &bro_nodes, const NodePtr &last_comp_node) { | ||
| 642 | + const auto last_comp_opdesc = last_comp_node->GetOpDesc(); | ||
| 643 | + GE_ASSERT_NOTNULL(last_comp_opdesc); | ||
| 644 | + const auto comp_output_tensor_desc = last_comp_opdesc->MutableOutputDesc(0); | ||
| 645 | + GE_ASSERT_NOTNULL(comp_output_tensor_desc); | ||
| 646 | + auto last_comp_dtype = comp_output_tensor_desc->GetDataType(); | ||
| 647 | + for (const auto &bro_node : bro_nodes) { | ||
| 648 | + const auto bro_opdesc = bro_node->GetOpDesc(); | ||
| 649 | + GE_ASSERT_NOTNULL(bro_opdesc); | ||
| 650 | + const auto bro_output_tensor_desc = bro_opdesc->MutableOutputDesc(0); | ||
| 651 | + GE_ASSERT_NOTNULL(bro_output_tensor_desc); | ||
| 652 | + bro_output_tensor_desc->SetDataType(last_comp_dtype); | ||
| 653 | + } | ||
| 654 | + return SUCCESS; | ||
| 655 | +} | ||
| 656 | + | ||
| 657 | +Status BroadcastBackwardReally(std::vector<NodePtr> &compute_nodes, std::vector<NodePtr> &bro_nodes, | ||
| 658 | + const NodePtr &pre_bro_node) { | ||
| 659 | + GE_ASSERT_SUCCESS(ReorderBroadcasts(compute_nodes, bro_nodes)); | ||
| 660 | + GE_ASSERT_SUCCESS(UpdateComputeNodesAscTensorAttr(bro_nodes, compute_nodes, pre_bro_node)); | ||
| 661 | + GE_ASSERT_SUCCESS(UpdateBroadcastNodesDataType(bro_nodes, compute_nodes.back())); | ||
| 662 | + return SUCCESS; | ||
| 663 | +} | ||
| 664 | + | ||
| 665 | +Status CollectBranchBroadcastNodes(const NodePtr &input_node, std::vector<NodePtr> &branch_bro_nodes) { | ||
| 666 | + NodePtr temp_node = input_node; | ||
| 667 | + while (temp_node->GetType() == kBroadcastType && IsSingleInAndOutNode(temp_node)) { | ||
| 668 | + branch_bro_nodes.push_back(temp_node); | ||
| 669 | + NodePtr next_temp_node; | ||
| 670 | + if (GetPeerOutNodeSafe(temp_node, next_temp_node, 0) != SUCCESS) { | ||
| 671 | + break; | ||
| 672 | + } | ||
| 673 | + temp_node = next_temp_node; | ||
| 674 | + } | ||
| 675 | + return SUCCESS; | ||
| 676 | +} | ||
| 677 | + | ||
| 678 | +bool HasCommonBroadcastAxis(const std::vector<int64_t> &axes1, const std::vector<int64_t> &axes2) { | ||
| 679 | + for (const auto &axis : axes1) { | ||
| 680 | + if (std::find(axes2.begin(), axes2.end(), axis) != axes2.end()) { | ||
| 681 | + return true; | ||
| 682 | + } | ||
| 683 | + } | ||
| 684 | + return false; | ||
| 685 | +} | ||
| 686 | + | ||
| 687 | +NodePtr GetPreBroadcastNode(const NodePtr &branch_bro_node) { | ||
| 688 | + auto bro_in_anchor = branch_bro_node->GetInDataAnchor(0); | ||
| 689 | + if (bro_in_anchor == nullptr) { | ||
| 690 | + return nullptr; | ||
| 691 | + } | ||
| 692 | + auto peer_out_anchor = bro_in_anchor->GetPeerOutAnchor(); | ||
| 693 | + if (peer_out_anchor == nullptr) { | ||
| 694 | + return nullptr; | ||
| 695 | + } | ||
| 696 | + return ToAscNode(peer_out_anchor->GetOwnerNode()); | ||
| 697 | +} | ||
| 698 | + | ||
| 699 | +Status ProcessSingleInputBranch(const NodePtr &input_node, const std::vector<int64_t> &bro_axes, bool &is_scalar) { | ||
| 700 | + std::vector<NodePtr> branch_bro_nodes; | ||
| 701 | + GE_ASSERT_SUCCESS(CollectBranchBroadcastNodes(input_node, branch_bro_nodes)); | ||
| 702 | + | ||
| 703 | + if (!branch_bro_nodes.empty()) { | ||
| 704 | + std::vector<int64_t> branch_bro_axes; | ||
| 705 | + GE_ASSERT_SUCCESS(GetBroAxises(branch_bro_nodes, branch_bro_axes)); | ||
| 706 | + if (HasCommonBroadcastAxis(bro_axes, branch_bro_axes)) { | ||
| 707 | + NodePtr pre_bro_node = GetPreBroadcastNode(branch_bro_nodes.front()); | ||
| 708 | + if (pre_bro_node != nullptr) { | ||
| 709 | + AscTensorAttr *pre_bro_attr = nullptr; | ||
| 710 | + if (GetOutputTensorAttr(pre_bro_node, pre_bro_attr) == SUCCESS) { | ||
| 711 | + const std::vector<Expression> repeats = pre_bro_attr->repeats; | ||
| 712 | + is_scalar = ascgen_utils::IsScalarInput(repeats); | ||
| 713 | + } | ||
| 714 | + } | ||
| 715 | + } | ||
| 716 | + } | ||
| 717 | + return SUCCESS; | ||
| 718 | +} | ||
| 719 | + | ||
| 720 | +Status ProcessOtherInputBranches(const NodePtr &next_comp_op, size_t current_idx, const std::vector<int64_t> &bro_axes, | ||
| 721 | + std::vector<bool> &is_scalar_list) { | ||
| 722 | + for (size_t i = 0; i < is_scalar_list.size(); ++i) { | ||
| 723 | + if (i == current_idx) { | ||
| 724 | + continue; | ||
| 725 | + } | ||
| 726 | + NodePtr input_node; | ||
| 727 | + if (GetPeerOutNodeSafe(next_comp_op, input_node, static_cast<int32_t>(i)) != SUCCESS) { | ||
| 728 | + continue; | ||
| 729 | + } | ||
| 730 | + bool scalar_flag = is_scalar_list[i]; | ||
| 731 | + GE_ASSERT_SUCCESS(ProcessSingleInputBranch(input_node, bro_axes, scalar_flag)); | ||
| 732 | + is_scalar_list[i] = scalar_flag; | ||
| 733 | + } | ||
| 734 | + return SUCCESS; | ||
| 735 | +} | ||
| 736 | + | ||
| 737 | +Status JudgeNextCompOpSupportsScalarInput(const NodePtr &node, bool &is_next_support_scalar) { | ||
| 738 | + is_next_support_scalar = false; | ||
| 739 | + | ||
| 740 | + GE_ASSERT_TRUE(node->GetAllOutDataAnchorsSize() == 1U); | ||
| 741 | + auto out_anchor = node->GetOutDataAnchor(0); | ||
| 742 | + GE_ASSERT_NOTNULL(out_anchor); | ||
| 743 | + | ||
| 744 | + auto peer_in_anchors = out_anchor->GetPeerInDataAnchors(); | ||
| 745 | + GE_ASSERT_TRUE(!peer_in_anchors.empty()); | ||
| 746 | + | ||
| 747 | + bool all_branches_support_scalar = true; | ||
| 748 | + for (const auto &peer_in_anchor : peer_in_anchors) { | ||
| 749 | + NodePtr branch_start_node = ToAscNode(peer_in_anchor->GetOwnerNode()); | ||
| 750 | + NodePtr cur_node = branch_start_node; | ||
| 751 | + | ||
| 752 | + std::vector<NodePtr> bro_nodes; | ||
| 753 | + NodePtr temp_cur_node = branch_start_node; | ||
| 754 | + NodePtr temp_next_node = branch_start_node; | ||
| 755 | + GE_ASSERT_SUCCESS(CollectBroNodes(temp_cur_node, temp_next_node, bro_nodes)); | ||
| 756 | + | ||
| 757 | + NodePtr compute_node = temp_next_node; | ||
| 758 | + if (compute_node == nullptr || bro_nodes.empty()) { | ||
| 759 | + all_branches_support_scalar = false; | ||
| 760 | + break; | ||
| 761 | + } | ||
| 762 | + | ||
| 763 | + bool is_support = false; | ||
| 764 | + std::vector<int64_t> bro_axes; | ||
| 765 | + if (!bro_nodes.empty()) { | ||
| 766 | + GE_ASSERT_SUCCESS(GetBroAxises(bro_nodes, bro_axes)); | ||
| 767 | + } | ||
| 768 | + GE_ASSERT_SUCCESS(CheckNodeSupportsScalarInput(compute_node, peer_in_anchor->GetIdx(), bro_axes, is_support)); | ||
| 769 | + if (!is_support) { | ||
| 770 | + all_branches_support_scalar = false; | ||
| 771 | + break; | ||
| 772 | + } | ||
| 773 | + } | ||
| 774 | + | ||
| 775 | + is_next_support_scalar = all_branches_support_scalar; | ||
| 776 | + return SUCCESS; | ||
| 777 | +} | ||
| 778 | + | ||
| 779 | +bool ContainsBroadcastNode(const NodePtr &node) { | ||
| 780 | + NodePtr pre_node; | ||
| 781 | + if (GetPeerOutNodeSafe(node, pre_node, 0) != SUCCESS) { | ||
| 782 | + return false; | ||
| 783 | + } | ||
| 784 | + return pre_node->GetType() == kBroadcastType; | ||
| 785 | +} | ||
| 786 | + | ||
| 787 | +Status CollectCandidateMultiRefNodes(const AscGraph &graph, std::vector<NodePtr> &candidate_nodes) { | ||
| 788 | + for (const auto &node : graph.GetAllNodes()) { | ||
| 789 | + if (node->GetAllOutDataAnchorsSize() != 1U) { | ||
| 790 | + continue; | ||
| 791 | + } | ||
| 792 | + auto out_anchor = node->GetOutDataAnchor(0); | ||
| 793 | + if (out_anchor == nullptr || out_anchor->GetPeerInDataAnchors().size() <= 1) { | ||
| 794 | + continue; | ||
| 795 | + } | ||
| 796 | + if (node->GetType() == kBroadcastType || ContainsBroadcastNode(node)) { | ||
| 797 | + candidate_nodes.push_back(node); | ||
| 798 | + } | ||
| 799 | + } | ||
| 800 | + return SUCCESS; | ||
| 801 | +} | ||
| 802 | + | ||
| 803 | +Status ExtractBroadcastChainFromNode(const NodePtr &node, std::vector<NodePtr> &bro_nodes) { | ||
| 804 | + NodePtr cur_node = node; | ||
| 805 | + | ||
| 806 | + if (cur_node != nullptr && cur_node->GetType() != kBroadcastType) { | ||
| 807 | + NodePtr pre_node; | ||
| 808 | + GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(cur_node, pre_node, 0)); | ||
| 809 | + cur_node = pre_node; | ||
| 810 | + } | ||
| 811 | + | ||
| 812 | + while (cur_node != nullptr && cur_node->GetType() == kBroadcastType) { | ||
| 813 | + bro_nodes.push_back(cur_node); | ||
| 814 | + NodePtr pre_node; | ||
| 815 | + if (GetPeerOutNodeSafe(cur_node, pre_node, 0) != SUCCESS) { | ||
| 816 | + break; | ||
| 817 | + } | ||
| 818 | + cur_node = pre_node; | ||
| 819 | + } | ||
| 820 | + | ||
| 821 | + std::reverse(bro_nodes.begin(), bro_nodes.end()); | ||
| 822 | + return SUCCESS; | ||
| 823 | +} | ||
| 824 | + | ||
| 825 | +Status TraceBranchToMergeNode(const NodePtr &start_node, NodePtr &merge_node, std::vector<NodePtr> &branch_nodes) { | ||
| 826 | + NodePtr cur_node = start_node; | ||
| 827 | + std::unordered_set<NodePtr> visited_nodes; | ||
| 828 | + | ||
| 829 | + while (cur_node != nullptr) { | ||
| 830 | + GE_ASSERT_TRUE(visited_nodes.count(cur_node) == 0, "Found cycle dependency in TraceBranchToMergeNode, node: %s", | ||
| 831 | + cur_node->GetName().c_str()); | ||
| 832 | + visited_nodes.insert(cur_node); | ||
| 833 | + | ||
| 834 | + if (!IsSingleInNode(cur_node)) { | ||
| 835 | + merge_node = cur_node; | ||
| 836 | + return SUCCESS; | ||
| 837 | + } | ||
| 838 | + | ||
| 839 | + if (cur_node->GetType() == kStoreType) { | ||
| 840 | + merge_node = cur_node; | ||
| 841 | + return SUCCESS; | ||
| 842 | + } | ||
| 843 | + | ||
| 844 | + if (!IsSingleOutNode(cur_node)) { | ||
| 845 | + return FAILED; | ||
| 846 | + } | ||
| 847 | + | ||
| 848 | + branch_nodes.push_back(cur_node); | ||
| 849 | + | ||
| 850 | + NodePtr next_node; | ||
| 851 | + if (GetSingleNextNode(cur_node, next_node) != SUCCESS) { | ||
| 852 | + return FAILED; | ||
| 853 | + } | ||
| 854 | + cur_node = next_node; | ||
| 855 | + } | ||
| 856 | + return FAILED; | ||
| 857 | +} | ||
| 858 | + | ||
| 859 | +bool CheckBranchNodesSupportBackward(const std::vector<NodePtr> &branch_nodes) { | ||
| 860 | + for (const auto &next_node : branch_nodes) { | ||
| 861 | + if (!CanBackwardSimplified(next_node)) { | ||
| 862 | + return false; | ||
| 863 | + } | ||
| 864 | + } | ||
| 865 | + return true; | ||
| 866 | +} | ||
| 867 | + | ||
| 868 | +bool CheckAllBranchesCanBackward(const NodePtr &multi_ref_node, NodePtr &merge_node, | ||
| 869 | + std::vector<std::vector<NodePtr>> &all_branch_nodes) { | ||
| 870 | + auto out_anchor = multi_ref_node->GetOutDataAnchor(0); | ||
| 871 | + GE_ASSERT_NOTNULL(out_anchor); | ||
| 872 | + auto peer_in_anchors = out_anchor->GetPeerInDataAnchors(); | ||
| 873 | + | ||
| 874 | + GE_ASSERT_TRUE(peer_in_anchors.size() > 1U); | ||
| 875 | + NodePtr first_merge_node = nullptr; | ||
| 876 | + | ||
| 877 | + for (const auto &in_anchor : peer_in_anchors) { | ||
| 878 | + NodePtr branch_start_node = ToAscNode(in_anchor->GetOwnerNode()); | ||
| 879 | + NodePtr current_merge_node = nullptr; | ||
| 880 | + std::vector<NodePtr> branch_nodes; | ||
| 881 | + | ||
| 882 | + if (TraceBranchToMergeNode(branch_start_node, current_merge_node, branch_nodes) != SUCCESS) { | ||
| 883 | + return false; | ||
| 884 | + } | ||
| 885 | + | ||
| 886 | + if (first_merge_node == nullptr) { | ||
| 887 | + first_merge_node = current_merge_node; | ||
| 888 | + } else if (first_merge_node != current_merge_node) { | ||
| 889 | + return false; | ||
| 890 | + } | ||
| 891 | + all_branch_nodes.push_back(branch_nodes); | ||
| 892 | + } | ||
| 893 | + | ||
| 894 | + merge_node = first_merge_node; | ||
| 895 | + return true; | ||
| 896 | +} | ||
| 897 | + | ||
| 898 | +bool CheckAllBranchesSupportBackward(const std::vector<std::vector<NodePtr>> &all_branch_nodes, | ||
| 899 | + const NodePtr &candidate_node) { | ||
| 900 | + if (candidate_node->GetType() != kBroadcastType) { | ||
| 901 | + if (!CanBackwardSimplified(candidate_node)) { | ||
| 902 | + return false; | ||
| 903 | + } | ||
| 904 | + } | ||
| 905 | + for (const auto &branch_nodes : all_branch_nodes) { | ||
| 906 | + if (!CheckBranchNodesSupportBackward(branch_nodes)) { | ||
| 907 | + return false; | ||
| 908 | + } | ||
| 909 | + } | ||
| 910 | + return true; | ||
| 911 | +} | ||
| 912 | + | ||
| 913 | +Status DisconnectBranchesFromBroadcast(const NodePtr &last_bro_node, | ||
| 914 | + std::vector<af::OutDataAnchorPtr> &branch_out_anchors, | ||
| 915 | + std::vector<af::InDataAnchorPtr> &branch_in_anchors) { | ||
| 916 | + auto bro_out_anchor = last_bro_node->GetOutDataAnchor(0); | ||
| 917 | + auto peer_in_anchors = bro_out_anchor->GetPeerInDataAnchors(); | ||
| 918 | + | ||
| 919 | + for (const auto &in_anchor : peer_in_anchors) { | ||
| 920 | + auto branch_node = in_anchor->GetOwnerNode(); | ||
| 921 | + auto branch_in_anchor = branch_node->GetInDataAnchor(in_anchor->GetIdx()); | ||
| 922 | + auto branch_out_anchor = branch_in_anchor->GetPeerOutAnchor(); | ||
| 923 | + | ||
| 924 | + branch_out_anchors.push_back(branch_out_anchor); | ||
| 925 | + branch_in_anchors.push_back(branch_in_anchor); | ||
| 926 | + | ||
| 927 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(branch_out_anchor, branch_in_anchor)); | ||
| 928 | + } | ||
| 929 | + return SUCCESS; | ||
| 930 | +} | ||
| 931 | + | ||
| 932 | +Status MoveBroadcastAfterMerge(const NodePtr &merge_node, const NodePtr &first_bro_node, const NodePtr &last_bro_node) { | ||
| 933 | + auto merge_out_anchor = merge_node->GetOutDataAnchor(0); | ||
| 934 | + if (merge_out_anchor == nullptr || merge_out_anchor->GetPeerInDataAnchors().empty()) { | ||
| 935 | + return FAILED; | ||
| 936 | + } | ||
| 937 | + auto peer_in_anchors = merge_out_anchor->GetPeerInDataAnchors(); | ||
| 938 | + std::vector<af::InDataAnchorPtr> merge_next_in_anchors(peer_in_anchors.begin(), peer_in_anchors.end()); | ||
| 939 | + for (const auto &in_anchor : merge_next_in_anchors) { | ||
| 940 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(merge_out_anchor, in_anchor)); | ||
| 941 | + } | ||
| 942 | + | ||
| 943 | + auto bro_in_anchor = first_bro_node->GetInDataAnchor(0); | ||
| 944 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(merge_out_anchor, bro_in_anchor)); | ||
| 945 | + | ||
| 946 | + auto bro_out_anchor = last_bro_node->GetOutDataAnchor(0); | ||
| 947 | + for (const auto &in_anchor : merge_next_in_anchors) { | ||
| 948 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(bro_out_anchor, in_anchor)); | ||
| 949 | + } | ||
| 950 | + return SUCCESS; | ||
| 951 | +} | ||
| 952 | + | ||
| 953 | +Status BackwardMultiRefBroadcast(const NodePtr &candidate_node, const NodePtr &merge_node, | ||
| 954 | + const std::vector<std::vector<NodePtr>> &all_branch_nodes, | ||
| 955 | + std::vector<NodePtr> &bro_nodes, [[maybe_unused]] AscGraph &graph) { | ||
| 956 | + auto bro_in_anchor = bro_nodes.front()->GetInDataAnchor(0); | ||
| 957 | + auto pre_bro_out_anchor = bro_in_anchor->GetPeerOutAnchor(); | ||
| 958 | + NodePtr pre_bro_node = ToAscNode(pre_bro_out_anchor->GetOwnerNode()); | ||
| 959 | + bool is_pre_scalar = (pre_bro_node->GetType() == kScalarType); | ||
| 960 | + | ||
| 961 | + if (is_pre_scalar) { | ||
| 962 | + std::vector<int64_t> bro_axes; | ||
| 963 | + if (!bro_nodes.empty()) { | ||
| 964 | + GE_ASSERT_SUCCESS(GetBroAxises(bro_nodes, bro_axes)); | ||
| 965 | + } | ||
| 966 | + auto bro_out_anchor = bro_nodes.back()->GetOutDataAnchor(0); | ||
| 967 | + auto peer_in_anchors = bro_out_anchor->GetPeerInDataAnchors(); | ||
| 968 | + for (const auto &in_anchor : peer_in_anchors) { | ||
| 969 | + NodePtr compute_node = ToAscNode(in_anchor->GetOwnerNode()); | ||
| 970 | + int32_t input_idx = in_anchor->GetIdx(); | ||
| 971 | + bool is_support = false; | ||
| 972 | + GE_ASSERT_SUCCESS(CheckNodeSupportsScalarInput(compute_node, input_idx, bro_axes, is_support)); | ||
| 973 | + if (!is_support) { | ||
| 974 | + return FAILED; | ||
| 975 | + } | ||
| 976 | + } | ||
| 977 | + } | ||
| 978 | + | ||
| 979 | + std::vector<af::OutDataAnchorPtr> branch_out_anchors; | ||
| 980 | + std::vector<af::InDataAnchorPtr> branch_in_anchors; | ||
| 981 | + GE_ASSERT_SUCCESS(DisconnectBranchesFromBroadcast(bro_nodes.back(), branch_out_anchors, branch_in_anchors)); | ||
| 982 | + | ||
| 983 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::RemoveEdge(pre_bro_out_anchor, bro_in_anchor)); | ||
| 984 | + GE_ASSERT_SUCCESS(MoveBroadcastAfterMerge(merge_node, bro_nodes.front(), bro_nodes.back())); | ||
| 985 | + | ||
| 986 | + for (size_t i = 0; i < branch_out_anchors.size(); ++i) { | ||
| 987 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::AddEdge(pre_bro_out_anchor, branch_in_anchors[i])); | ||
| 988 | + } | ||
| 989 | + | ||
| 990 | + std::vector<NodePtr> compute_nodes; | ||
| 991 | + if (candidate_node->GetType() != kBroadcastType) { | ||
| 992 | + compute_nodes.push_back(candidate_node); | ||
| 993 | + } | ||
| 994 | + for (const auto &branch : all_branch_nodes) { | ||
| 995 | + compute_nodes.insert(compute_nodes.end(), branch.begin(), branch.end()); | ||
| 996 | + } | ||
| 997 | + compute_nodes.push_back(merge_node); | ||
| 998 | + GE_ASSERT_SUCCESS(UpdateComputeNodesAscTensorAttr(bro_nodes, compute_nodes, pre_bro_node)); | ||
| 999 | + | ||
| 1000 | + if (!compute_nodes.empty()) { | ||
| 1001 | + GE_ASSERT_SUCCESS(UpdateBroadcastNodesDataType(bro_nodes, compute_nodes.back())); | ||
| 1002 | + } | ||
| 1003 | + return SUCCESS; | ||
| 1004 | +} | ||
| 1005 | + | ||
| 1006 | +Status ProcessMultiRefBroadcastBackward(AscGraph &graph, bool &is_changed) { | ||
| 1007 | + std::vector<NodePtr> candidate_nodes; | ||
| 1008 | + GE_ASSERT_SUCCESS(CollectCandidateMultiRefNodes(graph, candidate_nodes)); | ||
| 1009 | + | ||
| 1010 | + for (const auto &candidate_node : candidate_nodes) { | ||
| 1011 | + std::vector<NodePtr> bro_nodes; | ||
| 1012 | + GE_ASSERT_SUCCESS(ExtractBroadcastChainFromNode(candidate_node, bro_nodes)); | ||
| 1013 | + | ||
| 1014 | + if (bro_nodes.empty()) { | ||
| 1015 | + continue; | ||
| 1016 | + } | ||
| 1017 | + | ||
| 1018 | + NodePtr bro_node = bro_nodes.front(); | ||
| 1019 | + NodePtr merge_node = nullptr; | ||
| 1020 | + std::vector<std::vector<NodePtr>> all_branch_nodes; | ||
| 1021 | + | ||
| 1022 | + if (!CheckAllBranchesCanBackward(candidate_node, merge_node, all_branch_nodes)) { | ||
| 1023 | + continue; | ||
| 1024 | + } | ||
| 1025 | + if (!CheckAllBranchesSupportBackward(all_branch_nodes, candidate_node)) { | ||
| 1026 | + continue; | ||
| 1027 | + } | ||
| 1028 | + GELOGI("Move shared broadcast from node[%s] to merge[%s], broadcasts=%zu, branches=%zu.", | ||
| 1029 | + candidate_node->GetName().c_str(), merge_node->GetName().c_str(), bro_nodes.size(), all_branch_nodes.size()); | ||
| 1030 | + | ||
| 1031 | + is_changed = BackwardMultiRefBroadcast(candidate_node, merge_node, all_branch_nodes, bro_nodes, graph) == SUCCESS; | ||
| 1032 | + } | ||
| 1033 | + return SUCCESS; | ||
| 1034 | +} | ||
| 1035 | + | ||
| 1036 | +Status CollectBackwardStartNodes(const AscGraph &graph, std::vector<NodePtr> &pre_brc_nodes) { | ||
| 1037 | + for (const auto &node : graph.GetAllNodes()) { | ||
| 1038 | + NodePtr cur_node = node; | ||
| 1039 | + while ((cur_node->GetType() == kBroadcastType) && IsSingleOutNode(cur_node)) { | ||
| 1040 | + GE_ASSERT_SUCCESS(GetPeerOutNodeSafe(cur_node, cur_node, 0)); | ||
| 1041 | + } | ||
| 1042 | + bool is_next_support_scalar = true; | ||
| 1043 | + if (cur_node->GetType() == kScalarType) { | ||
| 1044 | + GE_ASSERT_SUCCESS(JudgeNextCompOpSupportsScalarInput(cur_node, is_next_support_scalar)); | ||
| 1045 | + } | ||
| 1046 | + | ||
| 1047 | + if ((cur_node != node) && is_next_support_scalar) { | ||
| 1048 | + pre_brc_nodes.push_back(cur_node); | ||
| 1049 | + } | ||
| 1050 | + } | ||
| 1051 | + return SUCCESS; | ||
| 1052 | +} | ||
| 1053 | + | ||
| 1054 | +Status CollectBackwardSatisfyStartNodes(const NodePtr &node, std::vector<NodePtr> &peer_in_nodes) { | ||
| 1055 | + auto output_size = node->GetAllOutDataAnchorsSize(); | ||
| 1056 | + for (uint32_t idx = 0U; idx < output_size; ++idx) { | ||
| 1057 | + std::vector<NodePtr> temp_nodes; | ||
| 1058 | + GE_ASSERT_SUCCESS(GetPeerInNodes(node, temp_nodes, static_cast<int32_t>(idx))); | ||
| 1059 | + if (!temp_nodes.empty() && (temp_nodes.front()->GetType() == kBroadcastType)) { | ||
| 1060 | + peer_in_nodes.insert(peer_in_nodes.end(), temp_nodes.begin(), temp_nodes.end()); | ||
| 1061 | + } | ||
| 1062 | + } | ||
| 1063 | + return SUCCESS; | ||
| 1064 | +} | ||
| 1065 | + | ||
| 1066 | +Status RemoveBroadcasts(AscGraph &graph, std::vector<std::vector<NodePtr>> &bro_nodes_list, | ||
| 1067 | + const std::set<int64_t> &common_axises) { | ||
| 1068 | + for (std::vector<NodePtr> &bro_nodes : bro_nodes_list) { | ||
| 1069 | + std::vector<NodePtr> remove_nodes; | ||
| 1070 | + for (auto it = bro_nodes.begin(); it != bro_nodes.end();) { | ||
| 1071 | + int64_t bro_axis; | ||
| 1072 | + GE_ASSERT_SUCCESS(GetBroAxisFromNode(*it, bro_axis)); | ||
| 1073 | + if ((bro_axis != -1) && common_axises.count(bro_axis) != 0) { | ||
| 1074 | + remove_nodes.push_back(*it); | ||
| 1075 | + it = bro_nodes.erase(it); | ||
| 1076 | + } else { | ||
| 1077 | + ++it; | ||
| 1078 | + } | ||
| 1079 | + } | ||
| 1080 | + GE_ASSERT_SUCCESS(RemoveBroadcastOneByOne(remove_nodes, graph)); | ||
| 1081 | + } | ||
| 1082 | + return SUCCESS; | ||
| 1083 | +} | ||
| 1084 | + | ||
| 1085 | +Status GetBackwardBrcNodes(const std::vector<NodePtr> &origin_bro_nodes, | ||
| 1086 | + std::vector<NodePtr> &origin_need_move_bro_nodes, const std::set<int64_t> &common_axises) { | ||
| 1087 | + for (auto &bro_node : origin_bro_nodes) { | ||
| 1088 | + int64_t bro_axis; | ||
| 1089 | + GE_ASSERT_SUCCESS(GetBroAxisFromNode(bro_node, bro_axis)); | ||
| 1090 | + if ((bro_axis != -1) && common_axises.count(bro_axis) != 0) { | ||
| 1091 | + origin_need_move_bro_nodes.push_back(bro_node); | ||
| 1092 | + } | ||
| 1093 | + } | ||
| 1094 | + return SUCCESS; | ||
| 1095 | +} | ||
| 1096 | + | ||
| 1097 | +struct TensorInfo { | ||
| 1098 | + std::vector<int64_t> axis; | ||
| 1099 | + std::vector<Expression> repeats; | ||
| 1100 | + std::vector<Expression> strides; | ||
| 1101 | + af::DataType dtype; | ||
| 1102 | + int64_t sched_axis; | ||
| 1103 | + std::vector<int64_t> broadcast_info; | ||
| 1104 | +}; | ||
| 1105 | + | ||
| 1106 | +Status GetTensorInfo(const NodePtr &node, TensorInfo &tensor_info) { | ||
| 1107 | + AscTensorAttr *attr = nullptr; | ||
| 1108 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(node, attr)); | ||
| 1109 | + tensor_info.axis = attr->axis; | ||
| 1110 | + tensor_info.repeats = attr->repeats; | ||
| 1111 | + tensor_info.strides = attr->strides; | ||
| 1112 | + tensor_info.dtype = attr->dtype; | ||
| 1113 | + tensor_info.sched_axis = attr->axis.back(); | ||
| 1114 | + return SUCCESS; | ||
| 1115 | +} | ||
| 1116 | + | ||
| 1117 | +Status UpdateBroadcastNodeAttrs(const NodePtr &b_node, const std::vector<int64_t> &axis, | ||
| 1118 | + const std::vector<Expression> &repeats, const std::vector<Expression> &strides, | ||
| 1119 | + int64_t broadcast_axis) { | ||
| 1120 | + AscTensorAttr *attr = nullptr; | ||
| 1121 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(b_node, attr)); | ||
| 1122 | + attr->axis = axis; | ||
| 1123 | + attr->repeats = repeats; | ||
| 1124 | + attr->strides = strides; | ||
| 1125 | + for (size_t i = 0U; i < attr->axis.size(); ++i) { | ||
| 1126 | + if (attr->axis[i] == broadcast_axis) { | ||
| 1127 | + attr->repeats[i] = af::sym::kSymbolOne; | ||
| 1128 | + attr->strides[i] = af::sym::kSymbolZero; | ||
| 1129 | + } | ||
| 1130 | + } | ||
| 1131 | + return SUCCESS; | ||
| 1132 | +} | ||
| 1133 | + | ||
| 1134 | +Status UpdateBroadcastNodeSchedInfo(const NodePtr &b_node, const NodePtr &ref_node) { | ||
| 1135 | + b_node->attr.sched = ref_node->attr.sched; | ||
| 1136 | + return SUCCESS; | ||
| 1137 | +} | ||
| 1138 | + | ||
| 1139 | +Status FromDtypeToOtherDtype(const NodePtr &b_node, af::DataType from_dtype, af::DataType to_dtype) { | ||
| 1140 | + std::vector<af::DataType> input_dtypes = {from_dtype}; | ||
| 1141 | + std::vector<af::DataType> expect_output_dtypes = {to_dtype}; | ||
| 1142 | + GE_ASSERT_SUCCESS(ScheduleUtils::CallAscirInferDataType<Broadcast>(input_dtypes, expect_output_dtypes)); | ||
| 1143 | + b_node->outputs[0].attr.dtype = to_dtype; | ||
| 1144 | + return SUCCESS; | ||
| 1145 | +} | ||
| 1146 | + | ||
| 1147 | +Status CreateAndUpdateBroadcastNode(AscGraph &asc_graph, const NodePtr &node, NodePtr &connect_node, | ||
| 1148 | + TensorInfo &tensor_info) { | ||
| 1149 | + const std::vector<int64_t> &broadcast_info = tensor_info.broadcast_info; | ||
| 1150 | + GE_ASSERT_TRUE(broadcast_info.size() > 0U); | ||
| 1151 | + for (size_t index = 0U; index < broadcast_info.size(); index++) { | ||
| 1152 | + const std::string brc_name = "backward_broadcast_" + node->GetName() + "_" + std::to_string(index); | ||
| 1153 | + Broadcast brc_op(brc_name.c_str()); | ||
| 1154 | + auto b_node = asc_graph.AddNode(brc_op); | ||
| 1155 | + GE_ASSERT_NOTNULL(b_node); | ||
| 1156 | + | ||
| 1157 | + brc_op.attr.sched = node->attr.sched; | ||
| 1158 | + brc_op.attr.api.compute_type = af::ComputeType::kComputeBroadcast; | ||
| 1159 | + brc_op.attr.api.type = af::ApiType::kAPITypeCompute; | ||
| 1160 | + | ||
| 1161 | + int32_t anchor_idx = 0; | ||
| 1162 | + GE_ASSERT_TRUE(node->GetOutDataAnchor(0)->GetPeerInDataAnchors().size() == 1U); | ||
| 1163 | + anchor_idx = node->GetOutDataAnchor(0)->GetPeerInDataAnchors().at(0)->GetIdx(); | ||
| 1164 | + GE_ASSERT_GRAPH_SUCCESS(af::GraphUtils::ReplaceNodeDataAnchors(b_node, connect_node, {anchor_idx}, {})); | ||
| 1165 | + GE_ASSERT_GRAPH_SUCCESS( | ||
| 1166 | + af::GraphUtils::AddEdge(b_node->GetOutDataAnchor(0), connect_node->GetInDataAnchor(anchor_idx))); | ||
| 1167 | + | ||
| 1168 | + GE_ASSERT_SUCCESS(UpdateBroadcastNodeAttrs(b_node, tensor_info.axis, tensor_info.repeats, tensor_info.strides, | ||
| 1169 | + broadcast_info[index])); | ||
| 1170 | + GE_ASSERT_SUCCESS(UpdateBroadcastNodeSchedInfo(b_node, node)); | ||
| 1171 | + GE_ASSERT_SUCCESS(FromDtypeToOtherDtype(b_node, tensor_info.dtype, tensor_info.dtype)); | ||
| 1172 | + connect_node = b_node; | ||
| 1173 | + } | ||
| 1174 | + return SUCCESS; | ||
| 1175 | +} | ||
| 1176 | + | ||
| 1177 | +Status InsertBroadcastNode(NodePtr &pre_bro_node, AscGraph &graph, std::set<int64_t> &broadcast_axis) { | ||
| 1178 | + std::vector<int64_t> broadcast_info(broadcast_axis.begin(), broadcast_axis.end()); | ||
| 1179 | + TensorInfo tensor_info; | ||
| 1180 | + GE_ASSERT_SUCCESS(GetTensorInfo(pre_bro_node, tensor_info)); | ||
| 1181 | + tensor_info.broadcast_info = broadcast_info; | ||
| 1182 | + NodePtr connect_node; | ||
| 1183 | + GE_ASSERT_SUCCESS(GetSingleNextNode(pre_bro_node, connect_node)); | ||
| 1184 | + GE_ASSERT_SUCCESS(CreateAndUpdateBroadcastNode(graph, pre_bro_node, connect_node, tensor_info)); | ||
| 1185 | + return SUCCESS; | ||
| 1186 | +} | ||
| 1187 | + | ||
| 1188 | +Status UpdateOutputTensor(std::vector<NodePtr> &nodes, std::set<int64_t> &common_axises) { | ||
| 1189 | + for (auto &node : nodes) { | ||
| 1190 | + AscTensorAttr *compute_output_attr = nullptr; | ||
| 1191 | + GE_ASSERT_SUCCESS(GetOutputTensorAttr(node, compute_output_attr)); | ||
| 1192 | + auto &repeats = compute_output_attr->repeats; | ||
| 1193 | + auto &strides = compute_output_attr->strides; | ||
| 1194 | + auto &attr_axis = compute_output_attr->axis; | ||
| 1195 | + for (auto common_axis : common_axises) { | ||
| 1196 | + auto it = std::find(attr_axis.begin(), attr_axis.end(), common_axis); | ||
| 1197 | + GE_ASSERT_TRUE(it != attr_axis.end()); | ||
| 1198 | + auto index = std::distance(attr_axis.begin(), it); | ||
| 1199 | + repeats[index] = af::sym::kSymbolOne; | ||
| 1200 | + strides[index] = af::sym::kSymbolZero; | ||
| 1201 | + } | ||
| 1202 | + GE_ASSERT_SUCCESS(ScheduleUtils::RecalculateStridesFromRepeats(repeats, strides)); | ||
| 1203 | + } | ||
| 1204 | + return SUCCESS; | ||
| 1205 | +} | ||
| 1206 | + | ||
| 1207 | +Status JudgePartBackward(std::set<NodePtr> &mul_input_nodes, bool &is_changed, AscGraph &graph) { | ||
| 1208 | + std::set<NodePtr> next_mul_input_nodes; | ||
| 1209 | + for (auto mul_input_node : mul_input_nodes) { | ||
| 1210 | + std::vector<std::vector<NodePtr>> bro_nodes_list; | ||
| 1211 | + std::vector<NodePtr> origin_bro_nodes; | ||
| 1212 | + std::vector<NodePtr> origin_need_move_bro_nodes; | ||
| 1213 | + std::vector<NodePtr> compute_nodes; | ||
| 1214 | + std::set<int64_t> common_axises; | ||
| 1215 | + NodePtr peer_out_node = nullptr; | ||
| 1216 | + if (!CollectSameBrcAxis(peer_out_node, mul_input_node, bro_nodes_list, common_axises, origin_bro_nodes)) { | ||
| 1217 | + continue; | ||
| 1218 | + } | ||
| 1219 | + if (bro_nodes_list.empty() || origin_bro_nodes.size() > 1U) { | ||
| 1220 | + GELOGD("Skip partial broadcast backward at node[%s]: branches=%zu, origin_broadcasts=%zu.", | ||
| 1221 | + mul_input_node->GetName().c_str(), bro_nodes_list.size(), origin_bro_nodes.size()); | ||
| 1222 | + continue; | ||
| 1223 | + } | ||
| 1224 | + | ||
| 1225 | + GE_ASSERT_SUCCESS(GetBackwardBrcNodes(origin_bro_nodes, origin_need_move_bro_nodes, common_axises)); | ||
| 1226 | + | ||
| 1227 | + auto cur_node = mul_input_node; | ||
| 1228 | + auto next_node = mul_input_node; | ||
| 1229 | + GE_ASSERT_SUCCESS(GetSingleNextNode(cur_node, next_node)); | ||
| 1230 | + compute_nodes.push_back(mul_input_node); | ||
| 1231 | + GE_ASSERT_SUCCESS( | ||
| 1232 | + CollectCmpNodes(cur_node, next_node, compute_nodes, origin_bro_nodes, graph, next_mul_input_nodes)); | ||
| 1233 | + GELOGD("Move partial broadcast at node[%s]: common_axes=%zu, compute_nodes=%zu, broadcasts=%zu.", | ||
| 1234 | + mul_input_node->GetName().c_str(), common_axises.size(), compute_nodes.size(), origin_bro_nodes.size()); | ||
| 1235 | + | ||
| 1236 | + is_changed = true; | ||
| 1237 | + GE_ASSERT_SUCCESS(RemoveBroadcasts(graph, bro_nodes_list, common_axises)); | ||
| 1238 | + GE_ASSERT_SUCCESS(InsertBroadcastNode(compute_nodes.back(), graph, common_axises)); | ||
| 1239 | + | ||
| 1240 | + std::vector<NodePtr> merged_nodes; | ||
| 1241 | + for (const auto &row : bro_nodes_list) { | ||
| 1242 | + merged_nodes.insert(merged_nodes.end(), row.begin(), row.end()); | ||
| 1243 | + } | ||
| 1244 | + merged_nodes.insert(merged_nodes.end(), compute_nodes.begin(), compute_nodes.end()); | ||
| 1245 | + GE_ASSERT_SUCCESS(UpdateOutputTensor(merged_nodes, common_axises)); | ||
| 1246 | + } | ||
| 1247 | + if (!next_mul_input_nodes.empty()) { | ||
| 1248 | + return JudgePartBackward(next_mul_input_nodes, is_changed, graph); | ||
| 1249 | + } | ||
| 1250 | + return SUCCESS; | ||
| 1251 | +} | ||
| 1252 | + | ||
| 1253 | +Status ProcessOriginalBackwardLogic(AscGraph &graph, bool &is_changed, std::set<NodePtr> &mul_input_nodes) { | ||
| 1254 | + std::vector<NodePtr> start_nodes; | ||
| 1255 | + GE_ASSERT_SUCCESS(CollectBackwardStartNodes(graph, start_nodes)); | ||
| 1256 | + RemoveDuplicates(start_nodes); | ||
| 1257 | + for (const auto &start_node : start_nodes) { | ||
| 1258 | + std::vector<NodePtr> peer_in_nodes; | ||
| 1259 | + GE_ASSERT_SUCCESS(CollectBackwardSatisfyStartNodes(start_node, peer_in_nodes)); | ||
| 1260 | + for (const auto &peer_in_node : peer_in_nodes) { | ||
| 1261 | + NodePtr cur_node = start_node; | ||
| 1262 | + NodePtr next_node = peer_in_node; | ||
| 1263 | + | ||
| 1264 | + std::vector<NodePtr> bro_nodes; | ||
| 1265 | + std::vector<NodePtr> compute_nodes; | ||
| 1266 | + NodePtr pre_bro_node = start_node; | ||
| 1267 | + GE_ASSERT_SUCCESS(CollectBroNodes(cur_node, next_node, bro_nodes)); | ||
| 1268 | + GE_ASSERT_SUCCESS(CollectCmpNodes(cur_node, next_node, compute_nodes, bro_nodes, graph, mul_input_nodes)); | ||
| 1269 | + | ||
| 1270 | + GELOGI("Move broadcast backward from node[%s]: broadcasts=%zu, computes=%zu.", peer_in_node->GetName().c_str(), | ||
| 1271 | + bro_nodes.size(), compute_nodes.size()); | ||
| 1272 | + | ||
| 1273 | + if (!bro_nodes.empty() && !compute_nodes.empty()) { | ||
| 1274 | + is_changed = true; | ||
| 1275 | + GE_ASSERT_SUCCESS(BroadcastBackwardReally(compute_nodes, bro_nodes, pre_bro_node)); | ||
| 1276 | + } | ||
| 1277 | + } | ||
| 1278 | + } | ||
| 1279 | + return SUCCESS; | ||
| 1280 | +} | ||
| 1281 | + | ||
| 1282 | +Status BroadcastBackward(AscGraph &graph) { | ||
| 1283 | + if (ScheduleUtils::HasComputeType(graph, af::ComputeType::kComputeCube)) { | ||
| 1284 | + GELOGI("graph %s fuse type is cube, don't backward broadcast.", graph.GetName().c_str()); | ||
| 1285 | + return SUCCESS; | ||
| 1286 | + } | ||
| 1287 | + | ||
| 1288 | + GE_ASSERT_SUCCESS(broadcast_backward_shared_split::SplitSharedBroadcastBranches(graph)); | ||
| 1289 | + GE_ASSERT_SUCCESS(broadcast_backward_shared_split::SplitSharedBroadcastConsumers(graph)); | ||
| 1290 | + | ||
| 1291 | + bool is_changed = false; | ||
| 1292 | + bool has_multi_ref_change = true; | ||
| 1293 | + while (has_multi_ref_change) { | ||
| 1294 | + has_multi_ref_change = false; | ||
| 1295 | + | ||
| 1296 | + std::set<NodePtr> mul_input_nodes; | ||
| 1297 | + GE_ASSERT_SUCCESS(ProcessOriginalBackwardLogic(graph, is_changed, mul_input_nodes)); | ||
| 1298 | + | ||
| 1299 | + if (!mul_input_nodes.empty()) { | ||
| 1300 | + GE_ASSERT_SUCCESS(JudgePartBackward(mul_input_nodes, is_changed, graph)); | ||
| 1301 | + } | ||
| 1302 | + | ||
| 1303 | + bool multi_ref_changed = false; | ||
| 1304 | + GE_ASSERT_SUCCESS(ProcessMultiRefBroadcastBackward(graph, multi_ref_changed)); | ||
| 1305 | + if (multi_ref_changed) { | ||
| 1306 | + is_changed = true; | ||
| 1307 | + has_multi_ref_change = true; | ||
| 1308 | + } | ||
| 1309 | + } | ||
| 1310 | + | ||
| 1311 | + if (is_changed) { | ||
| 1312 | + GE_ASSERT_SUCCESS(ScheduleUtils::TopologicalSorting(graph)); | ||
| 1313 | + } | ||
| 1314 | + return SUCCESS; | ||
| 1315 | +} | ||
| 1316 | +} // namespace | ||
| 1317 | + | ||
| 1318 | +Status BroadcastBackwardPass::RunPass(af::AscGraph &graph) { | ||
| 1319 | + GE_ASSERT_SUCCESS(ScheduleUtils::TopologicalSorting(graph)); | ||
| 1320 | + GE_ASSERT_SUCCESS(BroadcastBackward(graph)); | ||
| 1321 | + GELOGI("Graph %s completed BroadcastBackward successfully.", graph.GetName().c_str()); | ||
| 1322 | + return SUCCESS; | ||
| 1323 | +} | ||
| 1324 | +} // namespace optimize | ||
| @@ -0,0 +1,24 @@ | |||
| 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,6 +12,7 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | + | ||
| 15 | 16 | ||
| 16 | 17 | ||
| 17 | 18 | ||
| @@ -27,6 +28,8 @@ class PassRunnerV1 final : public BasePassRunner { | |||
| 27 | this->RegisterPass<PowEquivSubstitutionPass>(); | 28 | this->RegisterPass<PowEquivSubstitutionPass>(); |
| 28 | this->RegisterPass<BroadcastConstToStorePass>(); | 29 | this->RegisterPass<BroadcastConstToStorePass>(); |
| 29 | this->RegisterPass<ScalarTo1DTensorPass>(); | 30 | this->RegisterPass<ScalarTo1DTensorPass>(); |
| 31 | + // The sched/tensor axes must be complete before moving Broadcasts; scalar Broadcast cleanup runs afterward. | ||
| 32 | + this->RegisterPass<BroadcastBackwardPass>(); | ||
| 30 | this->RegisterPass<ScalarBroadcastOptimizationPass>(); | 33 | this->RegisterPass<ScalarBroadcastOptimizationPass>(); |
| 31 | this->RegisterPass<MaskedFillInputReorderPass>(); | 34 | this->RegisterPass<MaskedFillInputReorderPass>(); |
| 32 | this->RegisterPass<ExpandDimsForAllReducePass>(); | 35 | this->RegisterPass<ExpandDimsForAllReducePass>(); |
| @@ -0,0 +1,12 @@ | |||
| 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}) | ||
| @@ -0,0 +1,99 @@ | |||
| 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 | + | ||
| @@ -0,0 +1,378 @@ | |||
| 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 | + | ||
| @@ -1894,7 +1894,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) { | |||
| 1894 | auto impl_grp_0_brc4 = impl_graphs[0].FindNode("brc4"); | 1894 | auto impl_grp_0_brc4 = impl_graphs[0].FindNode("brc4"); |
| 1895 | EXPECT_NE(impl_grp_0_brc4, nullptr); | 1895 | EXPECT_NE(impl_grp_0_brc4, nullptr); |
| 1896 | EXPECT_EQ(impl_grp_0_brc4->GetAllInDataAnchorsSize(), 1); | 1896 | EXPECT_EQ(impl_grp_0_brc4->GetAllInDataAnchorsSize(), 1); |
| 1897 | - EXPECT_EQ(impl_grp_0_brc4->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0"); | 1897 | + EXPECT_EQ(impl_grp_0_brc4->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0"); |
| 1898 | 1898 | ||
| 1899 | EXPECT_EQ(impl_graphs[1].FindNode("brc0"), nullptr); | 1899 | EXPECT_EQ(impl_graphs[1].FindNode("brc0"), nullptr); |
| 1900 | EXPECT_EQ(impl_graphs[1].FindNode("brc1"), nullptr); | 1900 | EXPECT_EQ(impl_graphs[1].FindNode("brc1"), nullptr); |
| @@ -1903,7 +1903,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) { | |||
| 1903 | auto impl_grp_1_brc3 = impl_graphs[1].FindNode("brc3"); | 1903 | auto impl_grp_1_brc3 = impl_graphs[1].FindNode("brc3"); |
| 1904 | EXPECT_NE(impl_grp_1_brc3, nullptr); | 1904 | EXPECT_NE(impl_grp_1_brc3, nullptr); |
| 1905 | EXPECT_EQ(impl_grp_1_brc3->GetAllInDataAnchorsSize(), 1); | 1905 | EXPECT_EQ(impl_grp_1_brc3->GetAllInDataAnchorsSize(), 1); |
| 1906 | - EXPECT_EQ(impl_grp_1_brc3->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0"); | 1906 | + EXPECT_EQ(impl_grp_1_brc3->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0"); |
| 1907 | 1907 | ||
| 1908 | EXPECT_EQ(impl_graphs[2].FindNode("brc0"), nullptr); | 1908 | EXPECT_EQ(impl_graphs[2].FindNode("brc0"), nullptr); |
| 1909 | EXPECT_EQ(impl_graphs[2].FindNode("brc3"), nullptr); | 1909 | EXPECT_EQ(impl_graphs[2].FindNode("brc3"), nullptr); |
| @@ -1911,7 +1911,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) { | |||
| 1911 | auto impl_grp_2_brc2 = impl_graphs[2].FindNode("brc2"); | 1911 | auto impl_grp_2_brc2 = impl_graphs[2].FindNode("brc2"); |
| 1912 | EXPECT_NE(impl_grp_2_brc2, nullptr); | 1912 | EXPECT_NE(impl_grp_2_brc2, nullptr); |
| 1913 | EXPECT_EQ(impl_grp_2_brc2->GetAllInDataAnchorsSize(), 1); | 1913 | EXPECT_EQ(impl_grp_2_brc2->GetAllInDataAnchorsSize(), 1); |
| 1914 | - EXPECT_EQ(impl_grp_2_brc2->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0"); | 1914 | + EXPECT_EQ(impl_grp_2_brc2->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0"); |
| 1915 | 1915 | ||
| 1916 | EXPECT_EQ(impl_graphs[3].FindNode("brc0"), nullptr); | 1916 | EXPECT_EQ(impl_graphs[3].FindNode("brc0"), nullptr); |
| 1917 | EXPECT_EQ(impl_graphs[3].FindNode("brc2"), nullptr); | 1917 | EXPECT_EQ(impl_graphs[3].FindNode("brc2"), nullptr); |
| @@ -1920,7 +1920,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) { | |||
| 1920 | auto impl_grp_3_brc1 = impl_graphs[3].FindNode("brc1"); | 1920 | auto impl_grp_3_brc1 = impl_graphs[3].FindNode("brc1"); |
| 1921 | EXPECT_NE(impl_grp_3_brc1, nullptr); | 1921 | EXPECT_NE(impl_grp_3_brc1, nullptr); |
| 1922 | EXPECT_EQ(impl_grp_3_brc1->GetAllInDataAnchorsSize(), 1); | 1922 | EXPECT_EQ(impl_grp_3_brc1->GetAllInDataAnchorsSize(), 1); |
| 1923 | - EXPECT_EQ(impl_grp_3_brc1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0"); | 1923 | + EXPECT_EQ(impl_grp_3_brc1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "add0"); |
| 1924 | 1924 | ||
| 1925 | EXPECT_EQ(impl_graphs[4].FindNode("brc1"), nullptr); | 1925 | EXPECT_EQ(impl_graphs[4].FindNode("brc1"), nullptr); |
| 1926 | EXPECT_EQ(impl_graphs[4].FindNode("brc2"), nullptr); | 1926 | EXPECT_EQ(impl_graphs[4].FindNode("brc2"), nullptr); |
| @@ -1929,7 +1929,7 @@ TEST_F(OptimizerSt, MultiBroadcastCancellation_All_One) { | |||
| 1929 | auto impl_grp_4_exp0 = impl_graphs[4].FindNode("exp0"); | 1929 | auto impl_grp_4_exp0 = impl_graphs[4].FindNode("exp0"); |
| 1930 | EXPECT_NE(impl_grp_4_exp0, nullptr); | 1930 | EXPECT_NE(impl_grp_4_exp0, nullptr); |
| 1931 | EXPECT_EQ(impl_grp_4_exp0->GetAllInDataAnchorsSize(), 1); | 1931 | EXPECT_EQ(impl_grp_4_exp0->GetAllInDataAnchorsSize(), 1); |
| 1932 | - EXPECT_EQ(impl_grp_4_exp0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc0"); | 1932 | + EXPECT_EQ(impl_grp_4_exp0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0"); |
| 1933 | } | 1933 | } |
| 1934 | 1934 | ||
| 1935 | TEST_F(OptimizerSt, ScalarBroadcastOptimization_Two_Scalar) { | 1935 | TEST_F(OptimizerSt, ScalarBroadcastOptimization_Two_Scalar) { |
| @@ -1965,10 +1965,10 @@ TEST_F(OptimizerSt, ScalarBroadcastOptimization_Two_Scalar) { | |||
| 1965 | EXPECT_EQ(impl_graphs.size(), 3); | 1965 | EXPECT_EQ(impl_graphs.size(), 3); |
| 1966 | auto impl_graph0 = af::AscGraphUtils::GetComputeGraph(impl_graphs[0]); | 1966 | auto impl_graph0 = af::AscGraphUtils::GetComputeGraph(impl_graphs[0]); |
| 1967 | EXPECT_EQ(impl_graph0->GetAllNodesSize(), 8); | 1967 | EXPECT_EQ(impl_graph0->GetAllNodesSize(), 8); |
| 1968 | - EXPECT_EQ(impl_graph0->FindNode("brc1"), nullptr); | 1968 | + EXPECT_NE(impl_graph0->FindNode("brc1"), nullptr); |
| 1969 | EXPECT_EQ(impl_graph0->FindNode("brc2"), nullptr); | 1969 | EXPECT_EQ(impl_graph0->FindNode("brc2"), nullptr); |
| 1970 | EXPECT_EQ(impl_graph0->FindNode("brc3"), nullptr); | 1970 | EXPECT_EQ(impl_graph0->FindNode("brc3"), nullptr); |
| 1971 | - EXPECT_NE(impl_graph0->FindNode("brc4"), nullptr); | 1971 | + EXPECT_EQ(impl_graph0->FindNode("brc4"), nullptr); |
| 1972 | EXPECT_EQ(impl_graph0->FindNode("brc5"), nullptr); | 1972 | EXPECT_EQ(impl_graph0->FindNode("brc5"), nullptr); |
| 1973 | EXPECT_EQ(impl_graph0->FindNode("brc6"), nullptr); | 1973 | EXPECT_EQ(impl_graph0->FindNode("brc6"), nullptr); |
| 1974 | } | 1974 | } |
| @@ -1999,23 +1999,23 @@ TEST_F(OptimizerSt, RemoveRedundantBroadcast) { | |||
| 1999 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0].size(), 1UL); | 1999 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0].size(), 1UL); |
| 2000 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL); | 2000 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL); |
| 2001 | auto impl_graphs = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs; | 2001 | auto impl_graphs = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs; |
| 2002 | - EXPECT_EQ(impl_graphs.size(), 3); | 2002 | + EXPECT_EQ(impl_graphs.size(), 4); |
| 2003 | - // check don't remove brc | 2003 | + // consumer split creates clone for exp1; common-axis backward removes brc0/brc1 from add0's chain |
| 2004 | auto impl_grp_0_exp1 = impl_graphs[0].FindNode("exp1"); | 2004 | auto impl_grp_0_exp1 = impl_graphs[0].FindNode("exp1"); |
| 2005 | EXPECT_NE(impl_grp_0_exp1, nullptr); | 2005 | EXPECT_NE(impl_grp_0_exp1, nullptr); |
| 2006 | EXPECT_EQ(impl_grp_0_exp1->GetAllInDataAnchorsSize(), 1); | 2006 | EXPECT_EQ(impl_grp_0_exp1->GetAllInDataAnchorsSize(), 1); |
| 2007 | - EXPECT_EQ(impl_grp_0_exp1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc0"); | 2007 | + EXPECT_EQ(impl_grp_0_exp1->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), |
| 2008 | + "brc0_consumer_split_1"); | ||
| 2008 | 2009 | ||
| 2009 | auto impol_grp_0_add0 = impl_graphs[0].FindNode("add0"); | 2010 | auto impol_grp_0_add0 = impl_graphs[0].FindNode("add0"); |
| 2010 | EXPECT_NE(impol_grp_0_add0, nullptr); | 2011 | EXPECT_NE(impol_grp_0_add0, nullptr); |
| 2011 | EXPECT_EQ(impol_grp_0_add0->GetAllInDataAnchorsSize(), 2); | 2012 | EXPECT_EQ(impol_grp_0_add0->GetAllInDataAnchorsSize(), 2); |
| 2012 | - EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc0"); | 2013 | + EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "abs0"); |
| 2013 | - EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "brc1"); | 2014 | + EXPECT_EQ(impol_grp_0_add0->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "exp0"); |
| 2014 | 2015 | ||
| 2015 | - EXPECT_NE(impl_graphs[0].FindNode("brc0"), nullptr); | 2016 | + EXPECT_NE(impl_graphs[0].FindNode("brc0_consumer_split_1"), nullptr); |
| 2016 | - EXPECT_NE(impl_graphs[0].FindNode("brc1"), nullptr); | ||
| 2017 | 2017 | ||
| 2018 | - // check remove brc | 2018 | + // check remove brc in unaligned template |
| 2019 | auto impl_grp_1_exp1 = impl_graphs[1].FindNode("exp1"); | 2019 | auto impl_grp_1_exp1 = impl_graphs[1].FindNode("exp1"); |
| 2020 | EXPECT_NE(impl_grp_1_exp1, nullptr); | 2020 | EXPECT_NE(impl_grp_1_exp1, nullptr); |
| 2021 | EXPECT_EQ(impl_grp_1_exp1->GetAllInDataAnchorsSize(), 1); | 2021 | EXPECT_EQ(impl_grp_1_exp1->GetAllInDataAnchorsSize(), 1); |
| @@ -2260,20 +2260,9 @@ TEST_F(OptimizerSt, BufQueAllocator_RemovePad_MemUnique) { | |||
| 2260 | broadcast1.y.dtype = af::DataType::DT_FLOAT; | 2260 | broadcast1.y.dtype = af::DataType::DT_FLOAT; |
| 2261 | broadcast1.attr.api.unit = ComputeUnit::kUnitVector; | 2261 | broadcast1.attr.api.unit = ComputeUnit::kUnitVector; |
| 2262 | 2262 | ||
| 2263 | - af::ascir_op::Abs abs0("abs0"); | ||
| 2264 | - abs0.x = broadcast1.y; | ||
| 2265 | - abs0.attr.api.compute_type = ComputeType::kComputeElewise; | ||
| 2266 | - abs0.attr.api.type = af::ApiType::kAPITypeCompute; | ||
| 2267 | - abs0.attr.sched.axis = {z0.id, z1.id}; | ||
| 2268 | - *abs0.y.axis = {z0.id, z1.id}; | ||
| 2269 | - *abs0.y.repeats = {s0, s1}; | ||
| 2270 | - *abs0.y.strides = {s1, One}; | ||
| 2271 | - abs0.y.dtype = af::DataType::DT_FLOAT; | ||
| 2272 | - abs0.attr.api.unit = ComputeUnit::kUnitVector; | ||
| 2273 | - | ||
| 2274 | af::ascir_op::Add add0("add0"); | 2263 | af::ascir_op::Add add0("add0"); |
| 2275 | add0.x1 = load0.y; | 2264 | add0.x1 = load0.y; |
| 2276 | - add0.x2 = abs0.y; | 2265 | + add0.x2 = broadcast1.y; |
| 2277 | add0.attr.api.compute_type = ComputeType::kComputeElewise; | 2266 | add0.attr.api.compute_type = ComputeType::kComputeElewise; |
| 2278 | add0.attr.api.type = af::ApiType::kAPITypeCompute; | 2267 | add0.attr.api.type = af::ApiType::kAPITypeCompute; |
| 2279 | add0.attr.sched.axis = {z0.id, z1.id}; | 2268 | add0.attr.sched.axis = {z0.id, z1.id}; |
| @@ -2337,23 +2326,20 @@ TEST_F(OptimizerSt, BufQueAllocator_RemovePad_MemUnique) { | |||
| 2337 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL); | 2326 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.size(), 1UL); |
| 2338 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs.size(), 3UL); | 2327 | EXPECT_EQ(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs.size(), 3UL); |
| 2339 | 2328 | ||
| 2340 | - auto impl_graph2 = af::AscGraphUtils::GetComputeGraph( | 2329 | + auto impl_graph1 = af::AscGraphUtils::GetComputeGraph( |
| 2341 | - fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[2]); | 2330 | + fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[1]); |
| 2342 | - EXPECT_EQ(impl_graph2->GetAllNodesSize(), 13); | 2331 | + EXPECT_EQ(impl_graph1->GetAllNodesSize(), 12); |
| 2343 | - EXPECT_NE(impl_graph2->FindNode("broadcast1"), nullptr); | 2332 | + EXPECT_NE(impl_graph1->FindNode("broadcast1"), nullptr); |
| 2344 | - EXPECT_NE(impl_graph2->FindNode("broadcast1_remove_pad_0"), nullptr); | 2333 | + EXPECT_NE(impl_graph1->FindNode("broadcast1_remove_pad_0"), nullptr); |
| 2345 | - EXPECT_NE(impl_graph2->FindNode("add0"), nullptr); | 2334 | + EXPECT_NE(impl_graph1->FindNode("add0"), nullptr); |
| 2346 | - EXPECT_NE(impl_graph2->FindNode("abs0"), nullptr); | 2335 | + const auto &impl_graph1_brc1 = std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("broadcast1")); |
| 2347 | - const auto &impl_graph2_brc1 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("broadcast1")); | 2336 | + const auto &impl_graph1_rpd = |
| 2348 | - const auto &impl_graph2_rpd = | 2337 | + std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("broadcast1_remove_pad_0")); |
| 2349 | - std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("broadcast1_remove_pad_0")); | 2338 | + const auto &impl_graph1_add0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("add0")); |
| 2350 | - const auto &impl_graph2_add0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("add0")); | 2339 | + const auto &impl_graph1_mul0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph1->FindNode("mul0")); |
| 2351 | - const auto &impl_graph2_abs0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("abs0")); | 2340 | + EXPECT_EQ(impl_graph1_brc1->outputs[0].attr.buf.id, 1); |
| 2352 | - const auto &impl_graph2_mul0 = std::dynamic_pointer_cast<af::AscNode>(impl_graph2->FindNode("mul0")); | 2341 | + EXPECT_EQ(impl_graph1_rpd->outputs[0].attr.buf.id, 2); |
| 2353 | - EXPECT_EQ(impl_graph2_brc1->outputs[0].attr.buf.id, 1); | 2342 | + EXPECT_EQ(impl_graph1_add0->outputs[0].attr.que.id, impl_graph1_mul0->outputs[0].attr.que.id); |
| 2354 | - EXPECT_EQ(impl_graph2_rpd->outputs[0].attr.buf.id, 2); | ||
| 2355 | - EXPECT_EQ(impl_graph2_abs0->outputs[0].attr.buf.id, 3); | ||
| 2356 | - EXPECT_EQ(impl_graph2_add0->outputs[0].attr.que.id, impl_graph2_mul0->outputs[0].attr.que.id); | ||
| 2357 | } | 2343 | } |
| 2358 | 2344 | ||
| 2359 | TEST_F(OptimizerSt, BufQueAllocator_Inplace) { | 2345 | TEST_F(OptimizerSt, BufQueAllocator_Inplace) { |
| @@ -0,0 +1,132 @@ | |||
| 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 | +} | ||
| @@ -0,0 +1,1884 @@ | |||
| 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 | + | ||
| 24 | + | ||
| 25 | +namespace { | ||
| 26 | +using af::AscGraph; | ||
| 27 | +using af::testing::AscGraphBuilder; | ||
| 28 | +using af::testing::Sym; | ||
| 29 | +using namespace broadcast_backward_test; | ||
| 30 | + | ||
| 31 | +class ScopedTestPlatform { | ||
| 32 | + public: | ||
| 33 | + explicit ScopedTestPlatform(const char *platform) { | ||
| 34 | + ge::PlatformContext::GetInstance().SetPlatform(platform); | ||
| 35 | + } | ||
| 36 | + | ||
| 37 | + ~ScopedTestPlatform() { | ||
| 38 | + ge::PlatformContext::GetInstance().Reset(); | ||
| 39 | + } | ||
| 40 | +}; | ||
| 41 | +} // namespace | ||
| 42 | + | ||
| 43 | +TEST(BroadcastBackwardPass, MovesSingleBroadcastChain) { | ||
| 44 | + auto graph = BuildUnaryGraph("Abs"); | ||
| 45 | + CompleteApiInfo(graph); | ||
| 46 | + const auto load_node = FindNode(graph, "load"); | ||
| 47 | + ASSERT_NE(load_node, nullptr); | ||
| 48 | + std::vector<af::Expression> expected_strides; | ||
| 49 | + ASSERT_EQ( | ||
| 50 | + optimize::ScheduleUtils::RecalculateStridesFromRepeats(load_node->outputs[0].attr.repeats, expected_strides), | ||
| 51 | + af::SUCCESS); | ||
| 52 | + | ||
| 53 | + optimize::BroadcastBackwardPass pass; | ||
| 54 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 55 | + EXPECT_TRUE(IsConnected(graph, "load", "compute")); | ||
| 56 | + EXPECT_TRUE(IsConnected(graph, "compute", "broadcast")); | ||
| 57 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 58 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load", "compute")); | ||
| 59 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "compute", "broadcast")); | ||
| 60 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "store")); | ||
| 61 | + const auto compute_node = FindNode(graph, "compute"); | ||
| 62 | + ASSERT_NE(compute_node, nullptr); | ||
| 63 | + ExpectStaticEq(compute_node->outputs[0].attr.strides, expected_strides); | ||
| 64 | +} | ||
| 65 | + | ||
| 66 | +TEST(BroadcastBackwardPass, MovesAllSupportedUnaryOperators) { | ||
| 67 | + for (const auto &op_name : std::vector<std::string>{"Abs", "Neg", "Exp", "Sqrt", "Rsqrt", "Relu", "Reciprocal", "Erf", | ||
| 68 | + "Sign", "Tanh", "Ln"}) { | ||
| 69 | + auto graph = BuildUnaryGraph(op_name); | ||
| 70 | + CompleteApiInfo(graph); | ||
| 71 | + | ||
| 72 | + optimize::BroadcastBackwardPass pass; | ||
| 73 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS) << op_name; | ||
| 74 | + EXPECT_TRUE(IsConnected(graph, "load", "compute")) << op_name; | ||
| 75 | + EXPECT_TRUE(IsConnected(graph, "compute", "broadcast")) << op_name; | ||
| 76 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")) << op_name; | ||
| 77 | + } | ||
| 78 | +} | ||
| 79 | + | ||
| 80 | +TEST(BroadcastBackwardPass, MovesMultipleComputeNodes) { | ||
| 81 | + auto graph = AscGraphBuilder("broadcast_backward_compute_chain") | ||
| 82 | + .Loops({Sym("s0"), Sym("s1")}) | ||
| 83 | + .Data("data", 0) | ||
| 84 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 85 | + .Broadcast("broadcast", "load", {Sym("s0"), Sym("s1")}) | ||
| 86 | + .Abs("abs", "broadcast") | ||
| 87 | + .Relu("relu", "abs") | ||
| 88 | + .Store("store", "relu") | ||
| 89 | + .Output("output", "store") | ||
| 90 | + .Build(); | ||
| 91 | + CompleteApiInfo(graph); | ||
| 92 | + optimize::BroadcastBackwardPass pass; | ||
| 93 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 94 | + EXPECT_TRUE(IsConnected(graph, "load", "abs")); | ||
| 95 | + EXPECT_TRUE(IsConnected(graph, "abs", "relu")); | ||
| 96 | + EXPECT_TRUE(IsConnected(graph, "relu", "broadcast")); | ||
| 97 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 98 | +} | ||
| 99 | + | ||
| 100 | +TEST(BroadcastBackwardPass, MovesFormerCastBarrierWithDtypeAwareBackward) { | ||
| 101 | + auto graph = AscGraphBuilder("broadcast_backward_cast_barrier") | ||
| 102 | + .Loops({Sym("s0"), Sym("s1")}) | ||
| 103 | + .Data("data", 0) | ||
| 104 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 105 | + .Broadcast("broadcast", "load", {Sym("s0"), Sym("s1")}) | ||
| 106 | + .Cast("cast", "broadcast", af::DT_FLOAT16) | ||
| 107 | + .Store("store", "cast") | ||
| 108 | + .Output("output", "store") | ||
| 109 | + .Build(); | ||
| 110 | + CompleteApiInfo(graph); | ||
| 111 | + optimize::BroadcastBackwardPass pass; | ||
| 112 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 113 | + EXPECT_TRUE(IsConnected(graph, "load", "cast")); | ||
| 114 | + EXPECT_TRUE(IsConnected(graph, "cast", "broadcast")); | ||
| 115 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 116 | +} | ||
| 117 | + | ||
| 118 | +TEST(BroadcastBackwardPass, SkipsReduceBarrier) { | ||
| 119 | + auto graph = AscGraphBuilder("broadcast_backward_reduce_barrier") | ||
| 120 | + .Loops({Sym("s0"), Sym("s1")}) | ||
| 121 | + .Data("data", 0) | ||
| 122 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 123 | + .Broadcast("broadcast", "load", {Sym("s0"), Sym("s1")}) | ||
| 124 | + .Sum("sum", "broadcast", {0U}) | ||
| 125 | + .Store("store", "sum") | ||
| 126 | + .Output("output", "store") | ||
| 127 | + .Build(); | ||
| 128 | + CompleteApiInfo(graph); | ||
| 129 | + | ||
| 130 | + optimize::BroadcastBackwardPass pass; | ||
| 131 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 132 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "sum")); | ||
| 133 | + EXPECT_TRUE(IsConnected(graph, "sum", "store")); | ||
| 134 | +} | ||
| 135 | + | ||
| 136 | +TEST(BroadcastBackwardPass, SkipsVectorizedLayout) { | ||
| 137 | + // Not supported by the restored repository BRC implementation. | ||
| 138 | + GTEST_SKIP(); | ||
| 139 | + auto graph = BuildUnaryGraph("Abs"); | ||
| 140 | + CompleteApiInfo(graph); | ||
| 141 | + auto load_node = FindNode(graph, "load"); | ||
| 142 | + ASSERT_NE(load_node, nullptr); | ||
| 143 | + load_node->outputs[0].attr.vectorized_axis.push_back(0U); | ||
| 144 | + | ||
| 145 | + optimize::BroadcastBackwardPass pass; | ||
| 146 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 147 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "compute")); | ||
| 148 | + EXPECT_TRUE(IsConnected(graph, "compute", "store")); | ||
| 149 | +} | ||
| 150 | + | ||
| 151 | +TEST(BroadcastBackwardPass, SkipsMultipleBroadcastConsumers) { | ||
| 152 | + // Not supported by the restored repository BRC implementation. | ||
| 153 | + GTEST_SKIP(); | ||
| 154 | + auto graph = AscGraphBuilder("broadcast_backward_multiple_consumers") | ||
| 155 | + .Loops({Sym("s0"), Sym("s1")}) | ||
| 156 | + .Data("data", 0) | ||
| 157 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 158 | + .Broadcast("broadcast", "load", {Sym("s0"), Sym("s1")}) | ||
| 159 | + .Abs("abs0", "broadcast") | ||
| 160 | + .Abs("abs1", "broadcast") | ||
| 161 | + .Store("store0", "abs0") | ||
| 162 | + .Store("store1", "abs1") | ||
| 163 | + .Output("output0", "store0") | ||
| 164 | + .Output("output1", "store1") | ||
| 165 | + .Build(); | ||
| 166 | + CompleteApiInfo(graph); | ||
| 167 | + | ||
| 168 | + optimize::BroadcastBackwardPass pass; | ||
| 169 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 170 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "abs0")); | ||
| 171 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "abs1")); | ||
| 172 | +} | ||
| 173 | + | ||
| 174 | +TEST(BroadcastBackwardPass, SkipsIncompleteBroadcastLayout) { | ||
| 175 | + // Not supported by the restored repository BRC implementation. | ||
| 176 | + GTEST_SKIP(); | ||
| 177 | + auto graph = BuildUnaryGraph("Abs"); | ||
| 178 | + CompleteApiInfo(graph); | ||
| 179 | + auto broadcast_node = FindNode(graph, "broadcast"); | ||
| 180 | + ASSERT_NE(broadcast_node, nullptr); | ||
| 181 | + broadcast_node->outputs[0].attr.strides.clear(); | ||
| 182 | + | ||
| 183 | + optimize::BroadcastBackwardPass pass; | ||
| 184 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 185 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "compute")); | ||
| 186 | + EXPECT_TRUE(IsConnected(graph, "compute", "store")); | ||
| 187 | +} | ||
| 188 | + | ||
| 189 | +TEST(BroadcastBackwardPass, SkipsSchedMismatch) { | ||
| 190 | + // Not supported by the restored repository BRC implementation. | ||
| 191 | + GTEST_SKIP(); | ||
| 192 | + auto graph = BuildUnaryGraph("Abs"); | ||
| 193 | + CompleteApiInfo(graph); | ||
| 194 | + auto compute_node = FindNode(graph, "compute"); | ||
| 195 | + ASSERT_NE(compute_node, nullptr); | ||
| 196 | + compute_node->attr.sched.axis.clear(); | ||
| 197 | + | ||
| 198 | + optimize::BroadcastBackwardPass pass; | ||
| 199 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 200 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "compute")); | ||
| 201 | + EXPECT_TRUE(IsConnected(graph, "compute", "store")); | ||
| 202 | +} | ||
| 203 | + | ||
| 204 | +TEST(BroadcastBackwardPass, SkipsControlEdge) { | ||
| 205 | + // Not supported by the restored repository BRC implementation. | ||
| 206 | + GTEST_SKIP(); | ||
| 207 | + auto graph = BuildUnaryGraph("Abs"); | ||
| 208 | + CompleteApiInfo(graph); | ||
| 209 | + auto load_node = FindNode(graph, "load"); | ||
| 210 | + auto compute_node = FindNode(graph, "compute"); | ||
| 211 | + ASSERT_NE(load_node, nullptr); | ||
| 212 | + ASSERT_NE(compute_node, nullptr); | ||
| 213 | + ASSERT_EQ(af::GraphUtils::AddEdge(load_node->GetOutControlAnchor(), compute_node->GetInControlAnchor()), af::SUCCESS); | ||
| 214 | + | ||
| 215 | + optimize::BroadcastBackwardPass pass; | ||
| 216 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 217 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "compute")); | ||
| 218 | + EXPECT_TRUE(IsConnected(graph, "compute", "store")); | ||
| 219 | +} | ||
| 220 | + | ||
| 221 | +TEST(BroadcastBackwardPass, SkipsScalarBroadcastSource) { | ||
| 222 | + auto graph = AscGraphBuilder("broadcast_backward_scalar") | ||
| 223 | + .Loops({Sym("s0"), Sym("s1")}) | ||
| 224 | + .Scalar("scalar", "1.0") | ||
| 225 | + .Broadcast("broadcast0", "scalar", {Sym("s0"), Sym("s1")}) | ||
| 226 | + .Broadcast("broadcast1", "broadcast0", {Sym("s0"), Sym("s1")}) | ||
| 227 | + .Abs("abs", "broadcast1") | ||
| 228 | + .Store("store", "abs") | ||
| 229 | + .Output("output", "store") | ||
| 230 | + .Build(); | ||
| 231 | + CompleteApiInfo(graph); | ||
| 232 | + | ||
| 233 | + optimize::BroadcastBackwardPass pass; | ||
| 234 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 235 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "broadcast1")); | ||
| 236 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "abs")); | ||
| 237 | + EXPECT_TRUE(IsConnected(graph, "abs", "store")); | ||
| 238 | +} | ||
| 239 | + | ||
| 240 | +TEST(BroadcastBackwardPass, MovesSupportedScalarBroadcastBranches) { | ||
| 241 | + // Not supported by the restored repository BRC implementation. | ||
| 242 | + GTEST_SKIP(); | ||
| 243 | + const auto s0 = Sym("s0"); | ||
| 244 | + const auto s1 = Sym("s1"); | ||
| 245 | + auto graph = AscGraphBuilder("broadcast_backward_scalar_supported") | ||
| 246 | + .Loops({s0, s1}) | ||
| 247 | + .Scalar("scalar0", "1.0") | ||
| 248 | + .Scalar("scalar1", "2.0") | ||
| 249 | + .Broadcast("broadcast0", "scalar0", {s0, s1}) | ||
| 250 | + .Broadcast("broadcast1", "scalar1", {s0, s1}) | ||
| 251 | + .Sub("compute", "broadcast0", "broadcast1") | ||
| 252 | + .Store("store", "compute") | ||
| 253 | + .Output("output", "store") | ||
| 254 | + .Build(); | ||
| 255 | + CompleteApiInfo(graph); | ||
| 256 | + | ||
| 257 | + optimize::BroadcastBackwardPass pass; | ||
| 258 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 259 | + EXPECT_TRUE(IsConnected(graph, "scalar0", "compute")); | ||
| 260 | + EXPECT_TRUE(IsConnected(graph, "scalar1", "compute")); | ||
| 261 | + EXPECT_TRUE(IsConnected(graph, "compute", "compute_broadcast_backward_common")); | ||
| 262 | + EXPECT_TRUE(IsConnected(graph, "compute_broadcast_backward_common", "store")); | ||
| 263 | + EXPECT_FALSE(HasNode(graph, "broadcast0")); | ||
| 264 | + EXPECT_FALSE(HasNode(graph, "broadcast1")); | ||
| 265 | + const auto compute = FindNode(graph, "compute"); | ||
| 266 | + ASSERT_NE(compute, nullptr); | ||
| 267 | + ExpectStaticEq(compute->outputs[0].attr.repeats, {af::sym::kSymbolOne, af::sym::kSymbolOne}); | ||
| 268 | + ExpectStaticEq(compute->outputs[0].attr.strides, {af::sym::kSymbolZero, af::sym::kSymbolZero}); | ||
| 269 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "scalar0", "compute")); | ||
| 270 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "scalar1", "compute")); | ||
| 271 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "compute", "compute_broadcast_backward_common")); | ||
| 272 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "compute_broadcast_backward_common", "store")); | ||
| 273 | +} | ||
| 274 | + | ||
| 275 | +TEST(BroadcastBackwardPass, SkipsUnsupportedAllScalarInputs) { | ||
| 276 | + const auto s0 = Sym("s0"); | ||
| 277 | + const auto s1 = Sym("s1"); | ||
| 278 | + auto graph = AscGraphBuilder("broadcast_backward_scalar_unsupported") | ||
| 279 | + .Loops({s0, s1}) | ||
| 280 | + .Scalar("scalar0", "1.0") | ||
| 281 | + .Scalar("scalar1", "2.0") | ||
| 282 | + .Broadcast("broadcast0", "scalar0", {s0, s1}) | ||
| 283 | + .Broadcast("broadcast1", "scalar1", {s0, s1}) | ||
| 284 | + .Add("compute", "broadcast0", "broadcast1") | ||
| 285 | + .Store("store", "compute") | ||
| 286 | + .Output("output", "store") | ||
| 287 | + .Build(); | ||
| 288 | + CompleteApiInfo(graph); | ||
| 289 | + | ||
| 290 | + optimize::BroadcastBackwardPass pass; | ||
| 291 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 292 | + EXPECT_TRUE(IsConnected(graph, "scalar0", "broadcast0")); | ||
| 293 | + EXPECT_TRUE(IsConnected(graph, "scalar1", "broadcast1")); | ||
| 294 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "compute")); | ||
| 295 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "compute")); | ||
| 296 | + EXPECT_FALSE(HasNode(graph, "compute_broadcast_backward_common")); | ||
| 297 | +} | ||
| 298 | + | ||
| 299 | +TEST(BroadcastBackwardPass, SkipsScalarBranchWithResidualBroadcastAxis) { | ||
| 300 | + const auto s0 = Sym("s0"); | ||
| 301 | + const auto s1 = Sym("s1"); | ||
| 302 | + auto graph = AscGraphBuilder("broadcast_backward_scalar_mixed") | ||
| 303 | + .Loops({s0, s1}) | ||
| 304 | + .Data("data", 0) | ||
| 305 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 306 | + .Scalar("scalar", "1.0") | ||
| 307 | + .Broadcast("broadcast0", "load", {s0, s1}) | ||
| 308 | + .Broadcast("broadcast1", "scalar", {s0, s1}) | ||
| 309 | + .Add("compute", "broadcast0", "broadcast1") | ||
| 310 | + .Store("store", "compute") | ||
| 311 | + .Output("output", "store") | ||
| 312 | + .Build(); | ||
| 313 | + CompleteApiInfo(graph); | ||
| 314 | + | ||
| 315 | + optimize::BroadcastBackwardPass pass; | ||
| 316 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 317 | + EXPECT_TRUE(IsConnected(graph, "load", "broadcast0")); | ||
| 318 | + EXPECT_TRUE(IsConnected(graph, "scalar", "broadcast1")); | ||
| 319 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "compute")); | ||
| 320 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "compute")); | ||
| 321 | + EXPECT_TRUE(IsConnected(graph, "compute", "store")); | ||
| 322 | + EXPECT_FALSE(HasNode(graph, "compute_broadcast_backward_common")); | ||
| 323 | +} | ||
| 324 | + | ||
| 325 | +TEST(BroadcastBackwardPass, MovesSupportedScalarSameConsumerMultiReference) { | ||
| 326 | + // Not supported by the restored repository BRC implementation. | ||
| 327 | + GTEST_SKIP(); | ||
| 328 | + const auto s0 = Sym("s0"); | ||
| 329 | + const auto s1 = Sym("s1"); | ||
| 330 | + auto graph = AscGraphBuilder("broadcast_backward_scalar_same_consumer") | ||
| 331 | + .Loops({s0, s1}) | ||
| 332 | + .Scalar("scalar", "1.0") | ||
| 333 | + .Broadcast("broadcast", "scalar", {s0, s1}) | ||
| 334 | + .Sub("compute", "broadcast", "broadcast") | ||
| 335 | + .Store("store", "compute") | ||
| 336 | + .Output("output", "store") | ||
| 337 | + .Build(); | ||
| 338 | + CompleteApiInfo(graph); | ||
| 339 | + | ||
| 340 | + optimize::BroadcastBackwardPass pass; | ||
| 341 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 342 | + EXPECT_TRUE(IsConnected(graph, "scalar", "compute")); | ||
| 343 | + EXPECT_TRUE(IsConnected(graph, "compute", "broadcast")); | ||
| 344 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 345 | + const auto compute = FindNode(graph, "compute"); | ||
| 346 | + ASSERT_NE(compute, nullptr); | ||
| 347 | + ExpectStaticEq(compute->outputs[0].attr.repeats, {af::sym::kSymbolOne, af::sym::kSymbolOne}); | ||
| 348 | + ExpectStaticEq(compute->outputs[0].attr.strides, {af::sym::kSymbolZero, af::sym::kSymbolZero}); | ||
| 349 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "compute", "broadcast")); | ||
| 350 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "store")); | ||
| 351 | +} | ||
| 352 | + | ||
| 353 | +TEST(BroadcastBackwardPass, SkipsUnsupportedScalarSameConsumerMultiReference) { | ||
| 354 | + const auto s0 = Sym("s0"); | ||
| 355 | + const auto s1 = Sym("s1"); | ||
| 356 | + auto graph = AscGraphBuilder("broadcast_backward_scalar_same_consumer_unsupported") | ||
| 357 | + .Loops({s0, s1}) | ||
| 358 | + .Scalar("scalar", "1.0") | ||
| 359 | + .Broadcast("broadcast", "scalar", {s0, s1}) | ||
| 360 | + .Add("compute", "broadcast", "broadcast") | ||
| 361 | + .Store("store", "compute") | ||
| 362 | + .Output("output", "store") | ||
| 363 | + .Build(); | ||
| 364 | + CompleteApiInfo(graph); | ||
| 365 | + const auto compute = FindNode(graph, "compute"); | ||
| 366 | + ASSERT_NE(compute, nullptr); | ||
| 367 | + EXPECT_TRUE(IsConnected(graph, "scalar", "broadcast")); | ||
| 368 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "compute")); | ||
| 369 | + | ||
| 370 | + optimize::BroadcastBackwardPass pass; | ||
| 371 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 372 | + EXPECT_TRUE(HasNode(graph, "scalar")); | ||
| 373 | + EXPECT_TRUE(HasNode(graph, "broadcast")); | ||
| 374 | + EXPECT_TRUE(HasNode(graph, "compute")); | ||
| 375 | + EXPECT_TRUE(IsConnected(graph, "scalar", "broadcast")); | ||
| 376 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "compute")); | ||
| 377 | + EXPECT_TRUE(IsConnected(graph, "compute", "store")); | ||
| 378 | +} | ||
| 379 | + | ||
| 380 | +TEST(BroadcastBackwardPass, MovesSupportedMultiLevelScalarBroadcastChain) { | ||
| 381 | + // Not supported by the restored repository BRC implementation. | ||
| 382 | + GTEST_SKIP(); | ||
| 383 | + const auto s0 = Sym("s0"); | ||
| 384 | + const auto s1 = Sym("s1"); | ||
| 385 | + auto graph = AscGraphBuilder("broadcast_backward_scalar_multi_level") | ||
| 386 | + .Loops({s0, s1}) | ||
| 387 | + .Scalar("scalar", "1.0") | ||
| 388 | + .Broadcast("broadcast0", "scalar", {0}) | ||
| 389 | + .Broadcast("broadcast1", "broadcast0", {1}) | ||
| 390 | + .Sub("compute", "broadcast1", "broadcast1") | ||
| 391 | + .Store("store", "compute") | ||
| 392 | + .Output("output", "store") | ||
| 393 | + .Build(); | ||
| 394 | + CompleteApiInfo(graph); | ||
| 395 | + | ||
| 396 | + optimize::BroadcastBackwardPass pass; | ||
| 397 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 398 | + EXPECT_TRUE(IsConnected(graph, "scalar", "compute")); | ||
| 399 | + EXPECT_TRUE(IsConnected(graph, "compute", "broadcast0")); | ||
| 400 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "broadcast1")); | ||
| 401 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "store")); | ||
| 402 | + const auto compute = FindNode(graph, "compute"); | ||
| 403 | + ASSERT_NE(compute, nullptr); | ||
| 404 | + ExpectStaticEq(compute->outputs[0].attr.strides, {af::sym::kSymbolZero, af::sym::kSymbolZero}); | ||
| 405 | +} | ||
| 406 | + | ||
| 407 | +TEST(BroadcastBackwardPass, ScalarBranchAxisRequiresZeroInputStride) { | ||
| 408 | + GTEST_SKIP(); | ||
| 409 | + const auto s0 = Sym("s0"); | ||
| 410 | + const auto s1 = Sym("s1"); | ||
| 411 | + auto graph = AscGraphBuilder("broadcast_backward_scalar_stride_guard") | ||
| 412 | + .Loops({s0, s1}) | ||
| 413 | + .Scalar("scalar0", "1.0") | ||
| 414 | + .Scalar("scalar1", "2.0") | ||
| 415 | + .Broadcast("broadcast0", "scalar0", {0}) | ||
| 416 | + .Broadcast("broadcast1", "scalar1", {0}) | ||
| 417 | + .Sub("compute", "broadcast0", "broadcast1") | ||
| 418 | + .Store("store", "compute") | ||
| 419 | + .Output("output", "store") | ||
| 420 | + .Build(); | ||
| 421 | + CompleteApiInfo(graph); | ||
| 422 | + const auto scalar1 = FindNode(graph, "scalar1"); | ||
| 423 | + const auto broadcast0 = FindNode(graph, "broadcast0"); | ||
| 424 | + const auto compute = FindNode(graph, "compute"); | ||
| 425 | + ASSERT_NE(scalar1, nullptr); | ||
| 426 | + ASSERT_NE(broadcast0, nullptr); | ||
| 427 | + ASSERT_NE(compute, nullptr); | ||
| 428 | + scalar1->outputs[0].attr.strides[0] = af::sym::kSymbolOne; | ||
| 429 | +} | ||
| 430 | + | ||
| 431 | +TEST(BroadcastBackwardPass, MovesIdenticalMultiInputBroadcastChains) { | ||
| 432 | + // Not supported by the restored repository BRC implementation. | ||
| 433 | + GTEST_SKIP(); | ||
| 434 | + auto graph = BuildBinaryGraph("Add"); | ||
| 435 | + CompleteApiInfo(graph); | ||
| 436 | + | ||
| 437 | + optimize::BroadcastBackwardPass pass; | ||
| 438 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 439 | + ExpectBinaryBroadcastMove(graph); | ||
| 440 | +} | ||
| 441 | + | ||
| 442 | +TEST(BroadcastBackwardPass, MovesMultiNodeBroadcastAndComputeChains) { | ||
| 443 | + GTEST_SKIP(); | ||
| 444 | + const auto s0 = Sym("s0"); | ||
| 445 | + const auto s1 = Sym("s1"); | ||
| 446 | + const auto s2 = Sym("s2"); | ||
| 447 | + const std::vector<af::Expression> compact_repeats = {s0, af::sym::kSymbolOne, af::sym::kSymbolOne}; | ||
| 448 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolOne, af::sym::kSymbolZero, af::sym::kSymbolZero}; | ||
| 449 | + auto graph = AscGraphBuilder("broadcast_backward_multi_node_chains") | ||
| 450 | + .Loops({s0, s1, s2}) | ||
| 451 | + .Data("data0", 0) | ||
| 452 | + .Data("data1", 1) | ||
| 453 | + .Load("load0", "data0", compact_repeats, compact_strides) | ||
| 454 | + .Load("load1", "data1", compact_repeats, compact_strides) | ||
| 455 | + .Broadcast("broadcast00", "load0", {s0, s1, af::sym::kSymbolOne}) | ||
| 456 | + .Broadcast("broadcast01", "broadcast00", {s0, s1, s2}) | ||
| 457 | + .Broadcast("broadcast10", "load1", {s0, s1, af::sym::kSymbolOne}) | ||
| 458 | + .Broadcast("broadcast11", "broadcast10", {s0, s1, s2}) | ||
| 459 | + .Add("merge", "broadcast01", "broadcast11") | ||
| 460 | + .Abs("abs", "merge") | ||
| 461 | + .Relu("relu", "abs") | ||
| 462 | + .Store("store", "relu") | ||
| 463 | + .Output("output", "store") | ||
| 464 | + .Build(); | ||
| 465 | + CompleteApiInfo(graph); | ||
| 466 | + | ||
| 467 | + optimize::BroadcastBackwardPass pass; | ||
| 468 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 469 | + EXPECT_TRUE(IsConnected(graph, "load0", "merge")); | ||
| 470 | + EXPECT_TRUE(IsConnected(graph, "load1", "merge")); | ||
| 471 | + EXPECT_TRUE(IsConnected(graph, "merge", "abs")); | ||
| 472 | + EXPECT_TRUE(IsConnected(graph, "abs", "relu")); | ||
| 473 | + EXPECT_TRUE(IsConnected(graph, "relu", "broadcast00")); | ||
| 474 | + EXPECT_TRUE(IsConnected(graph, "broadcast00", "broadcast01")); | ||
| 475 | + EXPECT_TRUE(IsConnected(graph, "broadcast01", "store")); | ||
| 476 | + EXPECT_FALSE(HasNode(graph, "broadcast10")); | ||
| 477 | + EXPECT_FALSE(HasNode(graph, "broadcast11")); | ||
| 478 | +} | ||
| 479 | + | ||
| 480 | +TEST(BroadcastBackwardPass, MovesAllSupportedBinaryOperators) { | ||
| 481 | + // Not supported by the restored repository BRC implementation. | ||
| 482 | + GTEST_SKIP(); | ||
| 483 | + for (const auto &op_name : std::vector<std::string>{"Add", "Sub", "Mul", "Div", "Minimum", "Maximum"}) { | ||
| 484 | + auto graph = BuildBinaryGraph(op_name); | ||
| 485 | + CompleteApiInfo(graph); | ||
| 486 | + | ||
| 487 | + optimize::BroadcastBackwardPass pass; | ||
| 488 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS) << op_name; | ||
| 489 | + EXPECT_TRUE(IsConnected(graph, "load0", "compute")) << op_name; | ||
| 490 | + EXPECT_TRUE(IsConnected(graph, "load1", "compute")) << op_name; | ||
| 491 | + EXPECT_TRUE(IsConnected(graph, "compute", "broadcast0")) << op_name; | ||
| 492 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load0", "compute")) << op_name; | ||
| 493 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load1", "compute")) << op_name; | ||
| 494 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "compute", "broadcast0")) << op_name; | ||
| 495 | + EXPECT_FALSE(HasNode(graph, "broadcast1")) << op_name; | ||
| 496 | + } | ||
| 497 | +} | ||
| 498 | + | ||
| 499 | +TEST(BroadcastBackwardPass, MovesBeforeMultiInputBarrier) { | ||
| 500 | + auto graph = AscGraphBuilder("broadcast_backward_multi_input_barrier") | ||
| 501 | + .Loops({Sym("s0"), Sym("s1")}) | ||
| 502 | + .Data("data0", 0) | ||
| 503 | + .Data("data1", 1) | ||
| 504 | + .Load("load0", "data0", kCompactRepeats, kCompactStrides) | ||
| 505 | + .Load("load1", "data1") | ||
| 506 | + .Broadcast("broadcast", "load0", {Sym("s0"), Sym("s1")}) | ||
| 507 | + .Abs("abs", "broadcast") | ||
| 508 | + .Add("add", "load1", "abs") | ||
| 509 | + .Store("store", "add") | ||
| 510 | + .Output("output", "store") | ||
| 511 | + .Build(); | ||
| 512 | + CompleteApiInfo(graph); | ||
| 513 | + | ||
| 514 | + optimize::BroadcastBackwardPass pass; | ||
| 515 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 516 | + EXPECT_TRUE(IsConnected(graph, "load0", "abs")); | ||
| 517 | + EXPECT_TRUE(IsConnected(graph, "abs", "broadcast")); | ||
| 518 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "add")); | ||
| 519 | + EXPECT_TRUE(IsConnected(graph, "load1", "add")); | ||
| 520 | +} | ||
| 521 | + | ||
| 522 | +TEST(BroadcastBackwardPass, SkipsFollowingMultiInputBarrier) { | ||
| 523 | + // Not supported by the restored repository BRC implementation. | ||
| 524 | + GTEST_SKIP(); | ||
| 525 | + auto graph = AscGraphBuilder("broadcast_backward_following_multi_input_barrier") | ||
| 526 | + .Loops({Sym("s0"), Sym("s1")}) | ||
| 527 | + .Data("data0", 0) | ||
| 528 | + .Data("data1", 1) | ||
| 529 | + .Data("data2", 2) | ||
| 530 | + .Load("load0", "data0", kCompactRepeats, kCompactStrides) | ||
| 531 | + .Load("load1", "data1", kCompactRepeats, kCompactStrides) | ||
| 532 | + .Load("load2", "data2") | ||
| 533 | + .Broadcast("broadcast0", "load0", {Sym("s0"), Sym("s1")}) | ||
| 534 | + .Broadcast("broadcast1", "load1", {Sym("s0"), Sym("s1")}) | ||
| 535 | + .Add("merge", "broadcast0", "broadcast1") | ||
| 536 | + .Add("barrier", "load2", "merge") | ||
| 537 | + .Store("store", "barrier") | ||
| 538 | + .Output("output", "store") | ||
| 539 | + .Build(); | ||
| 540 | + CompleteApiInfo(graph); | ||
| 541 | + | ||
| 542 | + optimize::BroadcastBackwardPass pass; | ||
| 543 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 544 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "merge")); | ||
| 545 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "merge")); | ||
| 546 | + EXPECT_TRUE(IsConnected(graph, "merge", "barrier")); | ||
| 547 | +} | ||
| 548 | + | ||
| 549 | +TEST(BroadcastBackwardPass, SkipsMultiInputSchedMismatch) { | ||
| 550 | + // Not supported by the restored repository BRC implementation. | ||
| 551 | + GTEST_SKIP(); | ||
| 552 | + auto graph = BuildBinaryGraph("Add"); | ||
| 553 | + CompleteApiInfo(graph); | ||
| 554 | + auto broadcast1_node = FindNode(graph, "broadcast1"); | ||
| 555 | + ASSERT_NE(broadcast1_node, nullptr); | ||
| 556 | + broadcast1_node->attr.sched.axis.clear(); | ||
| 557 | + | ||
| 558 | + optimize::BroadcastBackwardPass pass; | ||
| 559 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 560 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "compute")); | ||
| 561 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "compute")); | ||
| 562 | + EXPECT_TRUE(HasNode(graph, "broadcast1")); | ||
| 563 | +} | ||
| 564 | + | ||
| 565 | +TEST(BroadcastBackwardPass, SkipsTailLayoutMismatch) { | ||
| 566 | + // Not supported by the restored repository BRC implementation. | ||
| 567 | + GTEST_SKIP(); | ||
| 568 | + auto graph = BuildBinaryGraph("Add"); | ||
| 569 | + CompleteApiInfo(graph); | ||
| 570 | + auto broadcast0_node = FindNode(graph, "broadcast0"); | ||
| 571 | + auto broadcast1_node = FindNode(graph, "broadcast1"); | ||
| 572 | + auto compute_node = FindNode(graph, "compute"); | ||
| 573 | + ASSERT_NE(broadcast0_node, nullptr); | ||
| 574 | + ASSERT_NE(broadcast1_node, nullptr); | ||
| 575 | + ASSERT_NE(compute_node, nullptr); | ||
| 576 | + broadcast0_node->outputs[0].attr.strides[0] = af::sym::kSymbolZero; | ||
| 577 | + broadcast1_node->outputs[0].attr.strides[0] = af::sym::kSymbolZero; | ||
| 578 | + compute_node->inputs[0].attr.strides[0] = af::sym::kSymbolZero; | ||
| 579 | + compute_node->inputs[1].attr.strides[0] = af::sym::kSymbolZero; | ||
| 580 | + | ||
| 581 | + optimize::BroadcastBackwardPass pass; | ||
| 582 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 583 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "compute")); | ||
| 584 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "compute")); | ||
| 585 | + EXPECT_TRUE(IsConnected(graph, "compute", "store")); | ||
| 586 | +} | ||
| 587 | + | ||
| 588 | +TEST(BroadcastBackwardPass, SkipsEdgeDtypeMismatch) { | ||
| 589 | + // Not supported by the restored repository BRC implementation. | ||
| 590 | + GTEST_SKIP(); | ||
| 591 | + auto graph = BuildBinaryGraph("Add"); | ||
| 592 | + CompleteApiInfo(graph); | ||
| 593 | + auto broadcast1_node = FindNode(graph, "broadcast1"); | ||
| 594 | + auto add_node = FindNode(graph, "compute"); | ||
| 595 | + ASSERT_NE(broadcast1_node, nullptr); | ||
| 596 | + ASSERT_NE(add_node, nullptr); | ||
| 597 | + broadcast1_node->outputs[0].attr.dtype = af::DT_FLOAT16; | ||
| 598 | + add_node->inputs[1].attr.dtype = af::DT_FLOAT16; | ||
| 599 | + | ||
| 600 | + optimize::BroadcastBackwardPass pass; | ||
| 601 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 602 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "compute")); | ||
| 603 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "compute")); | ||
| 604 | +} | ||
| 605 | + | ||
| 606 | +TEST(BroadcastBackwardPass, SkipsDifferentInputLayouts) { | ||
| 607 | + // Not supported by the restored repository BRC implementation. | ||
| 608 | + GTEST_SKIP(); | ||
| 609 | + auto graph = BuildBinaryGraph("Add"); | ||
| 610 | + CompleteApiInfo(graph); | ||
| 611 | + auto load1_node = FindNode(graph, "load1"); | ||
| 612 | + auto broadcast1_node = FindNode(graph, "broadcast1"); | ||
| 613 | + ASSERT_NE(load1_node, nullptr); | ||
| 614 | + ASSERT_NE(broadcast1_node, nullptr); | ||
| 615 | + load1_node->outputs[0].attr.strides[0] = af::sym::kSymbolZero; | ||
| 616 | + broadcast1_node->inputs[0].attr.strides[0] = af::sym::kSymbolZero; | ||
| 617 | + | ||
| 618 | + optimize::BroadcastBackwardPass pass; | ||
| 619 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 620 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "compute")); | ||
| 621 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "compute")); | ||
| 622 | +} | ||
| 623 | + | ||
| 624 | +TEST(BroadcastBackwardPass, HandlesEmptyGraph) { | ||
| 625 | + af::AscGraph graph("broadcast_backward_empty"); | ||
| 626 | + optimize::BroadcastBackwardPass pass; | ||
| 627 | + EXPECT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 628 | +} | ||
| 629 | + | ||
| 630 | +// ===== Dtype-aware backward tests for Cast and dtype-changing operators ===== | ||
| 631 | + | ||
| 632 | +TEST(BroadcastBackwardPass, MovesCastAcrossBroadcast) { | ||
| 633 | + ScopedTestPlatform platform("3510"); | ||
| 634 | + const auto s0 = Sym("s0"); | ||
| 635 | + const auto s1 = Sym("s1"); | ||
| 636 | + auto graph = AscGraphBuilder("broadcast_backward_cast") | ||
| 637 | + .Loops({s0, s1}) | ||
| 638 | + .Data("data", 0) | ||
| 639 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 640 | + .Broadcast("broadcast", "load", {s0, s1}) | ||
| 641 | + .Cast("cast", "broadcast", af::DT_FLOAT16) | ||
| 642 | + .Store("store", "cast") | ||
| 643 | + .Output("output", "store") | ||
| 644 | + .Build(); | ||
| 645 | + CompleteApiInfo(graph); | ||
| 646 | + SetNodeDtype(graph, "store", af::DT_FLOAT16); | ||
| 647 | + | ||
| 648 | + optimize::BroadcastBackwardPass pass; | ||
| 649 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 650 | + EXPECT_TRUE(IsConnected(graph, "load", "cast")); | ||
| 651 | + EXPECT_TRUE(IsConnected(graph, "cast", "broadcast")); | ||
| 652 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 653 | + const auto broadcast_node = FindNode(graph, "broadcast"); | ||
| 654 | + ASSERT_NE(broadcast_node, nullptr); | ||
| 655 | + EXPECT_EQ(broadcast_node->outputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 656 | + EXPECT_EQ(broadcast_node->inputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 657 | +} | ||
| 658 | + | ||
| 659 | +TEST(BroadcastBackwardPass, MovesComputeAndCastAcrossBroadcast) { | ||
| 660 | + ScopedTestPlatform platform("3510"); | ||
| 661 | + const auto s0 = Sym("s0"); | ||
| 662 | + const auto s1 = Sym("s1"); | ||
| 663 | + auto graph = AscGraphBuilder("broadcast_backward_abs_cast") | ||
| 664 | + .Loops({s0, s1}) | ||
| 665 | + .Data("data", 0) | ||
| 666 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 667 | + .Broadcast("broadcast", "load", {s0, s1}) | ||
| 668 | + .Abs("abs", "broadcast") | ||
| 669 | + .Cast("cast", "abs", af::DT_FLOAT16) | ||
| 670 | + .Store("store", "cast") | ||
| 671 | + .Output("output", "store") | ||
| 672 | + .Build(); | ||
| 673 | + CompleteApiInfo(graph); | ||
| 674 | + SetNodeDtype(graph, "store", af::DT_FLOAT16); | ||
| 675 | + | ||
| 676 | + optimize::BroadcastBackwardPass pass; | ||
| 677 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 678 | + EXPECT_TRUE(IsConnected(graph, "load", "abs")); | ||
| 679 | + EXPECT_TRUE(IsConnected(graph, "abs", "cast")); | ||
| 680 | + EXPECT_TRUE(IsConnected(graph, "cast", "broadcast")); | ||
| 681 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 682 | + const auto broadcast_node = FindNode(graph, "broadcast"); | ||
| 683 | + ASSERT_NE(broadcast_node, nullptr); | ||
| 684 | + EXPECT_EQ(broadcast_node->outputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 685 | +} | ||
| 686 | + | ||
| 687 | +TEST(BroadcastBackwardPass, MovesComparisonAcrossIdenticalBroadcastBranches) { | ||
| 688 | + // Not supported by the restored repository BRC implementation. | ||
| 689 | + GTEST_SKIP(); | ||
| 690 | + const auto s0 = Sym("s0"); | ||
| 691 | + const auto s1 = Sym("s1"); | ||
| 692 | + auto graph = AscGraphBuilder("broadcast_backward_comparison") | ||
| 693 | + .Loops({s0, s1}) | ||
| 694 | + .Data("data0", 0) | ||
| 695 | + .Data("data1", 1) | ||
| 696 | + .Load("load0", "data0", kCompactRepeats, kCompactStrides) | ||
| 697 | + .Load("load1", "data1", kCompactRepeats, kCompactStrides) | ||
| 698 | + .Broadcast("broadcast0", "load0", {s0, s1}) | ||
| 699 | + .Broadcast("broadcast1", "load1", {s0, s1}) | ||
| 700 | + .Op<af::ascir_op::Ge>("compare", {"broadcast0", "broadcast1"}) | ||
| 701 | + .Store("store", "compare") | ||
| 702 | + .Output("output", "store") | ||
| 703 | + .Build(); | ||
| 704 | + CompleteApiInfo(graph); | ||
| 705 | + | ||
| 706 | + optimize::BroadcastBackwardPass pass; | ||
| 707 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 708 | + EXPECT_TRUE(IsConnected(graph, "load0", "compare")); | ||
| 709 | + EXPECT_TRUE(IsConnected(graph, "load1", "compare")); | ||
| 710 | + EXPECT_TRUE(IsConnected(graph, "compare", "broadcast0")); | ||
| 711 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "store")); | ||
| 712 | + EXPECT_FALSE(HasNode(graph, "broadcast1")); | ||
| 713 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load0", "compare")); | ||
| 714 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load1", "compare")); | ||
| 715 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "compare", "broadcast0")); | ||
| 716 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast0", "store")); | ||
| 717 | +} | ||
| 718 | + | ||
| 719 | +TEST(BroadcastBackwardPass, MovesAdditionalDtypeAwareBinaryOps) { | ||
| 720 | + GTEST_SKIP(); | ||
| 721 | + for (const auto &op_name : {std::string("Eq"), std::string("TrueDiv")}) { | ||
| 722 | + SCOPED_TRACE(op_name); | ||
| 723 | + auto graph = BuildDtypeAwareBinaryGraph(op_name); | ||
| 724 | + ExpectDtypeAwareBinaryMove(graph, op_name); | ||
| 725 | + } | ||
| 726 | +} | ||
| 727 | + | ||
| 728 | +TEST(BroadcastBackwardPass, DtypeAwareBackwardEnablesCommonAxisAtMultiInputTail) { | ||
| 729 | + // Not supported by the restored repository BRC implementation. | ||
| 730 | + GTEST_SKIP(); | ||
| 731 | + auto graph = BuildDtypeAwareCommonAxisGraph("broadcast_backward_dtype_aware_common_axis"); | ||
| 732 | + CompleteApiInfo(graph); | ||
| 733 | + SetNodeDtype(graph, "relu", af::DT_FLOAT16); | ||
| 734 | + SetNodeDtype(graph, "merge", af::DT_FLOAT16); | ||
| 735 | + SetNodeDtype(graph, "store", af::DT_FLOAT16); | ||
| 736 | + | ||
| 737 | + optimize::BroadcastBackwardPass pass; | ||
| 738 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 739 | + EXPECT_TRUE(IsConnected(graph, "load0", "abs")); | ||
| 740 | + EXPECT_TRUE(IsConnected(graph, "abs", "cast0")); | ||
| 741 | + EXPECT_TRUE(IsConnected(graph, "load1", "cast1")); | ||
| 742 | + EXPECT_TRUE(IsConnected(graph, "cast1", "relu")); | ||
| 743 | + EXPECT_EQ(FindNode(graph, "broadcast0"), nullptr); | ||
| 744 | + EXPECT_EQ(FindNode(graph, "broadcast1"), nullptr); | ||
| 745 | + const auto residual0 = FindNode(graph, "broadcast0_residual_0"); | ||
| 746 | + const auto residual1 = FindNode(graph, "broadcast1_residual_1"); | ||
| 747 | + const auto merge = FindNode(graph, "merge"); | ||
| 748 | + const auto common = FindNode(graph, "merge_broadcast_backward_common"); | ||
| 749 | + ASSERT_NE(residual0, nullptr); | ||
| 750 | + ASSERT_NE(residual1, nullptr); | ||
| 751 | + ASSERT_NE(merge, nullptr); | ||
| 752 | + ASSERT_NE(common, nullptr); | ||
| 753 | + EXPECT_TRUE(IsConnected(graph, "cast0", "broadcast0_residual_0")); | ||
| 754 | + EXPECT_TRUE(IsConnected(graph, "broadcast0_residual_0", "merge")); | ||
| 755 | + EXPECT_TRUE(IsConnected(graph, "relu", "broadcast1_residual_1")); | ||
| 756 | + EXPECT_TRUE(IsConnected(graph, "broadcast1_residual_1", "merge")); | ||
| 757 | + EXPECT_TRUE(IsConnected(graph, "merge", "merge_broadcast_backward_common")); | ||
| 758 | + EXPECT_TRUE(IsConnected(graph, "merge_broadcast_backward_common", "store")); | ||
| 759 | + EXPECT_EQ(residual0->inputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 760 | + EXPECT_EQ(residual0->outputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 761 | + EXPECT_EQ(residual1->inputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 762 | + EXPECT_EQ(residual1->outputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 763 | + EXPECT_EQ(merge->inputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 764 | + EXPECT_EQ(merge->inputs[1].attr.dtype, af::DT_FLOAT16); | ||
| 765 | + EXPECT_EQ(merge->outputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 766 | + EXPECT_EQ(common->inputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 767 | + EXPECT_EQ(common->outputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 768 | +} | ||
| 769 | + | ||
| 770 | +// ===== Common broadcast-axis partial backward tests ===== | ||
| 771 | + | ||
| 772 | +TEST(BroadcastBackwardPass, CommonAxisBackwardSkipsBroadcastControlEdges) { | ||
| 773 | + for (size_t case_index = 0U; case_index < 2U; ++case_index) { | ||
| 774 | + auto graph = BuildCommonAxisGraph("broadcast_backward_common_axis_control_" + std::to_string(case_index)); | ||
| 775 | + CompleteApiInfo(graph); | ||
| 776 | + const auto broadcast = FindNode(graph, "broadcast0"); | ||
| 777 | + const auto peer = case_index == 0U ? FindNode(graph, "load0") : FindNode(graph, "store"); | ||
| 778 | + ASSERT_NE(broadcast, nullptr); | ||
| 779 | + ASSERT_NE(peer, nullptr); | ||
| 780 | + const auto status = case_index == 0U | ||
| 781 | + ? af::GraphUtils::AddEdge(peer->GetOutControlAnchor(), broadcast->GetInControlAnchor()) | ||
| 782 | + : af::GraphUtils::AddEdge(broadcast->GetOutControlAnchor(), peer->GetInControlAnchor()); | ||
| 783 | + ASSERT_EQ(status, af::SUCCESS); | ||
| 784 | + | ||
| 785 | + optimize::BroadcastBackwardPass pass; | ||
| 786 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 787 | + EXPECT_TRUE(HasNode(graph, "broadcast0")); | ||
| 788 | + EXPECT_TRUE(HasNode(graph, "broadcast1")); | ||
| 789 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 790 | + EXPECT_EQ(broadcast->GetInControlNodesSize() + broadcast->GetOutControlNodesSize(), 1U); | ||
| 791 | + } | ||
| 792 | +} | ||
| 793 | + | ||
| 794 | +TEST(BroadcastBackwardPass, CommonAxisBackwardSkipsSourceControlEdge) { | ||
| 795 | + auto graph = BuildCommonAxisGraph("broadcast_backward_common_axis_source_control"); | ||
| 796 | + CompleteApiInfo(graph); | ||
| 797 | + const auto load = FindNode(graph, "load0"); | ||
| 798 | + const auto store = FindNode(graph, "store"); | ||
| 799 | + ASSERT_NE(load, nullptr); | ||
| 800 | + ASSERT_NE(store, nullptr); | ||
| 801 | + ASSERT_EQ(af::GraphUtils::AddEdge(load->GetOutControlAnchor(), store->GetInControlAnchor()), af::SUCCESS); | ||
| 802 | + | ||
| 803 | + optimize::BroadcastBackwardPass pass; | ||
| 804 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 805 | + EXPECT_TRUE(HasNode(graph, "broadcast0")); | ||
| 806 | + EXPECT_TRUE(HasNode(graph, "broadcast1")); | ||
| 807 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 808 | + EXPECT_EQ(load->GetOutControlNodesSize(), 1U); | ||
| 809 | +} | ||
| 810 | + | ||
| 811 | +TEST(BroadcastBackwardPass, CommonAxisBackwardSkipsSourceBroadcastEdgeAttrMismatch) { | ||
| 812 | + auto graph = BuildCommonAxisGraph("broadcast_backward_common_axis_source_edge_mismatch"); | ||
| 813 | + CompleteApiInfo(graph); | ||
| 814 | + const auto broadcast = FindNode(graph, "broadcast0"); | ||
| 815 | + ASSERT_NE(broadcast, nullptr); | ||
| 816 | + const auto input_desc = broadcast->GetOpDesc()->MutableInputDesc(0U); | ||
| 817 | + ASSERT_NE(input_desc, nullptr); | ||
| 818 | + input_desc->SetDataType(af::DT_FLOAT16); | ||
| 819 | + | ||
| 820 | + optimize::BroadcastBackwardPass pass; | ||
| 821 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 822 | + EXPECT_TRUE(HasNode(graph, "broadcast0")); | ||
| 823 | + EXPECT_TRUE(HasNode(graph, "broadcast1")); | ||
| 824 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 825 | +} | ||
| 826 | + | ||
| 827 | +TEST(BroadcastBackwardPass, MovesCommonAxisFromMultiNodeBroadcastChains) { | ||
| 828 | + // Not supported by the restored repository BRC implementation. | ||
| 829 | + GTEST_SKIP(); | ||
| 830 | + const auto s0 = Sym("s0"); | ||
| 831 | + const auto s1 = Sym("s1"); | ||
| 832 | + const auto s2 = Sym("s2"); | ||
| 833 | + auto graph = AscGraphBuilder("broadcast_backward_common_axis_multi_node") | ||
| 834 | + .Loops({s0, s1, s2}) | ||
| 835 | + .Data("data0", 0) | ||
| 836 | + .Data("data1", 1) | ||
| 837 | + .Load("load0", "data0", {af::sym::kSymbolOne, af::sym::kSymbolOne, s2}, | ||
| 838 | + {af::sym::kSymbolZero, af::sym::kSymbolZero, af::sym::kSymbolOne}) | ||
| 839 | + .Load("load1", "data1", {s0, af::sym::kSymbolOne, af::sym::kSymbolOne}, | ||
| 840 | + {af::sym::kSymbolOne, af::sym::kSymbolZero, af::sym::kSymbolZero}) | ||
| 841 | + .Broadcast("residual0", "load0", {s0, af::sym::kSymbolOne, s2}) | ||
| 842 | + .Broadcast("common0", "residual0", {s0, s1, s2}) | ||
| 843 | + .Broadcast("residual1", "load1", {s0, af::sym::kSymbolOne, s2}) | ||
| 844 | + .Broadcast("common1", "residual1", {s0, s1, s2}) | ||
| 845 | + .Add("merge", "common0", "common1") | ||
| 846 | + .Store("store", "merge") | ||
| 847 | + .Output("output", "store") | ||
| 848 | + .Build(); | ||
| 849 | + CompleteApiInfo(graph); | ||
| 850 | + | ||
| 851 | + optimize::BroadcastBackwardPass pass; | ||
| 852 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 853 | + EXPECT_TRUE(IsConnected(graph, "residual0", "merge")); | ||
| 854 | + EXPECT_TRUE(IsConnected(graph, "residual1", "merge")); | ||
| 855 | + EXPECT_TRUE(IsConnected(graph, "merge", "merge_broadcast_backward_common")); | ||
| 856 | + EXPECT_FALSE(HasNode(graph, "common0_residual_0")); | ||
| 857 | + EXPECT_FALSE(HasNode(graph, "common1_residual_1")); | ||
| 858 | + EXPECT_FALSE(HasNode(graph, "common0")); | ||
| 859 | + EXPECT_FALSE(HasNode(graph, "common1")); | ||
| 860 | +} | ||
| 861 | + | ||
| 862 | +TEST(BroadcastBackwardPass, MovesCreatedCommonBroadcastPastUnaryChainWithoutIdentityResiduals) { | ||
| 863 | + // Not supported by the restored repository BRC implementation. | ||
| 864 | + GTEST_SKIP(); | ||
| 865 | + const auto s0 = Sym("s0"); | ||
| 866 | + const auto s1 = Sym("s1"); | ||
| 867 | + const auto s2 = Sym("s2"); | ||
| 868 | + const std::vector<af::Expression> compact = {s0, af::sym::kSymbolOne, s2}; | ||
| 869 | + const std::vector<af::Expression> compact_strides = {s2, af::sym::kSymbolZero, af::sym::kSymbolOne}; | ||
| 870 | + auto graph = AscGraphBuilder("broadcast_backward_created_common_unary_chain") | ||
| 871 | + .Loops({s0, s1, s2}) | ||
| 872 | + .Data("data0", 0) | ||
| 873 | + .Data("data1", 1) | ||
| 874 | + .Load("load0", "data0", compact, compact_strides) | ||
| 875 | + .Load("load1", "data1", compact, compact_strides) | ||
| 876 | + .Broadcast("broadcast0", "load0", {1}) | ||
| 877 | + .Broadcast("broadcast1", "load1", {1}) | ||
| 878 | + .Abs("abs", "broadcast0") | ||
| 879 | + .Cast("cast0", "abs", af::DT_FLOAT16) | ||
| 880 | + .Cast("cast1", "broadcast1", af::DT_FLOAT16) | ||
| 881 | + .Relu("relu", "cast1") | ||
| 882 | + .Add("merge", "cast0", "relu") | ||
| 883 | + .Sqrt("sqrt", "merge") | ||
| 884 | + .Op<af::ascir_op::Sigmoid>("sigmoid", {"sqrt"}) | ||
| 885 | + .Store("store", "sigmoid") | ||
| 886 | + .Output("output", "store") | ||
| 887 | + .Build(); | ||
| 888 | + CompleteApiInfo(graph); | ||
| 889 | + for (const auto *node_name : {"relu", "merge", "sqrt", "sigmoid", "store"}) { | ||
| 890 | + SetNodeDtype(graph, node_name, af::DT_FLOAT16); | ||
| 891 | + } | ||
| 892 | + | ||
| 893 | + optimize::BroadcastBackwardPass pass; | ||
| 894 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 895 | + EXPECT_FALSE(HasNode(graph, "broadcast0_residual_0")); | ||
| 896 | + EXPECT_FALSE(HasNode(graph, "broadcast1_residual_1")); | ||
| 897 | + EXPECT_TRUE(IsConnected(graph, "cast0", "merge")); | ||
| 898 | + EXPECT_TRUE(IsConnected(graph, "relu", "merge")); | ||
| 899 | + EXPECT_TRUE(IsConnected(graph, "merge", "sqrt")); | ||
| 900 | + EXPECT_TRUE(IsConnected(graph, "sqrt", "sigmoid")); | ||
| 901 | + EXPECT_TRUE(IsConnected(graph, "sigmoid", "merge_broadcast_backward_common")); | ||
| 902 | + EXPECT_TRUE(IsConnected(graph, "merge_broadcast_backward_common", "store")); | ||
| 903 | +} | ||
| 904 | + | ||
| 905 | +TEST(BroadcastBackwardPass, MovesCommonAxisBeforeResidualBroadcasts) { | ||
| 906 | + // Not supported by the restored repository BRC implementation. | ||
| 907 | + GTEST_SKIP(); | ||
| 908 | + const auto s0 = Sym("s0"); | ||
| 909 | + const auto s1 = Sym("s1"); | ||
| 910 | + const auto s2 = Sym("s2"); | ||
| 911 | + auto graph = AscGraphBuilder("broadcast_backward_common_axis_before_residual") | ||
| 912 | + .Loops({s0, s1, s2}) | ||
| 913 | + .Data("data0", 0) | ||
| 914 | + .Data("data1", 1) | ||
| 915 | + .Load("load0", "data0", {af::sym::kSymbolOne, af::sym::kSymbolOne, s2}, | ||
| 916 | + {af::sym::kSymbolZero, af::sym::kSymbolZero, af::sym::kSymbolOne}) | ||
| 917 | + .Load("load1", "data1", {s0, af::sym::kSymbolOne, af::sym::kSymbolOne}, | ||
| 918 | + {af::sym::kSymbolOne, af::sym::kSymbolZero, af::sym::kSymbolZero}) | ||
| 919 | + .Broadcast("common0", "load0", {af::sym::kSymbolOne, s1, s2}) | ||
| 920 | + .Broadcast("residual0", "common0", {s0, s1, s2}) | ||
| 921 | + .Broadcast("common1", "load1", {s0, s1, af::sym::kSymbolOne}) | ||
| 922 | + .Broadcast("residual1", "common1", {s0, s1, s2}) | ||
| 923 | + .Add("merge", "residual0", "residual1") | ||
| 924 | + .Store("store", "merge") | ||
| 925 | + .Output("output", "store") | ||
| 926 | + .Build(); | ||
| 927 | + CompleteApiInfo(graph); | ||
| 928 | + | ||
| 929 | + optimize::BroadcastBackwardPass pass; | ||
| 930 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 931 | + EXPECT_TRUE(IsConnected(graph, "load0", "residual0_residual_0")); | ||
| 932 | + EXPECT_TRUE(IsConnected(graph, "load1", "residual1_residual_1")); | ||
| 933 | + EXPECT_TRUE(IsConnected(graph, "merge", "merge_broadcast_backward_common")); | ||
| 934 | + EXPECT_FALSE(HasNode(graph, "common0")); | ||
| 935 | + EXPECT_FALSE(HasNode(graph, "common1")); | ||
| 936 | + EXPECT_FALSE(HasNode(graph, "residual0")); | ||
| 937 | + EXPECT_FALSE(HasNode(graph, "residual1")); | ||
| 938 | +} | ||
| 939 | + | ||
| 940 | +TEST(BroadcastBackwardPass, SkipsCommonAxesDistributedAcrossBroadcastChain) { | ||
| 941 | + // Not supported by the restored repository BRC implementation. | ||
| 942 | + GTEST_SKIP(); | ||
| 943 | + const auto s0 = Sym("s0"); | ||
| 944 | + const auto s1 = Sym("s1"); | ||
| 945 | + auto graph = AscGraphBuilder("broadcast_backward_distributed_common_axis") | ||
| 946 | + .Loops({s0, s1}) | ||
| 947 | + .Data("data0", 0) | ||
| 948 | + .Data("data1", 1) | ||
| 949 | + .Load("load0", "data0", {af::sym::kSymbolOne, af::sym::kSymbolOne}, | ||
| 950 | + {af::sym::kSymbolZero, af::sym::kSymbolZero}) | ||
| 951 | + .Load("load1", "data1", {af::sym::kSymbolOne, af::sym::kSymbolOne}, | ||
| 952 | + {af::sym::kSymbolZero, af::sym::kSymbolZero}) | ||
| 953 | + .Broadcast("broadcast00", "load0", {s0, af::sym::kSymbolOne}) | ||
| 954 | + .Broadcast("broadcast01", "broadcast00", {s0, s1}) | ||
| 955 | + .Broadcast("broadcast10", "load1", {s0, s1}) | ||
| 956 | + .Add("merge", "broadcast01", "broadcast10") | ||
| 957 | + .Store("store", "merge") | ||
| 958 | + .Output("output", "store") | ||
| 959 | + .Build(); | ||
| 960 | + CompleteApiInfo(graph); | ||
| 961 | + | ||
| 962 | + optimize::BroadcastBackwardPass pass; | ||
| 963 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 964 | + EXPECT_TRUE(IsConnected(graph, "broadcast01", "merge")); | ||
| 965 | + EXPECT_TRUE(IsConnected(graph, "broadcast10", "merge")); | ||
| 966 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 967 | +} | ||
| 968 | + | ||
| 969 | +TEST(BroadcastBackwardPass, MovesCommonAxisPartialBackward) { | ||
| 970 | + // Not supported by the restored repository BRC implementation. | ||
| 971 | + GTEST_SKIP(); | ||
| 972 | + const auto s0 = Sym("s0"); | ||
| 973 | + const auto s1 = Sym("s1"); | ||
| 974 | + const auto s2 = Sym("s2"); | ||
| 975 | + auto graph = BuildCommonAxisGraph("broadcast_backward_common_axis"); | ||
| 976 | + CompleteApiInfo(graph); | ||
| 977 | + | ||
| 978 | + optimize::BroadcastBackwardPass pass; | ||
| 979 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 980 | + auto b_common = FindNode(graph, "merge_broadcast_backward_common"); | ||
| 981 | + EXPECT_NE(b_common, nullptr); | ||
| 982 | + EXPECT_TRUE(IsConnected(graph, "merge", "merge_broadcast_backward_common")); | ||
| 983 | + EXPECT_TRUE(IsConnected(graph, "merge_broadcast_backward_common", "store")); | ||
| 984 | + auto b0_residual = FindNode(graph, "broadcast0_residual_0"); | ||
| 985 | + EXPECT_NE(b0_residual, nullptr); | ||
| 986 | + auto b1_residual = FindNode(graph, "broadcast1_residual_1"); | ||
| 987 | + EXPECT_NE(b1_residual, nullptr); | ||
| 988 | + ExpectResidualConnections(graph); | ||
| 989 | + EXPECT_FALSE(HasNode(graph, "broadcast0")); | ||
| 990 | + EXPECT_FALSE(HasNode(graph, "broadcast1")); | ||
| 991 | + std::vector<af::Expression> expected_compact_strides; | ||
| 992 | + ASSERT_EQ(optimize::ScheduleUtils::RecalculateStridesFromRepeats( | ||
| 993 | + std::vector<af::Expression>{s0, af::sym::kSymbolOne, s2}, expected_compact_strides), | ||
| 994 | + af::SUCCESS); | ||
| 995 | + std::vector<af::Expression> expected_expanded_strides; | ||
| 996 | + ASSERT_EQ(optimize::ScheduleUtils::RecalculateStridesFromRepeats(std::vector<af::Expression>{s0, s1, s2}, | ||
| 997 | + expected_expanded_strides), | ||
| 998 | + af::SUCCESS); | ||
| 999 | + ExpectCommonAxisLayouts(graph, {s0, af::sym::kSymbolOne, s2}, {s0, s1, s2}, expected_compact_strides, | ||
| 1000 | + expected_expanded_strides); | ||
| 1001 | +} | ||
| 1002 | + | ||
| 1003 | +TEST(BroadcastBackwardPass, MovesCommonAxisToActualSuccessorInput) { | ||
| 1004 | + // Not supported by the restored repository BRC implementation. | ||
| 1005 | + GTEST_SKIP(); | ||
| 1006 | + const auto s0 = Sym("s0"); | ||
| 1007 | + const auto s1 = Sym("s1"); | ||
| 1008 | + const auto s2 = Sym("s2"); | ||
| 1009 | + const std::vector<af::Expression> compact0 = {af::sym::kSymbolOne, af::sym::kSymbolOne, s2}; | ||
| 1010 | + const std::vector<af::Expression> strides0 = {af::sym::kSymbolZero, af::sym::kSymbolZero, af::sym::kSymbolOne}; | ||
| 1011 | + const std::vector<af::Expression> compact1 = {s0, af::sym::kSymbolOne, af::sym::kSymbolOne}; | ||
| 1012 | + const std::vector<af::Expression> strides1 = {af::sym::kSymbolOne, af::sym::kSymbolZero, af::sym::kSymbolZero}; | ||
| 1013 | + auto graph = AscGraphBuilder("broadcast_backward_successor_input") | ||
| 1014 | + .Loops({s0, s1, s2}) | ||
| 1015 | + .Data("data0", 0) | ||
| 1016 | + .Data("data1", 1) | ||
| 1017 | + .Load("load0", "data0", compact0, strides0) | ||
| 1018 | + .Load("load1", "data1", compact1, strides1) | ||
| 1019 | + .Broadcast("broadcast0", "load0", {0, 1}) | ||
| 1020 | + .Broadcast("broadcast1", "load1", {1, 2}) | ||
| 1021 | + .Add("merge", "broadcast0", "broadcast1") | ||
| 1022 | + .Add("succ", "merge", "merge") | ||
| 1023 | + .Store("store", "succ") | ||
| 1024 | + .Output("output", "store") | ||
| 1025 | + .Build(); | ||
| 1026 | + CompleteApiInfo(graph); | ||
| 1027 | + const auto merge = FindNode(graph, "merge"); | ||
| 1028 | + const auto succ = FindNode(graph, "succ"); | ||
| 1029 | + ASSERT_NE(merge, nullptr); | ||
| 1030 | + ASSERT_NE(succ, nullptr); | ||
| 1031 | + ASSERT_EQ(af::GraphUtils::RemoveEdge(merge->GetOutDataAnchor(0), succ->GetInDataAnchor(0)), af::SUCCESS); | ||
| 1032 | + | ||
| 1033 | + optimize::BroadcastBackwardPass pass; | ||
| 1034 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1035 | + const auto common = FindNode(graph, "merge_broadcast_backward_common"); | ||
| 1036 | + ASSERT_NE(common, nullptr); | ||
| 1037 | + EXPECT_EQ(succ->GetInDataAnchor(0)->GetPeerOutAnchor(), nullptr); | ||
| 1038 | + const auto succ_input1_peer = succ->GetInDataAnchor(1)->GetPeerOutAnchor(); | ||
| 1039 | + ASSERT_NE(succ_input1_peer, nullptr); | ||
| 1040 | + EXPECT_EQ(succ_input1_peer->GetOwnerNode()->GetName(), "merge_broadcast_backward_common"); | ||
| 1041 | + EXPECT_TRUE(AreConnectedTensorAttrsEqual(common, succ, 1U)); | ||
| 1042 | +} | ||
| 1043 | + | ||
| 1044 | +TEST(BroadcastBackwardPass, SkipsCommonAxisNoOverlap) { | ||
| 1045 | + const auto s0 = Sym("s0"); | ||
| 1046 | + const auto s1 = Sym("s1"); | ||
| 1047 | + const auto s2 = Sym("s2"); | ||
| 1048 | + const std::vector<af::Expression> compact0 = {af::sym::kSymbolOne, af::sym::kSymbolOne, s2}; | ||
| 1049 | + const std::vector<af::Expression> strides0 = {af::sym::kSymbolZero, af::sym::kSymbolZero, af::sym::kSymbolOne}; | ||
| 1050 | + const std::vector<af::Expression> compact1 = {s0, af::sym::kSymbolOne, af::sym::kSymbolOne}; | ||
| 1051 | + const std::vector<af::Expression> strides1 = {af::sym::kSymbolOne, af::sym::kSymbolZero, af::sym::kSymbolZero}; | ||
| 1052 | + auto graph = AscGraphBuilder("broadcast_backward_no_common_axis") | ||
| 1053 | + .Loops({s0, s1, s2}) | ||
| 1054 | + .Data("data0", 0) | ||
| 1055 | + .Data("data1", 1) | ||
| 1056 | + .Load("load0", "data0", compact0, strides0) | ||
| 1057 | + .Load("load1", "data1", compact1, strides1) | ||
| 1058 | + .Broadcast("broadcast0", "load0", {0}) | ||
| 1059 | + .Broadcast("broadcast1", "load1", {2}) | ||
| 1060 | + .Add("merge", "broadcast0", "broadcast1") | ||
| 1061 | + .Store("store", "merge") | ||
| 1062 | + .Output("output", "store") | ||
| 1063 | + .Build(); | ||
| 1064 | + CompleteApiInfo(graph); | ||
| 1065 | + | ||
| 1066 | + optimize::BroadcastBackwardPass pass; | ||
| 1067 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1068 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 1069 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "merge")); | ||
| 1070 | + EXPECT_TRUE(IsConnected(graph, "broadcast1", "merge")); | ||
| 1071 | +} | ||
| 1072 | + | ||
| 1073 | +TEST(BroadcastBackwardPass, MovesNoResidualCaseInIdenticalMultiInputBackward) { | ||
| 1074 | + const auto s0 = Sym("s0"); | ||
| 1075 | + const auto s1 = Sym("s1"); | ||
| 1076 | + const std::vector<af::Expression> compact = {af::sym::kSymbolOne, af::sym::kSymbolOne}; | ||
| 1077 | + const std::vector<af::Expression> strides = {af::sym::kSymbolZero, af::sym::kSymbolZero}; | ||
| 1078 | + auto graph = AscGraphBuilder("broadcast_backward_common_only") | ||
| 1079 | + .Loops({s0, s1}) | ||
| 1080 | + .Data("data0", 0) | ||
| 1081 | + .Data("data1", 1) | ||
| 1082 | + .Load("load0", "data0", compact, strides) | ||
| 1083 | + .Load("load1", "data1", compact, strides) | ||
| 1084 | + .Broadcast("broadcast0", "load0", {0, 1}) | ||
| 1085 | + .Broadcast("broadcast1", "load1", {0, 1}) | ||
| 1086 | + .Add("merge", "broadcast0", "broadcast1") | ||
| 1087 | + .Store("store", "merge") | ||
| 1088 | + .Output("output", "store") | ||
| 1089 | + .Build(); | ||
| 1090 | + CompleteApiInfo(graph); | ||
| 1091 | + | ||
| 1092 | + optimize::BroadcastBackwardPass pass; | ||
| 1093 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1094 | + EXPECT_EQ(FindNode(graph, "merge_broadcast_backward_common"), nullptr); | ||
| 1095 | + const bool has_broadcast0 = HasNode(graph, "broadcast0"); | ||
| 1096 | + const bool has_broadcast1 = HasNode(graph, "broadcast1"); | ||
| 1097 | + ASSERT_NE(has_broadcast0, has_broadcast1); | ||
| 1098 | + const char *kept_broadcast = has_broadcast0 ? "broadcast0" : "broadcast1"; | ||
| 1099 | + EXPECT_TRUE(IsConnected(graph, "load0", "merge")); | ||
| 1100 | + EXPECT_TRUE(IsConnected(graph, "load1", "merge")); | ||
| 1101 | + EXPECT_TRUE(IsConnected(graph, "merge", kept_broadcast)); | ||
| 1102 | + EXPECT_TRUE(IsConnected(graph, kept_broadcast, "store")); | ||
| 1103 | +} | ||
| 1104 | + | ||
| 1105 | +TEST(BroadcastBackwardPass, MovesCommonAxisOldBrcDeletedWithResidual) { | ||
| 1106 | + // Not supported by the restored repository BRC implementation. | ||
| 1107 | + GTEST_SKIP(); | ||
| 1108 | + auto graph = BuildCommonAxisGraph("broadcast_backward_residual_delete_old"); | ||
| 1109 | + CompleteApiInfo(graph); | ||
| 1110 | + | ||
| 1111 | + optimize::BroadcastBackwardPass pass; | ||
| 1112 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1113 | + EXPECT_FALSE(HasNode(graph, "broadcast0")); | ||
| 1114 | + EXPECT_FALSE(HasNode(graph, "broadcast1")); | ||
| 1115 | + auto b0_residual = FindNode(graph, "broadcast0_residual_0"); | ||
| 1116 | + ASSERT_NE(b0_residual, nullptr); | ||
| 1117 | + auto b1_residual = FindNode(graph, "broadcast1_residual_1"); | ||
| 1118 | + ASSERT_NE(b1_residual, nullptr); | ||
| 1119 | + ExpectResidualConnections(graph); | ||
| 1120 | +} | ||
| 1121 | + | ||
| 1122 | +TEST(BroadcastBackwardPass, MultiReferenceBackwardMovesSharedBroadcastInputs) { | ||
| 1123 | + const auto s0 = Sym("s0"); | ||
| 1124 | + const auto s1 = Sym("s1"); | ||
| 1125 | + const std::vector<af::Expression> compact = {s0, af::sym::kSymbolOne}; | ||
| 1126 | + const std::vector<af::Expression> strides = {af::sym::kSymbolOne, af::sym::kSymbolZero}; | ||
| 1127 | + auto graph = AscGraphBuilder("broadcast_backward_shared_common_axis_input") | ||
| 1128 | + .Loops({s0, s1}) | ||
| 1129 | + .Data("data", 0) | ||
| 1130 | + .Load("load", "data", compact, strides) | ||
| 1131 | + .Broadcast("broadcast", "load", {1}) | ||
| 1132 | + .Add("merge", "broadcast", "broadcast") | ||
| 1133 | + .Store("store", "merge") | ||
| 1134 | + .Output("output", "store") | ||
| 1135 | + .Build(); | ||
| 1136 | + CompleteApiInfo(graph); | ||
| 1137 | + | ||
| 1138 | + optimize::BroadcastBackwardPass pass; | ||
| 1139 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1140 | + EXPECT_TRUE(HasNode(graph, "broadcast")); | ||
| 1141 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 1142 | + const auto broadcast = FindNode(graph, "broadcast"); | ||
| 1143 | + ASSERT_NE(broadcast, nullptr); | ||
| 1144 | + EXPECT_TRUE(IsConnected(graph, "load", "merge")); | ||
| 1145 | + EXPECT_TRUE(IsConnected(graph, "merge", "broadcast")); | ||
| 1146 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 1147 | + EXPECT_EQ(broadcast->GetOutDataNodesSize(), 1U); | ||
| 1148 | +} | ||
| 1149 | + | ||
| 1150 | +TEST(BroadcastBackwardPass, SplitsSharedBroadcastBeforeDistinctMultiInputBarriers) { | ||
| 1151 | + // Not supported by the restored repository BRC implementation. | ||
| 1152 | + GTEST_SKIP(); | ||
| 1153 | + const auto s0 = Sym("s0"); | ||
| 1154 | + const auto s1 = Sym("s1"); | ||
| 1155 | + const auto s2 = Sym("s2"); | ||
| 1156 | + const auto s3 = Sym("s3"); | ||
| 1157 | + const std::vector<af::Expression> compact = {af::sym::kSymbolOne, s1, af::sym::kSymbolOne, s3}; | ||
| 1158 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolZero, s3, af::sym::kSymbolZero, | ||
| 1159 | + af::sym::kSymbolOne}; | ||
| 1160 | + const std::vector<af::Expression> expanded = {s0, s1, s2, s3}; | ||
| 1161 | + const std::vector<af::Expression> expanded_strides = {s1 * s2 * s3, s2 * s3, s3, af::sym::kSymbolOne}; | ||
| 1162 | + auto graph = AscGraphBuilder("broadcast_backward_split_shared_branch") | ||
| 1163 | + .Loops({s0, s1, s2, s3}) | ||
| 1164 | + .Data("data0", 0) | ||
| 1165 | + .Load("load0", "data0", compact, compact_strides) | ||
| 1166 | + .Broadcast("broadcast0", "load0", {s0, s1, af::sym::kSymbolOne, s3}) | ||
| 1167 | + .Broadcast("broadcast1", "broadcast0", expanded) | ||
| 1168 | + .Sqrt("sqrt", "broadcast1") | ||
| 1169 | + .Abs("abs", "broadcast1") | ||
| 1170 | + .Data("data1", 1) | ||
| 1171 | + .Load("load1", "data1", expanded, expanded_strides) | ||
| 1172 | + .Sub("sub", "sqrt", "load1") | ||
| 1173 | + .Add("add", "abs", "load1") | ||
| 1174 | + .Neg("neg", "sub") | ||
| 1175 | + .Relu("relu", "add") | ||
| 1176 | + .Mul("mul", "relu", "neg") | ||
| 1177 | + .Store("store", "mul") | ||
| 1178 | + .Output("output", "store") | ||
| 1179 | + .Build(); | ||
| 1180 | + CompleteApiInfo(graph); | ||
| 1181 | + | ||
| 1182 | + optimize::BroadcastBackwardPass pass; | ||
| 1183 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1184 | + EXPECT_TRUE(IsConnected(graph, "load0", "sqrt")); | ||
| 1185 | + EXPECT_TRUE(IsConnected(graph, "load0", "abs")); | ||
| 1186 | + EXPECT_FALSE(IsConnected(graph, "broadcast1", "sqrt")); | ||
| 1187 | + EXPECT_FALSE(IsConnected(graph, "broadcast1", "abs")); | ||
| 1188 | + | ||
| 1189 | + for (const auto &compute_name : {"sqrt", "abs"}) { | ||
| 1190 | + const auto compute = FindNode(graph, compute_name); | ||
| 1191 | + ASSERT_NE(compute, nullptr); | ||
| 1192 | + const auto peers = compute->GetOutDataAnchor(0)->GetPeerInDataAnchors(); | ||
| 1193 | + ASSERT_EQ(peers.size(), 1U); | ||
| 1194 | + EXPECT_EQ((*peers.begin())->GetOwnerNode()->GetType(), "Broadcast"); | ||
| 1195 | + } | ||
| 1196 | +} | ||
| 1197 | + | ||
| 1198 | +TEST(BroadcastBackwardPass, SharedBroadcastSplitSupportsSuccessorInputOne) { | ||
| 1199 | + // Not supported by the restored repository BRC implementation. | ||
| 1200 | + GTEST_SKIP(); | ||
| 1201 | + const auto s0 = Sym("s0"); | ||
| 1202 | + const auto s1 = Sym("s1"); | ||
| 1203 | + const auto s2 = Sym("s2"); | ||
| 1204 | + const std::vector<af::Expression> compact = {af::sym::kSymbolOne, s1, af::sym::kSymbolOne}; | ||
| 1205 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolZero, s2, af::sym::kSymbolZero}; | ||
| 1206 | + const std::vector<af::Expression> expanded = {s0, s1, s2}; | ||
| 1207 | + const std::vector<af::Expression> expanded_strides = {s1 * s2, s2, af::sym::kSymbolOne}; | ||
| 1208 | + auto graph = AscGraphBuilder("broadcast_backward_split_successor_input_one") | ||
| 1209 | + .Loops({s0, s1, s2}) | ||
| 1210 | + .Data("data0", 0) | ||
| 1211 | + .Data("data1", 1) | ||
| 1212 | + .Data("data2", 2) | ||
| 1213 | + .Load("load0", "data0", compact, compact_strides) | ||
| 1214 | + .Load("load1", "data1", expanded, expanded_strides) | ||
| 1215 | + .Load("load2", "data2", expanded, expanded_strides) | ||
| 1216 | + .Broadcast("broadcast0", "load0", expanded) | ||
| 1217 | + .Sqrt("sqrt", "broadcast0") | ||
| 1218 | + .Abs("abs", "broadcast0") | ||
| 1219 | + .Add("successor0", "load1", "sqrt") | ||
| 1220 | + .Add("successor1", "load2", "abs") | ||
| 1221 | + .Add("join", "successor0", "successor1") | ||
| 1222 | + .Store("store", "join") | ||
| 1223 | + .Output("output", "store") | ||
| 1224 | + .Build(); | ||
| 1225 | + CompleteApiInfo(graph); | ||
| 1226 | + | ||
| 1227 | + optimize::BroadcastBackwardPass pass; | ||
| 1228 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1229 | + EXPECT_TRUE(IsConnected(graph, "load0", "sqrt")); | ||
| 1230 | + EXPECT_TRUE(IsConnected(graph, "load0", "abs")); | ||
| 1231 | + EXPECT_FALSE(IsConnected(graph, "broadcast0", "sqrt")); | ||
| 1232 | + EXPECT_FALSE(IsConnected(graph, "broadcast0", "abs")); | ||
| 1233 | + EXPECT_TRUE(IsConnected(graph, "broadcast0_branch_split_1", "successor0")); | ||
| 1234 | + EXPECT_TRUE(HasNode(graph, "broadcast0_branch_split_1")); | ||
| 1235 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "successor1")); | ||
| 1236 | +} | ||
| 1237 | + | ||
| 1238 | +TEST(BroadcastBackwardPass, SplitsAndMovesSharedBroadcastPerBranch) { | ||
| 1239 | + // Not supported by the restored repository BRC implementation. | ||
| 1240 | + GTEST_SKIP(); | ||
| 1241 | + const auto s0 = Sym("s0"); | ||
| 1242 | + const auto s1 = Sym("s1"); | ||
| 1243 | + const auto s2 = Sym("s2"); | ||
| 1244 | + const std::vector<af::Expression> compact = {af::sym::kSymbolOne, s1, af::sym::kSymbolOne}; | ||
| 1245 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolZero, s2, af::sym::kSymbolZero}; | ||
| 1246 | + const std::vector<af::Expression> expanded = {s0, s1, s2}; | ||
| 1247 | + const std::vector<af::Expression> expanded_strides = {s1 * s2, s2, af::sym::kSymbolOne}; | ||
| 1248 | + auto graph = AscGraphBuilder("broadcast_backward_split_and_move_per_branch") | ||
| 1249 | + .Loops({s0, s1, s2}) | ||
| 1250 | + .Data("data0", 0) | ||
| 1251 | + .Data("data1", 1) | ||
| 1252 | + .Load("load0", "data0", compact, compact_strides) | ||
| 1253 | + .Load("load1", "data1", expanded, expanded_strides) | ||
| 1254 | + .Broadcast("broadcast", "load0", expanded) | ||
| 1255 | + .Sqrt("sqrt", "broadcast") | ||
| 1256 | + .Add("add", "sqrt", "load1") | ||
| 1257 | + .Neg("neg", "add") | ||
| 1258 | + .Abs("abs", "broadcast") | ||
| 1259 | + .Relu("relu", "abs") | ||
| 1260 | + .Mul("mul", "relu", "neg") | ||
| 1261 | + .Store("store", "mul") | ||
| 1262 | + .Output("output", "store") | ||
| 1263 | + .Build(); | ||
| 1264 | + CompleteApiInfo(graph); | ||
| 1265 | + | ||
| 1266 | + optimize::BroadcastBackwardPass pass; | ||
| 1267 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1268 | + EXPECT_TRUE(IsConnected(graph, "load0", "sqrt")); | ||
| 1269 | + EXPECT_TRUE(IsConnected(graph, "load0", "abs")); | ||
| 1270 | + EXPECT_FALSE(IsConnected(graph, "broadcast", "sqrt")); | ||
| 1271 | + EXPECT_FALSE(IsConnected(graph, "broadcast", "abs")); | ||
| 1272 | + | ||
| 1273 | + const auto sqrt = FindNode(graph, "sqrt"); | ||
| 1274 | + const auto relu = FindNode(graph, "relu"); | ||
| 1275 | + ASSERT_NE(sqrt, nullptr); | ||
| 1276 | + ASSERT_NE(relu, nullptr); | ||
| 1277 | + ASSERT_EQ(sqrt->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 1U); | ||
| 1278 | + ASSERT_EQ(relu->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 1U); | ||
| 1279 | + EXPECT_EQ((*sqrt->GetOutDataAnchor(0)->GetPeerInDataAnchors().begin())->GetOwnerNode()->GetType(), "Broadcast"); | ||
| 1280 | + EXPECT_EQ((*relu->GetOutDataAnchor(0)->GetPeerInDataAnchors().begin())->GetOwnerNode()->GetType(), "Broadcast"); | ||
| 1281 | +} | ||
| 1282 | + | ||
| 1283 | +TEST(BroadcastBackwardPass, SharedBroadcastSplitSkipsCommonMerge) { | ||
| 1284 | + // Not supported by the restored repository BRC implementation. | ||
| 1285 | + GTEST_SKIP(); | ||
| 1286 | + const auto s0 = Sym("s0"); | ||
| 1287 | + const auto s1 = Sym("s1"); | ||
| 1288 | + const std::vector<af::Expression> compact = {s0, af::sym::kSymbolOne}; | ||
| 1289 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolOne, af::sym::kSymbolZero}; | ||
| 1290 | + const std::vector<af::Expression> expanded = {s0, s1}; | ||
| 1291 | + const std::vector<af::Expression> expanded_strides = {s1, af::sym::kSymbolOne}; | ||
| 1292 | + auto graph = AscGraphBuilder("broadcast_backward_split_skips_common_merge") | ||
| 1293 | + .Loops({s0, s1}) | ||
| 1294 | + .Data("data", 0) | ||
| 1295 | + .Load("load", "data", compact, compact_strides) | ||
| 1296 | + .Broadcast("broadcast", "load", expanded) | ||
| 1297 | + .Abs("abs", "broadcast") | ||
| 1298 | + .Neg("neg", "broadcast") | ||
| 1299 | + .Add("merge", "abs", "neg") | ||
| 1300 | + .Store("store", "merge") | ||
| 1301 | + .Output("output", "store") | ||
| 1302 | + .Build(); | ||
| 1303 | + CompleteApiInfo(graph); | ||
| 1304 | + | ||
| 1305 | + optimize::BroadcastBackwardPass pass; | ||
| 1306 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1307 | + EXPECT_FALSE(HasNode(graph, "broadcast_consumer_split_1")); | ||
| 1308 | + EXPECT_TRUE(IsConnected(graph, "merge", "broadcast")); | ||
| 1309 | +} | ||
| 1310 | + | ||
| 1311 | +TEST(BroadcastBackwardPass, SharedBroadcastSplitSkipsWithoutMovableBranch) { | ||
| 1312 | + const auto s0 = Sym("s0"); | ||
| 1313 | + const auto s1 = Sym("s1"); | ||
| 1314 | + const std::vector<af::Expression> compact = {s0, af::sym::kSymbolOne}; | ||
| 1315 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolOne, af::sym::kSymbolZero}; | ||
| 1316 | + const std::vector<af::Expression> expanded = {s0, s1}; | ||
| 1317 | + auto graph = AscGraphBuilder("broadcast_backward_split_skips_non_movable") | ||
| 1318 | + .Loops({s0, s1}) | ||
| 1319 | + .Data("data", 0) | ||
| 1320 | + .Load("load", "data", compact, compact_strides) | ||
| 1321 | + .Broadcast("broadcast", "load", expanded) | ||
| 1322 | + .Store("store0", "broadcast") | ||
| 1323 | + .Output("output0", "store0") | ||
| 1324 | + .Store("store1", "broadcast") | ||
| 1325 | + .Output("output1", "store1") | ||
| 1326 | + .Build(); | ||
| 1327 | + CompleteApiInfo(graph); | ||
| 1328 | + | ||
| 1329 | + optimize::BroadcastBackwardPass pass; | ||
| 1330 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1331 | + EXPECT_FALSE(HasNode(graph, "broadcast_consumer_split_1")); | ||
| 1332 | + EXPECT_EQ(FindNode(graph, "broadcast")->GetOutDataNodesSize(), 2U); | ||
| 1333 | +} | ||
| 1334 | + | ||
| 1335 | +TEST(BroadcastBackwardPass, SharedBroadcastSplitSkipsMoreThanEightBranches) { | ||
| 1336 | + // Not supported by the restored repository BRC implementation. | ||
| 1337 | + GTEST_SKIP(); | ||
| 1338 | + const auto s0 = Sym("s0"); | ||
| 1339 | + const auto s1 = Sym("s1"); | ||
| 1340 | + const std::vector<af::Expression> compact = {s0, af::sym::kSymbolOne}; | ||
| 1341 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolOne, af::sym::kSymbolZero}; | ||
| 1342 | + const std::vector<af::Expression> expanded = {s0, s1}; | ||
| 1343 | + const std::vector<af::Expression> expanded_strides = {s1, af::sym::kSymbolOne}; | ||
| 1344 | + AscGraphBuilder builder("broadcast_backward_split_skips_more_than_eight_branches"); | ||
| 1345 | + builder.Loops({s0, s1}) | ||
| 1346 | + .Data("data0", 0) | ||
| 1347 | + .Data("data1", 1) | ||
| 1348 | + .Load("load0", "data0", compact, compact_strides) | ||
| 1349 | + .Load("load1", "data1", expanded, expanded_strides) | ||
| 1350 | + .Broadcast("broadcast", "load0", expanded); | ||
| 1351 | + for (size_t index = 0U; index < 9U; ++index) { | ||
| 1352 | + const auto suffix = std::to_string(index); | ||
| 1353 | + builder.Abs("abs" + suffix, "broadcast").Add("add" + suffix, "load1", "abs" + suffix); | ||
| 1354 | + } | ||
| 1355 | + builder.Store("store", "add0").Output("output", "store"); | ||
| 1356 | + auto graph = builder.Build(); | ||
| 1357 | + CompleteApiInfo(graph); | ||
| 1358 | + | ||
| 1359 | + optimize::BroadcastBackwardPass pass; | ||
| 1360 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1361 | + EXPECT_FALSE(HasNode(graph, "broadcast_consumer_split_1")); | ||
| 1362 | + EXPECT_EQ(FindNode(graph, "broadcast")->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 9U); | ||
| 1363 | +} | ||
| 1364 | + | ||
| 1365 | +TEST(BroadcastBackwardPass, MovesBroadcastToNonZeroMultiInputTail) { | ||
| 1366 | + const auto s0 = Sym("s0"); | ||
| 1367 | + const auto s1 = Sym("s1"); | ||
| 1368 | + const std::vector<af::Expression> compact = {s0, af::sym::kSymbolOne}; | ||
| 1369 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolOne, af::sym::kSymbolZero}; | ||
| 1370 | + const std::vector<af::Expression> expanded = {s0, s1}; | ||
| 1371 | + const std::vector<af::Expression> expanded_strides = {s1, af::sym::kSymbolOne}; | ||
| 1372 | + auto graph = AscGraphBuilder("broadcast_backward_non_zero_tail_input") | ||
| 1373 | + .Loops({s0, s1}) | ||
| 1374 | + .Data("data0", 0) | ||
| 1375 | + .Data("data1", 1) | ||
| 1376 | + .Load("load0", "data0", compact, compact_strides) | ||
| 1377 | + .Load("load1", "data1", expanded, expanded_strides) | ||
| 1378 | + .Broadcast("broadcast", "load0", expanded) | ||
| 1379 | + .Abs("abs", "broadcast") | ||
| 1380 | + .Add("add", "load1", "abs") | ||
| 1381 | + .Store("store", "add") | ||
| 1382 | + .Output("output", "store") | ||
| 1383 | + .Build(); | ||
| 1384 | + CompleteApiInfo(graph); | ||
| 1385 | + | ||
| 1386 | + optimize::BroadcastBackwardPass pass; | ||
| 1387 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1388 | + EXPECT_TRUE(IsConnected(graph, "abs", "broadcast")); | ||
| 1389 | + EXPECT_EQ(FindNode(graph, "add")->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "broadcast"); | ||
| 1390 | +} | ||
| 1391 | + | ||
| 1392 | +TEST(BroadcastBackwardPass, SharedBroadcastSplitSkipsSingleInputSuccessor) { | ||
| 1393 | + // Not supported by the restored repository BRC implementation. | ||
| 1394 | + GTEST_SKIP(); | ||
| 1395 | + const auto s0 = Sym("s0"); | ||
| 1396 | + const auto s1 = Sym("s1"); | ||
| 1397 | + const auto s2 = Sym("s2"); | ||
| 1398 | + const std::vector<af::Expression> compact = {af::sym::kSymbolOne, s1, af::sym::kSymbolOne}; | ||
| 1399 | + const std::vector<af::Expression> compact_strides = {af::sym::kSymbolZero, s2, af::sym::kSymbolZero}; | ||
| 1400 | + const std::vector<af::Expression> expanded = {s0, s1, s2}; | ||
| 1401 | + const std::vector<af::Expression> expanded_strides = {s1 * s2, s2, af::sym::kSymbolOne}; | ||
| 1402 | + auto graph = AscGraphBuilder("broadcast_backward_split_single_input_successor") | ||
| 1403 | + .Loops({s0, s1, s2}) | ||
| 1404 | + .Data("data0", 0) | ||
| 1405 | + .Load("load0", "data0", compact, compact_strides) | ||
| 1406 | + .Broadcast("broadcast0", "load0", expanded) | ||
| 1407 | + .Sqrt("sqrt", "broadcast0") | ||
| 1408 | + .Abs("abs", "broadcast0") | ||
| 1409 | + .Neg("neg", "sqrt") | ||
| 1410 | + .Relu("relu", "abs") | ||
| 1411 | + .Mul("mul", "relu", "neg") | ||
| 1412 | + .Store("store", "mul") | ||
| 1413 | + .Output("output", "store") | ||
| 1414 | + .Build(); | ||
| 1415 | + CompleteApiInfo(graph); | ||
| 1416 | + | ||
| 1417 | + optimize::BroadcastBackwardPass pass; | ||
| 1418 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1419 | + EXPECT_FALSE(HasNode(graph, "broadcast0_branch_split_1")); | ||
| 1420 | + EXPECT_TRUE(IsConnected(graph, "mul", "broadcast0")); | ||
| 1421 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "store")); | ||
| 1422 | +} | ||
| 1423 | + | ||
| 1424 | +TEST(BroadcastBackwardPass, CommonAxisBackwardSkipsBroadcastWithAnotherConsumer) { | ||
| 1425 | + // Not supported by the restored repository BRC implementation. | ||
| 1426 | + GTEST_SKIP(); | ||
| 1427 | + const auto s0 = Sym("s0"); | ||
| 1428 | + const auto s1 = Sym("s1"); | ||
| 1429 | + const auto s2 = Sym("s2"); | ||
| 1430 | + const std::vector<af::Expression> compact0 = {af::sym::kSymbolOne, af::sym::kSymbolOne, s2}; | ||
| 1431 | + const std::vector<af::Expression> strides0 = {af::sym::kSymbolZero, af::sym::kSymbolZero, af::sym::kSymbolOne}; | ||
| 1432 | + const std::vector<af::Expression> compact1 = {s0, af::sym::kSymbolOne, af::sym::kSymbolOne}; | ||
| 1433 | + const std::vector<af::Expression> strides1 = {af::sym::kSymbolOne, af::sym::kSymbolZero, af::sym::kSymbolZero}; | ||
| 1434 | + auto graph = AscGraphBuilder("broadcast_backward_extra_consumer") | ||
| 1435 | + .Loops({s0, s1, s2}) | ||
| 1436 | + .Data("data0", 0) | ||
| 1437 | + .Data("data1", 1) | ||
| 1438 | + .Load("load0", "data0", compact0, strides0) | ||
| 1439 | + .Load("load1", "data1", compact1, strides1) | ||
| 1440 | + .Broadcast("broadcast0", "load0", {0, 1}) | ||
| 1441 | + .Broadcast("broadcast1", "load1", {1, 2}) | ||
| 1442 | + .Add("merge", "broadcast0", "broadcast1") | ||
| 1443 | + .Abs("side", "broadcast0") | ||
| 1444 | + .Store("store", "merge") | ||
| 1445 | + .Store("side_store", "side") | ||
| 1446 | + .Output("output", "store") | ||
| 1447 | + .Build(); | ||
| 1448 | + CompleteApiInfo(graph); | ||
| 1449 | + | ||
| 1450 | + optimize::BroadcastBackwardPass pass; | ||
| 1451 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1452 | + EXPECT_TRUE(HasNode(graph, "broadcast0")); | ||
| 1453 | + EXPECT_TRUE(HasNode(graph, "broadcast1")); | ||
| 1454 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 1455 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "merge")); | ||
| 1456 | + EXPECT_TRUE(IsConnected(graph, "broadcast0", "side")); | ||
| 1457 | +} | ||
| 1458 | + | ||
| 1459 | +TEST(BroadcastBackwardPass, SkipsCommonAxisEdgeAttrMismatch) { | ||
| 1460 | + auto graph = BuildCommonAxisGraph("broadcast_backward_edge_mismatch"); | ||
| 1461 | + CompleteApiInfo(graph); | ||
| 1462 | + auto merge_node = FindNode(graph, "merge"); | ||
| 1463 | + ASSERT_NE(merge_node, nullptr); | ||
| 1464 | + merge_node->inputs[0].attr.dtype = af::DT_FLOAT16; | ||
| 1465 | + | ||
| 1466 | + optimize::BroadcastBackwardPass pass; | ||
| 1467 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1468 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 1469 | + EXPECT_TRUE(HasNode(graph, "broadcast0")); | ||
| 1470 | + EXPECT_TRUE(HasNode(graph, "broadcast1")); | ||
| 1471 | +} | ||
| 1472 | + | ||
| 1473 | +TEST(BroadcastBackwardPass, SkipsCommonAxisDtypeMismatch) { | ||
| 1474 | + auto graph = BuildCommonAxisGraph("broadcast_backward_dtype_mismatch"); | ||
| 1475 | + CompleteApiInfo(graph); | ||
| 1476 | + auto load1_node = FindNode(graph, "load1"); | ||
| 1477 | + ASSERT_NE(load1_node, nullptr); | ||
| 1478 | + load1_node->outputs[0].attr.dtype = af::DT_FLOAT16; | ||
| 1479 | + auto b1_node = FindNode(graph, "broadcast1"); | ||
| 1480 | + ASSERT_NE(b1_node, nullptr); | ||
| 1481 | + b1_node->inputs[0].attr.dtype = af::DT_FLOAT16; | ||
| 1482 | + b1_node->outputs[0].attr.dtype = af::DT_FLOAT16; | ||
| 1483 | + | ||
| 1484 | + optimize::BroadcastBackwardPass pass; | ||
| 1485 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1486 | + EXPECT_FALSE(HasNode(graph, "merge_broadcast_backward_common")); | ||
| 1487 | + EXPECT_TRUE(HasNode(graph, "broadcast0")); | ||
| 1488 | + EXPECT_TRUE(HasNode(graph, "broadcast1")); | ||
| 1489 | +} | ||
| 1490 | + | ||
| 1491 | +// ===== Multi-reference fork-join backward tests ===== | ||
| 1492 | + | ||
| 1493 | +TEST(BroadcastBackwardPass, GraphUtilsReconnectsSameSourceMultipleInputs) { | ||
| 1494 | + const auto s0 = Sym("s0"); | ||
| 1495 | + const auto s1 = Sym("s1"); | ||
| 1496 | + auto graph = AscGraphBuilder("broadcast_backward_multi_edge_graph_utils") | ||
| 1497 | + .Loops({s0, s1}) | ||
| 1498 | + .Data("data", 0) | ||
| 1499 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1500 | + .Broadcast("broadcast", "load", {1}) | ||
| 1501 | + .Add("consumer", "broadcast", "broadcast") | ||
| 1502 | + .Store("store", "consumer") | ||
| 1503 | + .Output("output", "store") | ||
| 1504 | + .Build(); | ||
| 1505 | + CompleteApiInfo(graph); | ||
| 1506 | + const auto load = FindNode(graph, "load"); | ||
| 1507 | + const auto broadcast = FindNode(graph, "broadcast"); | ||
| 1508 | + const auto consumer = FindNode(graph, "consumer"); | ||
| 1509 | + ASSERT_NE(load, nullptr); | ||
| 1510 | + ASSERT_NE(broadcast, nullptr); | ||
| 1511 | + ASSERT_NE(consumer, nullptr); | ||
| 1512 | + const auto peer_copy = broadcast->GetOutDataAnchor(0)->GetPeerInDataAnchors(); | ||
| 1513 | + ASSERT_EQ(peer_copy.size(), 2U); | ||
| 1514 | + for (const auto &peer : peer_copy) { | ||
| 1515 | + ASSERT_EQ(af::GraphUtils::RemoveEdge(broadcast->GetOutDataAnchor(0), peer), af::SUCCESS); | ||
| 1516 | + } | ||
| 1517 | + EXPECT_TRUE(broadcast->GetOutDataAnchor(0)->GetPeerInDataAnchors().empty()); | ||
| 1518 | + ASSERT_EQ(af::GraphUtils::AddEdge(load->GetOutDataAnchor(0), consumer->GetInDataAnchor(0)), af::SUCCESS); | ||
| 1519 | + ASSERT_EQ(af::GraphUtils::AddEdge(load->GetOutDataAnchor(0), consumer->GetInDataAnchor(1)), af::SUCCESS); | ||
| 1520 | + ASSERT_EQ(af::GraphUtils::RemoveEdge(load->GetOutDataAnchor(0), consumer->GetInDataAnchor(0)), af::SUCCESS); | ||
| 1521 | + ASSERT_EQ(af::GraphUtils::RemoveEdge(load->GetOutDataAnchor(0), consumer->GetInDataAnchor(1)), af::SUCCESS); | ||
| 1522 | + ASSERT_EQ(af::GraphUtils::AddEdge(broadcast->GetOutDataAnchor(0), consumer->GetInDataAnchor(0)), af::SUCCESS); | ||
| 1523 | + ASSERT_EQ(af::GraphUtils::AddEdge(broadcast->GetOutDataAnchor(0), consumer->GetInDataAnchor(1)), af::SUCCESS); | ||
| 1524 | + EXPECT_EQ(broadcast->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 2U); | ||
| 1525 | + EXPECT_TRUE(AreConnectedTensorAttrsEqual(broadcast, consumer, 0U)); | ||
| 1526 | + EXPECT_TRUE(AreConnectedTensorAttrsEqual(broadcast, consumer, 1U)); | ||
| 1527 | +} | ||
| 1528 | + | ||
| 1529 | +TEST(BroadcastBackwardPass, MovesDirectFanOutAfterMerge) { | ||
| 1530 | + auto graph = BuildDirectFanOutGraph("broadcast_backward_multi_reference_direct"); | ||
| 1531 | + CompleteApiInfo(graph); | ||
| 1532 | + ExpectDirectFanOutCandidate(graph); | ||
| 1533 | + | ||
| 1534 | + optimize::BroadcastBackwardPass pass; | ||
| 1535 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1536 | + ExpectDirectFanOutMoved(graph); | ||
| 1537 | +} | ||
| 1538 | + | ||
| 1539 | +TEST(BroadcastBackwardPass, MovesPrefixFanOutAfterMerge) { | ||
| 1540 | + const auto s0 = Sym("s0"); | ||
| 1541 | + const auto s1 = Sym("s1"); | ||
| 1542 | + auto graph = AscGraphBuilder("broadcast_backward_multi_reference_prefix") | ||
| 1543 | + .Loops({s0, s1}) | ||
| 1544 | + .Data("data", 0) | ||
| 1545 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1546 | + .Broadcast("broadcast", "load", {1}) | ||
| 1547 | + .Relu("prefix", "broadcast") | ||
| 1548 | + .Abs("branch0", "prefix") | ||
| 1549 | + .Neg("branch1", "prefix") | ||
| 1550 | + .Add("merge", "branch0", "branch1") | ||
| 1551 | + .Store("store", "merge") | ||
| 1552 | + .Output("output", "store") | ||
| 1553 | + .Build(); | ||
| 1554 | + CompleteApiInfo(graph); | ||
| 1555 | + | ||
| 1556 | + optimize::BroadcastBackwardPass pass; | ||
| 1557 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1558 | + EXPECT_TRUE(IsConnected(graph, "load", "prefix")); | ||
| 1559 | + EXPECT_TRUE(IsConnected(graph, "prefix", "branch0")); | ||
| 1560 | + EXPECT_TRUE(IsConnected(graph, "prefix", "branch1")); | ||
| 1561 | + EXPECT_TRUE(IsConnected(graph, "branch0", "merge")); | ||
| 1562 | + EXPECT_TRUE(IsConnected(graph, "branch1", "merge")); | ||
| 1563 | + EXPECT_TRUE(IsConnected(graph, "merge", "broadcast")); | ||
| 1564 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 1565 | + const auto prefix = FindNode(graph, "prefix"); | ||
| 1566 | + ASSERT_NE(prefix, nullptr); | ||
| 1567 | + ExpectStaticEq(prefix->outputs[0].attr.repeats, kCompactRepeats); | ||
| 1568 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load", "prefix")); | ||
| 1569 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "prefix", "branch0")); | ||
| 1570 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "prefix", "branch1")); | ||
| 1571 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch0", "merge")); | ||
| 1572 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch1", "merge")); | ||
| 1573 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "merge", "broadcast")); | ||
| 1574 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "store")); | ||
| 1575 | +} | ||
| 1576 | + | ||
| 1577 | +TEST(BroadcastBackwardPass, MovesMultiNodeFanOutBranches) { | ||
| 1578 | + auto graph = BuildMultiNodeFanOutGraph("broadcast_backward_multi_reference_multi_node_branches"); | ||
| 1579 | + CompleteApiInfo(graph); | ||
| 1580 | + | ||
| 1581 | + optimize::BroadcastBackwardPass pass; | ||
| 1582 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1583 | + EXPECT_TRUE(IsConnected(graph, "load", "branch0_head")); | ||
| 1584 | + EXPECT_TRUE(IsConnected(graph, "load", "branch1_head")); | ||
| 1585 | + EXPECT_TRUE(IsConnected(graph, "branch0_head", "branch0_tail")); | ||
| 1586 | + EXPECT_TRUE(IsConnected(graph, "branch1_head", "branch1_tail")); | ||
| 1587 | + EXPECT_TRUE(IsConnected(graph, "branch0_tail", "merge")); | ||
| 1588 | + EXPECT_TRUE(IsConnected(graph, "branch1_tail", "merge")); | ||
| 1589 | + EXPECT_TRUE(IsConnected(graph, "merge", "broadcast")); | ||
| 1590 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 1591 | + const auto branch0_tail = FindNode(graph, "branch0_tail"); | ||
| 1592 | + const auto branch1_tail = FindNode(graph, "branch1_tail"); | ||
| 1593 | + ASSERT_NE(branch0_tail, nullptr); | ||
| 1594 | + ASSERT_NE(branch1_tail, nullptr); | ||
| 1595 | + ExpectStaticEq(branch0_tail->outputs[0].attr.repeats, kCompactRepeats); | ||
| 1596 | + ExpectStaticEq(branch1_tail->outputs[0].attr.repeats, kCompactRepeats); | ||
| 1597 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load", "branch0_head")); | ||
| 1598 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch0_head", "branch0_tail")); | ||
| 1599 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch0_tail", "merge")); | ||
| 1600 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "load", "branch1_head")); | ||
| 1601 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch1_head", "branch1_tail")); | ||
| 1602 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "branch1_tail", "merge")); | ||
| 1603 | +} | ||
| 1604 | + | ||
| 1605 | +TEST(BroadcastBackwardPass, MovesFanOutAcrossFollowingComputeChain) { | ||
| 1606 | + const auto s0 = Sym("s0"); | ||
| 1607 | + const auto s1 = Sym("s1"); | ||
| 1608 | + auto graph = AscGraphBuilder("broadcast_backward_multi_reference_following_compute") | ||
| 1609 | + .Loops({s0, s1}) | ||
| 1610 | + .Data("data", 0) | ||
| 1611 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1612 | + .Broadcast("broadcast", "load", {1}) | ||
| 1613 | + .Abs("branch0", "broadcast") | ||
| 1614 | + .Neg("branch1", "broadcast") | ||
| 1615 | + .Add("merge", "branch0", "branch1") | ||
| 1616 | + .Relu("following0", "merge") | ||
| 1617 | + .Exp("following1", "following0") | ||
| 1618 | + .Store("store", "following1") | ||
| 1619 | + .Output("output", "store") | ||
| 1620 | + .Build(); | ||
| 1621 | + CompleteApiInfo(graph); | ||
| 1622 | + | ||
| 1623 | + optimize::BroadcastBackwardPass pass; | ||
| 1624 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1625 | + EXPECT_TRUE(IsConnected(graph, "load", "branch0")); | ||
| 1626 | + EXPECT_TRUE(IsConnected(graph, "load", "branch1")); | ||
| 1627 | + EXPECT_TRUE(IsConnected(graph, "merge", "following0")); | ||
| 1628 | + EXPECT_TRUE(IsConnected(graph, "following0", "following1")); | ||
| 1629 | + EXPECT_TRUE(IsConnected(graph, "following1", "broadcast")); | ||
| 1630 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 1631 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "merge", "following0")); | ||
| 1632 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "following0", "following1")); | ||
| 1633 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "following1", "broadcast")); | ||
| 1634 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "store")); | ||
| 1635 | +} | ||
| 1636 | + | ||
| 1637 | +TEST(BroadcastBackwardPass, DtypeAwareBackwardEnablesPrefixFanOut) { | ||
| 1638 | + ScopedTestPlatform platform("3510"); | ||
| 1639 | + const auto s0 = Sym("s0"); | ||
| 1640 | + const auto s1 = Sym("s1"); | ||
| 1641 | + auto graph = AscGraphBuilder("broadcast_backward_dtype_aware_prefix_fanout") | ||
| 1642 | + .Loops({s0, s1}) | ||
| 1643 | + .Data("data", 0) | ||
| 1644 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1645 | + .Broadcast("broadcast", "load", {1}) | ||
| 1646 | + .Sqrt("before_cast", "broadcast") | ||
| 1647 | + .Cast("cast", "before_cast", af::DT_FLOAT16) | ||
| 1648 | + .Relu("prefix", "cast") | ||
| 1649 | + .Abs("branch0", "prefix") | ||
| 1650 | + .Neg("branch1", "prefix") | ||
| 1651 | + .Add("merge", "branch0", "branch1") | ||
| 1652 | + .Store("store", "merge") | ||
| 1653 | + .Output("output", "store") | ||
| 1654 | + .Build(); | ||
| 1655 | + CompleteApiInfo(graph); | ||
| 1656 | + SetNodeDtype(graph, "prefix", af::DT_FLOAT16); | ||
| 1657 | + SetNodeDtype(graph, "branch0", af::DT_FLOAT16); | ||
| 1658 | + SetNodeDtype(graph, "branch1", af::DT_FLOAT16); | ||
| 1659 | + SetNodeDtype(graph, "merge", af::DT_FLOAT16); | ||
| 1660 | + SetNodeDtype(graph, "store", af::DT_FLOAT16); | ||
| 1661 | + | ||
| 1662 | + optimize::BroadcastBackwardPass pass; | ||
| 1663 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1664 | + EXPECT_TRUE(IsConnected(graph, "load", "before_cast")); | ||
| 1665 | + EXPECT_TRUE(IsConnected(graph, "before_cast", "cast")); | ||
| 1666 | + EXPECT_TRUE(IsConnected(graph, "cast", "prefix")); | ||
| 1667 | + EXPECT_TRUE(IsConnected(graph, "prefix", "branch0")); | ||
| 1668 | + EXPECT_TRUE(IsConnected(graph, "prefix", "branch1")); | ||
| 1669 | + EXPECT_TRUE(IsConnected(graph, "merge", "broadcast")); | ||
| 1670 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 1671 | + const auto broadcast = FindNode(graph, "broadcast"); | ||
| 1672 | + ASSERT_NE(broadcast, nullptr); | ||
| 1673 | + EXPECT_EQ(broadcast->inputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 1674 | + EXPECT_EQ(broadcast->outputs[0].attr.dtype, af::DT_FLOAT16); | ||
| 1675 | +} | ||
| 1676 | + | ||
| 1677 | +TEST(BroadcastBackwardPass, MovesSharedDtypeAwareBranchesPastFollowingChain) { | ||
| 1678 | + // Not supported by the restored repository BRC implementation. | ||
| 1679 | + GTEST_SKIP(); | ||
| 1680 | + ScopedTestPlatform platform("3510"); | ||
| 1681 | + auto graph = BuildSharedDtypeAwareFanOutGraph("broadcast_backward_shared_dtype_aware_fanout"); | ||
| 1682 | + CompleteSharedDtypeAwareFanOutGraph(graph); | ||
| 1683 | + | ||
| 1684 | + optimize::BroadcastBackwardPass pass; | ||
| 1685 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1686 | + EXPECT_TRUE(IsConnected(graph, "load", "abs")); | ||
| 1687 | + EXPECT_TRUE(IsConnected(graph, "abs", "left_cast")); | ||
| 1688 | + EXPECT_TRUE(IsConnected(graph, "load", "right_cast")); | ||
| 1689 | + EXPECT_TRUE(IsConnected(graph, "right_cast", "relu")); | ||
| 1690 | + EXPECT_TRUE(IsConnected(graph, "relu", "add")); | ||
| 1691 | + EXPECT_TRUE(IsConnected(graph, "left_cast", "add")); | ||
| 1692 | + EXPECT_TRUE(IsConnected(graph, "add", "sqrt")); | ||
| 1693 | + EXPECT_TRUE(IsConnected(graph, "sqrt", "sigmoid")); | ||
| 1694 | + const auto store = FindNode(graph, "store"); | ||
| 1695 | + ASSERT_NE(store, nullptr); | ||
| 1696 | + const auto broadcast = std::dynamic_pointer_cast<af::AscNode>(store->GetInDataNodes().at(0)); | ||
| 1697 | + ASSERT_NE(broadcast, nullptr); | ||
| 1698 | + EXPECT_EQ(broadcast->GetType(), af::ascir_op::Broadcast::Type); | ||
| 1699 | + EXPECT_EQ(broadcast->outputs[0].attr.dtype, af::DT_FLOAT); | ||
| 1700 | + EXPECT_EQ(broadcast->GetInDataNodes().at(0)->GetName(), "sigmoid"); | ||
| 1701 | +} | ||
| 1702 | + | ||
| 1703 | +TEST(BroadcastBackwardPass, SkipsSharedDtypeAwareBranchWithMismatchedInputDtype) { | ||
| 1704 | + // Not supported by the restored repository BRC implementation. | ||
| 1705 | + GTEST_SKIP(); | ||
| 1706 | + ScopedTestPlatform platform("3510"); | ||
| 1707 | + auto graph = BuildSharedDtypeAwareFanOutGraph("broadcast_backward_shared_dtype_aware_mismatch"); | ||
| 1708 | + CompleteSharedDtypeAwareFanOutGraph(graph); | ||
| 1709 | + const auto right_cast = FindNode(graph, "right_cast"); | ||
| 1710 | + ASSERT_NE(right_cast, nullptr); | ||
| 1711 | + right_cast->GetOpDesc()->MutableInputDesc(0U)->SetDataType(af::DT_FLOAT16); | ||
| 1712 | + | ||
| 1713 | + optimize::BroadcastBackwardPass pass; | ||
| 1714 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1715 | + EXPECT_TRUE(IsConnected(graph, "load", "broadcast")); | ||
| 1716 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "abs")); | ||
| 1717 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "right_cast")); | ||
| 1718 | + EXPECT_FALSE(IsConnected(graph, "sigmoid", "broadcast")); | ||
| 1719 | +} | ||
| 1720 | + | ||
| 1721 | +TEST(BroadcastBackwardPass, MovesSameConsumerMultipleInputs) { | ||
| 1722 | + const auto s0 = Sym("s0"); | ||
| 1723 | + const auto s1 = Sym("s1"); | ||
| 1724 | + auto graph = AscGraphBuilder("broadcast_backward_multi_reference_same_consumer") | ||
| 1725 | + .Loops({s0, s1}) | ||
| 1726 | + .Data("data", 0) | ||
| 1727 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1728 | + .Broadcast("broadcast", "load", {1}) | ||
| 1729 | + .Add("consumer", "broadcast", "broadcast") | ||
| 1730 | + .Store("store", "consumer") | ||
| 1731 | + .Output("output", "store") | ||
| 1732 | + .Build(); | ||
| 1733 | + CompleteApiInfo(graph); | ||
| 1734 | + | ||
| 1735 | + optimize::BroadcastBackwardPass pass; | ||
| 1736 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1737 | + const auto load = FindNode(graph, "load"); | ||
| 1738 | + const auto broadcast = FindNode(graph, "broadcast"); | ||
| 1739 | + const auto consumer = FindNode(graph, "consumer"); | ||
| 1740 | + ASSERT_NE(load, nullptr); | ||
| 1741 | + ASSERT_NE(broadcast, nullptr); | ||
| 1742 | + ASSERT_NE(consumer, nullptr); | ||
| 1743 | + EXPECT_EQ(consumer->GetInDataAnchor(0)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "load"); | ||
| 1744 | + EXPECT_EQ(consumer->GetInDataAnchor(1)->GetPeerOutAnchor()->GetOwnerNode()->GetName(), "load"); | ||
| 1745 | + EXPECT_TRUE(IsConnected(graph, "consumer", "broadcast")); | ||
| 1746 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "store")); | ||
| 1747 | + ExpectStaticEq(consumer->inputs[0].attr.repeats, kCompactRepeats); | ||
| 1748 | + ExpectStaticEq(consumer->inputs[1].attr.repeats, kCompactRepeats); | ||
| 1749 | + ExpectStaticEq(consumer->outputs[0].attr.repeats, kCompactRepeats); | ||
| 1750 | + ExpectStaticEq(broadcast->inputs[0].attr.repeats, kCompactRepeats); | ||
| 1751 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "consumer", "broadcast")); | ||
| 1752 | + EXPECT_TRUE(IsEdgeAttrConsistent(graph, "broadcast", "store")); | ||
| 1753 | +} | ||
| 1754 | + | ||
| 1755 | +TEST(BroadcastBackwardPass, MovesSameConsumerWithThreeDimensionalLayout) { | ||
| 1756 | + const std::vector<af::Expression> compact_repeats = {Sym(83), af::sym::kSymbolOne, Sym(91)}; | ||
| 1757 | + const std::vector<af::Expression> compact_strides = {Sym(91), af::sym::kSymbolZero, af::sym::kSymbolOne}; | ||
| 1758 | + const std::vector<af::Expression> expanded_repeats = {Sym(83), Sym(18), Sym(91)}; | ||
| 1759 | + const std::vector<af::Expression> expanded_strides = {Sym(1638), Sym(91), af::sym::kSymbolOne}; | ||
| 1760 | + auto graph = AscGraphBuilder("broadcast_backward_multi_reference_same_consumer_3d") | ||
| 1761 | + .Loops({Sym(83), Sym(18), Sym(91)}) | ||
| 1762 | + .Data("data", 0) | ||
| 1763 | + .Load("load", "data", compact_repeats, compact_strides) | ||
| 1764 | + .Broadcast("broadcast", "load", {1}) | ||
| 1765 | + .Add("consumer", "broadcast", "broadcast") | ||
| 1766 | + .Store("store", "consumer") | ||
| 1767 | + .Output("output", "store") | ||
| 1768 | + .Build(); | ||
| 1769 | + CompleteApiInfo(graph); | ||
| 1770 | + | ||
| 1771 | + optimize::BroadcastBackwardPass pass; | ||
| 1772 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1773 | + const auto consumer = FindNode(graph, "consumer"); | ||
| 1774 | + const auto broadcast = FindNode(graph, "broadcast"); | ||
| 1775 | + ASSERT_NE(consumer, nullptr); | ||
| 1776 | + ASSERT_NE(broadcast, nullptr); | ||
| 1777 | + ExpectStaticEq(consumer->inputs[0].attr.repeats, compact_repeats); | ||
| 1778 | + ExpectStaticEq(consumer->inputs[0].attr.strides, compact_strides); | ||
| 1779 | + ExpectStaticEq(consumer->inputs[1].attr.repeats, compact_repeats); | ||
| 1780 | + ExpectStaticEq(consumer->inputs[1].attr.strides, compact_strides); | ||
| 1781 | + ExpectStaticEq(consumer->outputs[0].attr.repeats, compact_repeats); | ||
| 1782 | + ExpectStaticEq(consumer->outputs[0].attr.strides, compact_strides); | ||
| 1783 | + ExpectStaticEq(broadcast->inputs[0].attr.repeats, compact_repeats); | ||
| 1784 | + ExpectStaticEq(broadcast->inputs[0].attr.strides, compact_strides); | ||
| 1785 | + ExpectStaticEq(broadcast->outputs[0].attr.repeats, expanded_repeats); | ||
| 1786 | + ExpectStaticEq(broadcast->outputs[0].attr.strides, expanded_strides); | ||
| 1787 | +} | ||
| 1788 | + | ||
| 1789 | +TEST(BroadcastBackwardPass, MultiReferenceBackwardSkipsBarrierBranch) { | ||
| 1790 | + // Not supported by the restored repository BRC implementation. | ||
| 1791 | + GTEST_SKIP(); | ||
| 1792 | + const auto s0 = Sym("s0"); | ||
| 1793 | + const auto s1 = Sym("s1"); | ||
| 1794 | + auto graph = AscGraphBuilder("broadcast_backward_multi_reference_barrier") | ||
| 1795 | + .Loops({s0, s1}) | ||
| 1796 | + .Data("data", 0) | ||
| 1797 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1798 | + .Broadcast("broadcast", "load", {1}) | ||
| 1799 | + .Cast("barrier", "broadcast", af::DT_FLOAT16) | ||
| 1800 | + .Abs("branch", "broadcast") | ||
| 1801 | + .Add("merge", "barrier", "branch") | ||
| 1802 | + .Store("store", "merge") | ||
| 1803 | + .Output("output", "store") | ||
| 1804 | + .Build(); | ||
| 1805 | + CompleteApiInfo(graph); | ||
| 1806 | + | ||
| 1807 | + optimize::BroadcastBackwardPass pass; | ||
| 1808 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1809 | + EXPECT_TRUE(IsConnected(graph, "load", "broadcast")); | ||
| 1810 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "barrier")); | ||
| 1811 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "branch")); | ||
| 1812 | + EXPECT_FALSE(IsConnected(graph, "merge", "broadcast")); | ||
| 1813 | +} | ||
| 1814 | + | ||
| 1815 | +TEST(BroadcastBackwardPass, MultiReferenceBackwardSkipsBroadcastWithControlEdge) { | ||
| 1816 | + // Not supported by the restored repository BRC implementation. | ||
| 1817 | + GTEST_SKIP(); | ||
| 1818 | + auto graph = BuildDirectFanOutGraph("broadcast_backward_multi_reference_control_edge"); | ||
| 1819 | + CompleteApiInfo(graph); | ||
| 1820 | + const auto load = FindNode(graph, "load"); | ||
| 1821 | + const auto broadcast = FindNode(graph, "broadcast"); | ||
| 1822 | + ASSERT_NE(load, nullptr); | ||
| 1823 | + ASSERT_NE(broadcast, nullptr); | ||
| 1824 | + ASSERT_EQ(af::GraphUtils::AddEdge(load->GetOutControlAnchor(), broadcast->GetInControlAnchor()), af::SUCCESS); | ||
| 1825 | + | ||
| 1826 | + optimize::BroadcastBackwardPass pass; | ||
| 1827 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1828 | + EXPECT_TRUE(IsConnected(graph, "load", "broadcast")); | ||
| 1829 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "branch0")); | ||
| 1830 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "branch1")); | ||
| 1831 | + EXPECT_FALSE(IsConnected(graph, "merge", "broadcast")); | ||
| 1832 | +} | ||
| 1833 | + | ||
| 1834 | +TEST(BroadcastBackwardPass, MultiReferenceBackwardSkipsMultiInputSuccessor) { | ||
| 1835 | + // Not supported by the restored repository BRC implementation. | ||
| 1836 | + GTEST_SKIP(); | ||
| 1837 | + const auto s0 = Sym("s0"); | ||
| 1838 | + const auto s1 = Sym("s1"); | ||
| 1839 | + auto graph = AscGraphBuilder("broadcast_backward_multi_reference_multi_input_successor") | ||
| 1840 | + .Loops({s0, s1}) | ||
| 1841 | + .Data("data", 0) | ||
| 1842 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1843 | + .Broadcast("broadcast", "load", {1}) | ||
| 1844 | + .Abs("branch0", "broadcast") | ||
| 1845 | + .Neg("branch1", "broadcast") | ||
| 1846 | + .Add("merge", "branch0", "branch1") | ||
| 1847 | + .Add("succ", "merge", "merge") | ||
| 1848 | + .Store("store", "succ") | ||
| 1849 | + .Output("output", "store") | ||
| 1850 | + .Build(); | ||
| 1851 | + CompleteApiInfo(graph); | ||
| 1852 | + | ||
| 1853 | + optimize::BroadcastBackwardPass pass; | ||
| 1854 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1855 | + EXPECT_TRUE(IsConnected(graph, "load", "broadcast")); | ||
| 1856 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "branch0")); | ||
| 1857 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "branch1")); | ||
| 1858 | + EXPECT_FALSE(IsConnected(graph, "merge", "broadcast")); | ||
| 1859 | +} | ||
| 1860 | + | ||
| 1861 | +TEST(BroadcastBackwardPass, MultiReferenceBackwardSkipsSameConsumerWithAnotherSource) { | ||
| 1862 | + // Not supported by the restored repository BRC implementation. | ||
| 1863 | + GTEST_SKIP(); | ||
| 1864 | + const auto s0 = Sym("s0"); | ||
| 1865 | + const auto s1 = Sym("s1"); | ||
| 1866 | + auto graph = AscGraphBuilder("broadcast_backward_multi_reference_mixed_consumer") | ||
| 1867 | + .Loops({s0, s1}) | ||
| 1868 | + .Data("data", 0) | ||
| 1869 | + .Load("load", "data", kCompactRepeats, kCompactStrides) | ||
| 1870 | + .Broadcast("broadcast", "load", {1}) | ||
| 1871 | + .Abs("other", "broadcast") | ||
| 1872 | + .Add("consumer", "broadcast", "other") | ||
| 1873 | + .Store("store", "consumer") | ||
| 1874 | + .Output("output", "store") | ||
| 1875 | + .Build(); | ||
| 1876 | + CompleteApiInfo(graph); | ||
| 1877 | + | ||
| 1878 | + optimize::BroadcastBackwardPass pass; | ||
| 1879 | + ASSERT_EQ(pass.RunPass(graph), af::SUCCESS); | ||
| 1880 | + EXPECT_TRUE(IsConnected(graph, "load", "broadcast")); | ||
| 1881 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "other")); | ||
| 1882 | + EXPECT_TRUE(IsConnected(graph, "broadcast", "consumer")); | ||
| 1883 | + EXPECT_FALSE(IsConnected(graph, "consumer", "broadcast")); | ||
| 1884 | +} | ||
| @@ -3370,12 +3370,19 @@ TEST_F(TestOptimizer, ScalarBroadcastOptimization_Two_Scalar) { | |||
| 3370 | EXPECT_EQ(res, af::SUCCESS); | 3370 | EXPECT_EQ(res, af::SUCCESS); |
| 3371 | auto compute_graph = af::AscGraphUtils::GetComputeGraph(graph); | 3371 | auto compute_graph = af::AscGraphUtils::GetComputeGraph(graph); |
| 3372 | EXPECT_EQ(compute_graph->GetAllNodesSize(), 10); | 3372 | EXPECT_EQ(compute_graph->GetAllNodesSize(), 10); |
| 3373 | - EXPECT_EQ(compute_graph->FindNode("brc1"), nullptr); | 3373 | + const auto retained_brc1 = compute_graph->FindNode("brc1"); |
| 3374 | - EXPECT_EQ(compute_graph->FindNode("brc2"), nullptr); | 3374 | + const auto retained_brc2 = compute_graph->FindNode("brc2"); |
| 3375 | - EXPECT_EQ(compute_graph->FindNode("brc3"), nullptr); | 3375 | + const auto retained_brc3 = compute_graph->FindNode("brc3"); |
| 3376 | - EXPECT_NE(compute_graph->FindNode("brc4"), nullptr); | 3376 | + ASSERT_NE(retained_brc1, nullptr); |
| 3377 | - EXPECT_NE(compute_graph->FindNode("brc5"), nullptr); | 3377 | + ASSERT_NE(retained_brc2, nullptr); |
| 3378 | - EXPECT_NE(compute_graph->FindNode("brc6"), nullptr); | 3378 | + ASSERT_NE(retained_brc3, nullptr); |
| 3379 | + EXPECT_EQ(compute_graph->FindNode("brc4"), nullptr); | ||
| 3380 | + EXPECT_EQ(compute_graph->FindNode("brc5"), nullptr); | ||
| 3381 | + EXPECT_EQ(compute_graph->FindNode("brc6"), nullptr); | ||
| 3382 | + EXPECT_EQ(retained_brc1->GetInDataNodes().at(0)->GetName(), "add"); | ||
| 3383 | + EXPECT_EQ(retained_brc2->GetInDataNodes().at(0)->GetName(), "brc1"); | ||
| 3384 | + EXPECT_EQ(retained_brc3->GetInDataNodes().at(0)->GetName(), "brc2"); | ||
| 3385 | + EXPECT_EQ(compute_graph->FindNode("store")->GetInDataNodes().at(0)->GetName(), "brc3"); | ||
| 3379 | } | 3386 | } |
| 3380 | 3387 | ||
| 3381 | TEST_F(TestOptimizer, ScalarBroadcastOptimization_Same_Input) { | 3388 | TEST_F(TestOptimizer, ScalarBroadcastOptimization_Same_Input) { |
| @@ -406,25 +406,32 @@ 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 | - .Exp("exp0", "broadcast") | 409 | + .Scalar("scalar0", "0", af::DT_FLOAT) |
| 410 | - .Abs("abs0", "broadcast") | 410 | + .Add("exp0", "broadcast", "scalar0") |
| 411 | - .Mul("mul0", "exp0", "abs0") | 411 | + .Abs("abs0", "exp0") |
| 412 | - .Store("store", "mul0") | 412 | + .Store("store", "abs0") |
| 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 | - EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); | 417 | + ASSERT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); |
| 418 | - | 418 | + ASSERT_FALSE(fused_scheduled_result.node_idx_to_scheduled_results.empty()); |
| 419 | - for (const auto &node : | 419 | + ASSERT_FALSE(fused_scheduled_result.node_idx_to_scheduled_results[0].empty()); |
| 420 | - fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs[1].GetAllNodes()) { | 420 | + ASSERT_FALSE(fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups.empty()); |
| 421 | - if (node->GetOpDesc()->GetId() == 1) { | 421 | + const auto &impl_graphs = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0].impl_graphs; |
| 422 | - EXPECT_EQ(node->GetOpDesc()->GetType(), "Nddma"); | 422 | + ASSERT_GT(impl_graphs.size(), 1UL); |
| 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; | ||
| 423 | } | 428 | } |
| 424 | - if (node->GetOpDesc()->GetId() == 2) { | 429 | + if (node->GetOpDesc()->GetType() == "VectorFunc") { |
| 425 | - EXPECT_EQ(node->GetOpDesc()->GetType(), "VectorFunc"); | 430 | + has_vector_func = true; |
| 426 | } | 431 | } |
| 427 | } | 432 | } |
| 433 | + EXPECT_TRUE(has_nddma); | ||
| 434 | + EXPECT_TRUE(has_vector_func); | ||
| 428 | } | 435 | } |
| 429 | 436 | ||
| 430 | TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) { | 437 | TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) { |
| @@ -456,16 +463,24 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) { | |||
| 456 | *broadcast0.y.repeats = {s0, s1}; | 463 | *broadcast0.y.repeats = {s0, s1}; |
| 457 | *broadcast0.y.strides = {s1, af::ops::One}; | 464 | *broadcast0.y.strides = {s1, af::ops::One}; |
| 458 | 465 | ||
| 459 | - Exp exp0("exp0"); | 466 | + Scalar scalar0("scalar0", graph); |
| 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"); | ||
| 460 | exp0.attr.sched.axis = {z0.id, z1.id}; | 474 | exp0.attr.sched.axis = {z0.id, z1.id}; |
| 461 | - exp0.x = broadcast0.y; | 475 | + exp0.x1 = broadcast0.y; |
| 476 | + exp0.x2 = scalar0.y; | ||
| 462 | *exp0.y.axis = {z0.id, z1.id}; | 477 | *exp0.y.axis = {z0.id, z1.id}; |
| 463 | exp0.y.dtype = dtype; | 478 | exp0.y.dtype = dtype; |
| 464 | *exp0.y.repeats = {s0, s1}; | 479 | *exp0.y.repeats = {s0, s1}; |
| 465 | *exp0.y.strides = {s1, af::ops::One}; | 480 | *exp0.y.strides = {s1, af::ops::One}; |
| 466 | 481 | ||
| 467 | Abs abs0("abs0"); | 482 | Abs abs0("abs0"); |
| 468 | - abs0.x = broadcast0.y; | 483 | + abs0.x = exp0.y; |
| 469 | abs0.attr.sched.axis = {z0.id, z1.id}; | 484 | abs0.attr.sched.axis = {z0.id, z1.id}; |
| 470 | abs0.y.dtype = dtype; | 485 | abs0.y.dtype = dtype; |
| 471 | *abs0.y.axis = {z0.id, z1.id}; | 486 | *abs0.y.axis = {z0.id, z1.id}; |
| @@ -473,18 +488,9 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) { | |||
| 473 | *abs0.y.strides = {s1, One}; | 488 | *abs0.y.strides = {s1, One}; |
| 474 | abs0.attr.api.compute_type = ComputeType::kComputeElewise; | 489 | abs0.attr.api.compute_type = ComputeType::kComputeElewise; |
| 475 | 490 | ||
| 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 | - | ||
| 485 | Store store_op("store"); | 491 | Store store_op("store"); |
| 486 | store_op.attr.sched.axis = {z0.id, z1.id}; | 492 | store_op.attr.sched.axis = {z0.id, z1.id}; |
| 487 | - store_op.x = mul0.y; | 493 | + store_op.x = abs0.y; |
| 488 | *store_op.y.axis = {z0.id, z1.id}; | 494 | *store_op.y.axis = {z0.id, z1.id}; |
| 489 | store_op.y.dtype = dtype; | 495 | store_op.y.dtype = dtype; |
| 490 | *store_op.y.strides = {s1, af::ops::One}; | 496 | *store_op.y.strides = {s1, af::ops::One}; |
| @@ -499,15 +505,8 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc) { | |||
| 499 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); | 505 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); |
| 500 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; | 506 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; |
| 501 | 507 | ||
| 502 | - ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2); | 508 | + ASSERT_FALSE(schedule_group.impl_graphs.empty()); |
| 503 | - | 509 | + EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL); |
| 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); | ||
| 511 | } | 510 | } |
| 512 | 511 | ||
| 513 | TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) { | 512 | TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) { |
| @@ -539,16 +538,24 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) { | |||
| 539 | *broadcast0.y.repeats = {s0, s1}; | 538 | *broadcast0.y.repeats = {s0, s1}; |
| 540 | *broadcast0.y.strides = {s1, af::ops::One}; | 539 | *broadcast0.y.strides = {s1, af::ops::One}; |
| 541 | 540 | ||
| 542 | - Exp exp0("exp0"); | 541 | + Scalar scalar0("scalar0", graph); |
| 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"); | ||
| 543 | exp0.attr.sched.axis = {z0.id, z1.id}; | 549 | exp0.attr.sched.axis = {z0.id, z1.id}; |
| 544 | - exp0.x = broadcast0.y; | 550 | + exp0.x1 = broadcast0.y; |
| 551 | + exp0.x2 = scalar0.y; | ||
| 545 | *exp0.y.axis = {z0.id, z1.id}; | 552 | *exp0.y.axis = {z0.id, z1.id}; |
| 546 | exp0.y.dtype = dtype; | 553 | exp0.y.dtype = dtype; |
| 547 | *exp0.y.repeats = {s0, s1}; | 554 | *exp0.y.repeats = {s0, s1}; |
| 548 | *exp0.y.strides = {s1, af::ops::One}; | 555 | *exp0.y.strides = {s1, af::ops::One}; |
| 549 | 556 | ||
| 550 | Abs abs0("abs0"); | 557 | Abs abs0("abs0"); |
| 551 | - abs0.x = broadcast0.y; | 558 | + abs0.x = exp0.y; |
| 552 | abs0.attr.sched.axis = {z0.id, z1.id}; | 559 | abs0.attr.sched.axis = {z0.id, z1.id}; |
| 553 | abs0.y.dtype = dtype; | 560 | abs0.y.dtype = dtype; |
| 554 | *abs0.y.axis = {z0.id, z1.id}; | 561 | *abs0.y.axis = {z0.id, z1.id}; |
| @@ -556,18 +563,9 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) { | |||
| 556 | *abs0.y.strides = {s1, One}; | 563 | *abs0.y.strides = {s1, One}; |
| 557 | abs0.attr.api.compute_type = ComputeType::kComputeElewise; | 564 | abs0.attr.api.compute_type = ComputeType::kComputeElewise; |
| 558 | 565 | ||
| 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 | - | ||
| 568 | Store store_op("store"); | 566 | Store store_op("store"); |
| 569 | store_op.attr.sched.axis = {z0.id, z1.id}; | 567 | store_op.attr.sched.axis = {z0.id, z1.id}; |
| 570 | - store_op.x = mul0.y; | 568 | + store_op.x = abs0.y; |
| 571 | *store_op.y.axis = {z0.id, z1.id}; | 569 | *store_op.y.axis = {z0.id, z1.id}; |
| 572 | store_op.y.dtype = dtype; | 570 | store_op.y.dtype = dtype; |
| 573 | *store_op.y.strides = {s1, af::ops::One}; | 571 | *store_op.y.strides = {s1, af::ops::One}; |
| @@ -582,17 +580,8 @@ TEST_F(OptimizerStV2, NddmaCaseAlignTailBrcScoreFunc_Dynamic) { | |||
| 582 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); | 580 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); |
| 583 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; | 581 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; |
| 584 | 582 | ||
| 585 | - ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2); | 583 | + ASSERT_FALSE(schedule_group.impl_graphs.empty()); |
| 586 | - const auto score_func_iter = schedule_group.graph_name_to_score_funcs.find(schedule_group.impl_graphs[2].GetName()); | 584 | + EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL); |
| 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); | ||
| 596 | } | 585 | } |
| 597 | 586 | ||
| 598 | TEST_F(OptimizerStV2, NddmaCaseLargeTailBrcScoreFunc) { | 587 | TEST_F(OptimizerStV2, NddmaCaseLargeTailBrcScoreFunc) { |
| @@ -608,25 +597,18 @@ TEST_F(OptimizerStV2, NddmaCaseLargeTailBrcScoreFunc) { | |||
| 608 | .Data("data0", 0, af::DT_FLOAT) | 597 | .Data("data0", 0, af::DT_FLOAT) |
| 609 | .Load("load0", "data0", load0_shape, load0_strides) | 598 | .Load("load0", "data0", load0_shape, load0_strides) |
| 610 | .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1 | 599 | .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1 |
| 611 | - .Exp("exp0", "broadcast") | 600 | + .Scalar("scalar0", "0", af::DT_FLOAT) |
| 612 | - .Abs("abs0", "broadcast") | 601 | + .Add("exp0", "broadcast", "scalar0") |
| 613 | - .Mul("mul0", "exp0", "abs0") | 602 | + .Abs("abs0", "exp0") |
| 614 | - .Store("store", "mul0") | 603 | + .Store("store", "abs0") |
| 615 | .Output("output", "store", 8, af::DT_FLOAT) | 604 | .Output("output", "store", 8, af::DT_FLOAT) |
| 616 | .Build(); | 605 | .Build(); |
| 617 | 606 | ||
| 618 | ::ascir::FusedScheduledResult fused_scheduled_result; | 607 | ::ascir::FusedScheduledResult fused_scheduled_result; |
| 619 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); | 608 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); |
| 620 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; | 609 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; |
| 621 | - ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2); | 610 | + ASSERT_FALSE(schedule_group.impl_graphs.empty()); |
| 622 | - | 611 | + EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL); |
| 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); | ||
| 630 | } | 612 | } |
| 631 | 613 | ||
| 632 | TEST_F(OptimizerStV2, NddmaCaseLargeTailBrc_Dynamic) { | 614 | TEST_F(OptimizerStV2, NddmaCaseLargeTailBrc_Dynamic) { |
| @@ -642,17 +624,18 @@ TEST_F(OptimizerStV2, NddmaCaseLargeTailBrc_Dynamic) { | |||
| 642 | .Data("data0", 0, af::DT_FLOAT) | 624 | .Data("data0", 0, af::DT_FLOAT) |
| 643 | .Load("load0", "data0", load0_shape, load0_strides) | 625 | .Load("load0", "data0", load0_shape, load0_strides) |
| 644 | .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1 | 626 | .Broadcast("broadcast", "load0", {1}) // broadcast on axis 1 |
| 645 | - .Exp("exp0", "broadcast") | 627 | + .Scalar("scalar0", "0", af::DT_FLOAT) |
| 646 | - .Abs("abs0", "broadcast") | 628 | + .Add("exp0", "broadcast", "scalar0") |
| 647 | - .Mul("mul0", "exp0", "abs0") | 629 | + .Abs("abs0", "exp0") |
| 648 | - .Store("store", "mul0") | 630 | + .Store("store", "abs0") |
| 649 | .Output("output", "store", 8, af::DT_FLOAT) | 631 | .Output("output", "store", 8, af::DT_FLOAT) |
| 650 | .Build(); | 632 | .Build(); |
| 651 | 633 | ||
| 652 | ::ascir::FusedScheduledResult fused_scheduled_result; | 634 | ::ascir::FusedScheduledResult fused_scheduled_result; |
| 653 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); | 635 | EXPECT_EQ(optimizer.Optimize(graph, fused_scheduled_result), 0); |
| 654 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; | 636 | const auto schedule_group = fused_scheduled_result.node_idx_to_scheduled_results[0][0].schedule_groups[0]; |
| 655 | - ASSERT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2); | 637 | + ASSERT_FALSE(schedule_group.impl_graphs.empty()); |
| 638 | + EXPECT_EQ(schedule_group.graph_name_to_score_funcs.size(), 2UL); | ||
| 656 | } | 639 | } |
| 657 | 640 | ||
| 658 | /** | 641 | /** |
| @@ -64,14 +64,19 @@ 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"); | ||
| 68 | const auto add = graph.FindNode("add"); | 67 | const auto add = graph.FindNode("add"); |
| 69 | - ASSERT_NE(canonical, nullptr); | ||
| 70 | ASSERT_NE(add, nullptr); | 68 | ASSERT_NE(add, nullptr); |
| 71 | EXPECT_EQ(graph.FindNode("broadcast1"), nullptr); | 69 | EXPECT_EQ(graph.FindNode("broadcast1"), nullptr); |
| 72 | - EXPECT_EQ(add->GetInDataAnchor(0)->GetPeerOutAnchor(), canonical->GetOutDataAnchor(0)); | 70 | + |
| 73 | - EXPECT_EQ(add->GetInDataAnchor(1)->GetPeerOutAnchor(), canonical->GetOutDataAnchor(0)); | 71 | + const auto input0_peer = add->GetInDataAnchor(0)->GetPeerOutAnchor(); |
| 74 | - EXPECT_EQ(reduce->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 1UL); | 72 | + const auto input1_peer = add->GetInDataAnchor(1)->GetPeerOutAnchor(); |
| 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); | ||
| 75 | } | 80 | } |
| 76 | 81 | ||
| 77 | TEST_F(SameSourceBroadcastCseStTest, SkipsGraphWithoutNormStructureThroughGraphPassRunner) { | 82 | 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(), 2UL); | 1500 | + EXPECT_EQ(asc_graphs.size(), 3UL); |
| 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,6 +12,7 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | + | ||
| 15 | 16 | ||
| 16 | 17 | ||
| 17 | 18 | ||
| @@ -31,11 +32,12 @@ class PassRunnerV2 final : public BasePassRunner { | |||
| 31 | this->RegisterPass<PowEquivSubstitutionPass>(); | 32 | this->RegisterPass<PowEquivSubstitutionPass>(); |
| 32 | this->RegisterPass<BroadcastConstToStorePass>(); | 33 | this->RegisterPass<BroadcastConstToStorePass>(); |
| 33 | this->RegisterPass<ScalarTo1DTensorPass>(); | 34 | this->RegisterPass<ScalarTo1DTensorPass>(); |
| 35 | + this->RegisterPass<SameSourceBroadcastCsePass>(); | ||
| 36 | + this->RegisterPass<BroadcastBackwardPass>(); | ||
| 34 | this->RegisterPass<ScalarBroadcastOptimizationPass>(); | 37 | this->RegisterPass<ScalarBroadcastOptimizationPass>(); |
| 35 | this->RegisterPass<MaskedFillInputReorderPass>(); | 38 | this->RegisterPass<MaskedFillInputReorderPass>(); |
| 36 | this->RegisterPass<ExpandDimsForAllReducePass>(); | 39 | this->RegisterPass<ExpandDimsForAllReducePass>(); |
| 37 | this->RegisterPass<ContinuesBroadcastOptimizationPass>(); | 40 | this->RegisterPass<ContinuesBroadcastOptimizationPass>(); |
| 38 | - this->RegisterPass<SameSourceBroadcastCsePass>(); | ||
| 39 | this->RegisterPass<DuplicateElewiseCsePass>(); | 41 | this->RegisterPass<DuplicateElewiseCsePass>(); |
| 40 | this->RegisterPass<GatherToLoadPass>(); | 42 | this->RegisterPass<GatherToLoadPass>(); |
| 41 | this->RegisterPass<SplitConcatOptimizationPass>(); | 43 | this->RegisterPass<SplitConcatOptimizationPass>(); |