已合并
【PR】: 重构,使用InsertNode收编节点插入行为 #1677
shengnan创建于 8月7日
【PR】: 重构,使用InsertNode收编节点插入行为 #1677
已合并
共 13 个文件变更+359-288
| @@ -17,6 +17,15 @@ | |||
| 17 | 17 | ||
| 18 | namespace af { | 18 | namespace af { |
| 19 | namespace { | 19 | namespace { |
| 20 | +graphStatus InheritAutofuseAttr(GeTensorDesc &src_tensor_desc, GeTensorDesc &insert_tensor_desc) { | ||
| 21 | + auto src_tensor_attr = src_tensor_desc.GetOrCreateAttrsGroup<AscTensorAttr>(); | ||
| 22 | + auto insert_tensor_attr = insert_tensor_desc.GetOrCreateAttrsGroup<AscTensorAttr>(); | ||
| 23 | + insert_tensor_attr->axis = src_tensor_attr->axis; | ||
| 24 | + insert_tensor_attr->repeats = src_tensor_attr->repeats; | ||
| 25 | + insert_tensor_attr->strides = src_tensor_attr->strides; | ||
| 26 | + return GRAPH_SUCCESS; | ||
| 27 | +} | ||
| 28 | + | ||
| 20 | graphStatus EstablishAscNodeAndEdges(const ascendc_ir::proto::AscGraphDef &asc_graph_def, AscGraph &out_asc_graph) { | 29 | graphStatus EstablishAscNodeAndEdges(const ascendc_ir::proto::AscGraphDef &asc_graph_def, AscGraph &out_asc_graph) { |
| 21 | auto &asc_nodes = asc_graph_def.asc_node(); | 30 | auto &asc_nodes = asc_graph_def.asc_node(); |
| 22 | // 1. Add AscNodes to AscGraph | 31 | // 1. Add AscNodes to AscGraph |
| @@ -110,6 +119,100 @@ graphStatus EstablishAscNodeAndEdges(const ascendc_ir::proto::AscGraphDef &asc_g | |||
| 110 | return GRAPH_SUCCESS; | 119 | return GRAPH_SUCCESS; |
| 111 | } | 120 | } |
| 112 | } // namespace | 121 | } // namespace |
| 122 | + | ||
| 123 | +NodePtr AscGraphUtils::InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 124 | + const OpDescPtr &insert_op, const uint32_t input_index, | ||
| 125 | + const uint32_t output_index) { | ||
| 126 | + GE_ASSERT_NOTNULL(src); | ||
| 127 | + const NodePtr src_node = src->GetOwnerNode(); | ||
| 128 | + GE_ASSERT_NOTNULL(src_node); | ||
| 129 | + auto compute_graph = src_node->GetOwnerComputeGraphBarePtr(); | ||
| 130 | + GE_ASSERT_NOTNULL(compute_graph); | ||
| 131 | + auto insert_node = compute_graph->InsertNode(src_node, insert_op); | ||
| 132 | + GE_ASSERT_GRAPH_SUCCESS(InsertNodeAfter(src, dsts, insert_node, input_index, output_index)); | ||
| 133 | + return insert_node; | ||
| 134 | +} | ||
| 135 | + | ||
| 136 | +graphStatus AscGraphUtils::InsertNodeAfter(const OutDataAnchorPtr &src, const NodePtr &insert_node, | ||
| 137 | + const uint32_t input_index, const uint32_t output_index) { | ||
| 138 | + GE_CHECK_NOTNULL(src); | ||
| 139 | + const auto peer_in_anchor_range = src->GetPeerInDataAnchors(); | ||
| 140 | + const std::vector<InDataAnchorPtr> peer_in_anchors(peer_in_anchor_range.begin(), peer_in_anchor_range.end()); | ||
| 141 | + return InsertNodeAfter(src, peer_in_anchors, insert_node, input_index, output_index); | ||
| 142 | +} | ||
| 143 | + | ||
| 144 | +graphStatus AscGraphUtils::InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 145 | + const NodePtr &insert_node, const uint32_t input_index, | ||
| 146 | + const uint32_t output_index) { | ||
| 147 | + GE_CHECK_NOTNULL(src); | ||
| 148 | + GE_CHECK_NOTNULL(insert_node); | ||
| 149 | + | ||
| 150 | + const auto src_node = src->GetOwnerNodeBarePtr(); | ||
| 151 | + GE_CHECK_NOTNULL(src_node); | ||
| 152 | + GE_ASSERT_TRUE(src_node->GetOwnerComputeGraph() == insert_node->GetOwnerComputeGraph(), | ||
| 153 | + "src:%s and insert_node:%s does not exist in the same graph.", src_node->GetName().c_str(), | ||
| 154 | + insert_node->GetName().c_str()); | ||
| 155 | + | ||
| 156 | + GE_ASSERT_GRAPH_SUCCESS(GraphUtils::AddEdge(src, insert_node->GetInDataAnchor(static_cast<int32_t>(input_index)))); | ||
| 157 | + | ||
| 158 | + for (auto &dst : dsts) { | ||
| 159 | + const auto dst_node = dst->GetOwnerNodeBarePtr(); | ||
| 160 | + GELOGI("Insert node(after) %s between %s->%s.", insert_node->GetName().c_str(), src_node->GetName().c_str(), | ||
| 161 | + dst_node->GetName().c_str()); | ||
| 162 | + GE_ASSERT_GRAPH_SUCCESS(GraphUtils::RemoveEdge(src, dst)); | ||
| 163 | + GE_ASSERT_GRAPH_SUCCESS( | ||
| 164 | + GraphUtils::AddEdge(insert_node->GetOutDataAnchor(static_cast<int32_t>(output_index)), dst)); | ||
| 165 | + } | ||
| 166 | + insert_node->GetOpDesc()->GetOrCreateAttrsGroup<AscNodeAttr>()->sched = | ||
| 167 | + src_node->GetOpDesc()->GetOrCreateAttrsGroup<AscNodeAttr>()->sched; | ||
| 168 | + auto src_tensor_desc = src_node->GetOpDesc()->MutableOutputDesc(src->GetIdx()); | ||
| 169 | + auto insert_tensor_desc = insert_node->GetOpDesc()->MutableOutputDesc(output_index); | ||
| 170 | + GE_ASSERT_GRAPH_SUCCESS(InheritAutofuseAttr(*src_tensor_desc, *insert_tensor_desc)); | ||
| 171 | + return GRAPH_SUCCESS; | ||
| 172 | +} | ||
| 173 | + | ||
| 174 | +NodePtr AscGraphUtils::InsertNodeBefore(const InDataAnchorPtr &dst, const OpDescPtr &insert_op, | ||
| 175 | + const uint32_t input_index, const uint32_t output_index) { | ||
| 176 | + GE_ASSERT_NOTNULL(dst); | ||
| 177 | + const auto src_node_out_anchor = dst->GetPeerOutAnchor(); | ||
| 178 | + GE_ASSERT_NOTNULL(src_node_out_anchor); | ||
| 179 | + const auto src_node = src_node_out_anchor->GetOwnerNode(); | ||
| 180 | + GE_ASSERT_NOTNULL(src_node); | ||
| 181 | + auto compute_graph = src_node->GetOwnerComputeGraphBarePtr(); | ||
| 182 | + GE_ASSERT_NOTNULL(compute_graph); | ||
| 183 | + auto insert_node = compute_graph->InsertNode(src_node, insert_op); | ||
| 184 | + GE_ASSERT_GRAPH_SUCCESS(InsertNodeBefore(dst, insert_node, input_index, output_index)); | ||
| 185 | + return insert_node; | ||
| 186 | +} | ||
| 187 | + | ||
| 188 | +graphStatus AscGraphUtils::InsertNodeBefore(const InDataAnchorPtr &dst, const NodePtr &insert_node, | ||
| 189 | + const uint32_t input_index, const uint32_t output_index) { | ||
| 190 | + GE_CHECK_NOTNULL(dst); | ||
| 191 | + GE_CHECK_NOTNULL(insert_node); | ||
| 192 | + const auto dst_node = dst->GetOwnerNodeBarePtr(); | ||
| 193 | + GE_CHECK_NOTNULL(dst_node); | ||
| 194 | + GE_ASSERT_TRUE(dst_node->GetOwnerComputeGraph() == insert_node->GetOwnerComputeGraph(), | ||
| 195 | + "dst:%s and insert_node:%s does not exist in the same graph.", dst_node->GetName().c_str(), | ||
| 196 | + insert_node->GetName().c_str()); | ||
| 197 | + | ||
| 198 | + const auto src_node_out_anchor = dst->GetPeerOutAnchor(); | ||
| 199 | + GE_CHECK_NOTNULL(src_node_out_anchor); | ||
| 200 | + const auto src_node = src_node_out_anchor->GetOwnerNodeBarePtr(); | ||
| 201 | + GE_CHECK_NOTNULL(src_node); | ||
| 202 | + GE_ASSERT_GRAPH_SUCCESS(GraphUtils::RemoveEdge(src_node_out_anchor, dst)); | ||
| 203 | + GE_ASSERT_GRAPH_SUCCESS( | ||
| 204 | + GraphUtils::AddEdge(src_node_out_anchor, insert_node->GetInDataAnchor(static_cast<int32_t>(input_index)))); | ||
| 205 | + GE_ASSERT_GRAPH_SUCCESS(GraphUtils::AddEdge(insert_node->GetOutDataAnchor(static_cast<int32_t>(output_index)), dst)); | ||
| 206 | + GELOGI("Insert node(before) %s between %s->%s", insert_node->GetName().c_str(), src_node->GetName().c_str(), | ||
| 207 | + dst_node->GetName().c_str()); | ||
| 208 | + insert_node->GetOpDesc()->GetOrCreateAttrsGroup<AscNodeAttr>()->sched = | ||
| 209 | + dst_node->GetOpDesc()->GetOrCreateAttrsGroup<AscNodeAttr>()->sched; | ||
| 210 | + auto src_tensor_desc = src_node->GetOpDesc()->MutableOutputDesc(src_node_out_anchor->GetIdx()); | ||
| 211 | + auto insert_tensor_desc = insert_node->GetOpDesc()->MutableOutputDesc(output_index); | ||
| 212 | + GE_ASSERT_GRAPH_SUCCESS(InheritAutofuseAttr(*src_tensor_desc, *insert_tensor_desc)); | ||
| 213 | + return GRAPH_SUCCESS; | ||
| 214 | +} | ||
| 215 | + | ||
| 113 | ComputeGraphPtr AscGraphUtils::GetComputeGraph(const AscGraph &asc_graph) { | 216 | ComputeGraphPtr AscGraphUtils::GetComputeGraph(const AscGraph &asc_graph) { |
| 114 | return asc_graph.impl_->GetComputeGraph(); | 217 | return asc_graph.impl_->GetComputeGraph(); |
| 115 | } | 218 | } |
| @@ -78,23 +78,6 @@ const uint32_t kSubgraphIndexOfPartitionedCall = 0U; | |||
| 78 | const std::set<std::string> kMergeInputSkipTypes{STREAMACTIVE, STREAMSWITCH, CONSTANT, CONSTANTOP}; | 78 | const std::set<std::string> kMergeInputSkipTypes{STREAMACTIVE, STREAMSWITCH, CONSTANT, CONSTANTOP}; |
| 79 | constexpr int32_t kInvalidStream = -1; | 79 | constexpr int32_t kInvalidStream = -1; |
| 80 | constexpr size_t kNoOpOptimizeThreshold = 1000UL; | 80 | constexpr size_t kNoOpOptimizeThreshold = 1000UL; |
| 81 | -const std::string kSuperKernelScope = "_super_kernel_scope"; | ||
| 82 | -const std::string kSuperKernelOptions = "_super_kernel_options"; | ||
| 83 | -const std::vector<std::string> kNecessaryStrAttrWhitelist = { | ||
| 84 | - public_attr::USER_STREAM_LABEL, public_attr::OP_AI_CORE_NUM, public_attr::OP_VECTOR_CORE_NUM, kSuperKernelScope, | ||
| 85 | - kSuperKernelOptions}; | ||
| 86 | - | ||
| 87 | -Status InheritAttr(const OpDescPtr &node_op_desc, const OpDescPtr &insert_op_desc) { | ||
| 88 | - GE_ASSERT_NOTNULL(node_op_desc); | ||
| 89 | - for (const auto &attr : kNecessaryStrAttrWhitelist) { | ||
| 90 | - const std::string *attr_val = AttrUtils::GetStr(node_op_desc, attr); | ||
| 91 | - if (attr_val != nullptr) { | ||
| 92 | - GE_ASSERT_NOTNULL(insert_op_desc); | ||
| 93 | - GE_ASSERT_TRUE(AttrUtils::SetStr(insert_op_desc, attr, *attr_val)); | ||
| 94 | - } | ||
| 95 | - } | ||
| 96 | - return SUCCESS; | ||
| 97 | -} | ||
| 98 | 81 | ||
| 99 | graphStatus ReLinkInputDataEdge(const NodePtr &input_node, const NodePtr &target_node) { | 82 | graphStatus ReLinkInputDataEdge(const NodePtr &input_node, const NodePtr &target_node) { |
| 100 | GE_ASSERT_TRUE(input_node->GetType() == DATA, "Input node: %s should be Data", input_node->GetNamePtr()); | 83 | GE_ASSERT_TRUE(input_node->GetType() == DATA, "Input node: %s should be Data", input_node->GetNamePtr()); |
| @@ -501,168 +484,6 @@ GraphUtils::RemoveNodesWithoutRelink(const ComputeGraphPtr &compute_graph, const | |||
| 501 | return af::GRAPH_SUCCESS; | 484 | return af::GRAPH_SUCCESS; |
| 502 | } | 485 | } |
| 503 | 486 | ||
| 504 | -GE_FUNC_DEV_VISIBILITY GE_FUNC_HOST_VISIBILITY NodePtr | ||
| 505 | -GraphUtils::InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 506 | - const OpDescPtr &insert_op, const uint32_t input_index, const uint32_t output_index) { | ||
| 507 | - GE_ASSERT_NOTNULL(src); | ||
| 508 | - const NodePtr src_node = src->GetOwnerNode(); | ||
| 509 | - GE_ASSERT_NOTNULL(src_node); | ||
| 510 | - auto compute_graph = src_node->GetOwnerComputeGraphBarePtr(); | ||
| 511 | - GE_ASSERT_NOTNULL(compute_graph); | ||
| 512 | - auto insert_node = compute_graph->InsertNode(src_node, insert_op); | ||
| 513 | - GE_ASSERT_GRAPH_SUCCESS(GraphUtils::InsertNodeAfter(src, dsts, insert_node, input_index, output_index)); | ||
| 514 | - return insert_node; | ||
| 515 | -} | ||
| 516 | - | ||
| 517 | -/// @brief Insert node: src->insert_node:input_index, insert_node:output_index->dst | ||
| 518 | -/// @param [in] src | ||
| 519 | -/// @param [in] dsts | ||
| 520 | -/// @param [in] insert_node | ||
| 521 | -/// @param [in] input_index | ||
| 522 | -/// @param [in] output_index | ||
| 523 | -/// @return graphStatus | ||
| 524 | -GE_FUNC_DEV_VISIBILITY GE_FUNC_HOST_VISIBILITY graphStatus | ||
| 525 | -GraphUtils::InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 526 | - const NodePtr &insert_node, const uint32_t input_index, const uint32_t output_index) { | ||
| 527 | - GE_CHECK_NOTNULL(src); | ||
| 528 | - GE_CHECK_NOTNULL(insert_node); | ||
| 529 | - | ||
| 530 | - const auto src_node = src->GetOwnerNodeBarePtr(); | ||
| 531 | - GE_CHECK_NOTNULL(src_node); | ||
| 532 | - if (src_node->GetOwnerComputeGraph() != insert_node->GetOwnerComputeGraph()) { | ||
| 533 | - REPORT_INNER_ERR_MSG("E18888", "src:%s and insert_node:%s does not exist in the same graph.", | ||
| 534 | - src_node->GetName().c_str(), insert_node->GetName().c_str()); | ||
| 535 | - GELOGE(af::GRAPH_FAILED, "[Check][Param] src:%s and insert_node:%s does not exist in the same graph.", | ||
| 536 | - src_node->GetName().c_str(), insert_node->GetName().c_str()); | ||
| 537 | - return af::GRAPH_FAILED; | ||
| 538 | - } | ||
| 539 | - | ||
| 540 | - if (AddEdge(src, insert_node->GetInDataAnchor(static_cast<int32_t>(input_index))) != af::GRAPH_SUCCESS) { | ||
| 541 | - REPORT_INNER_ERR_MSG("E18888", "AddEdge %s->%s failed.", src_node->GetName().c_str(), | ||
| 542 | - insert_node->GetName().c_str()); | ||
| 543 | - GELOGE(af::GRAPH_FAILED, "[Add][Edge] %s->%s failed.", src_node->GetName().c_str(), insert_node->GetName().c_str()); | ||
| 544 | - return af::GRAPH_FAILED; | ||
| 545 | - } | ||
| 546 | - | ||
| 547 | - const OutControlAnchorPtr src_out_ctrl_anchor = src_node->GetOutControlAnchor(); | ||
| 548 | - GE_CHECK_NOTNULL(src_out_ctrl_anchor); | ||
| 549 | - | ||
| 550 | - bool ctrl_edge_flag = true; | ||
| 551 | - const std::string type = NodeUtils::GetNodeType(src->GetOwnerNode()); | ||
| 552 | - if ((type == SWITCH) || (type == REFSWITCH) || (type == SWITCHN)) { | ||
| 553 | - ctrl_edge_flag = false; | ||
| 554 | - } | ||
| 555 | - | ||
| 556 | - for (auto &dst : dsts) { | ||
| 557 | - GE_CHECK_NOTNULL(dst); | ||
| 558 | - const auto dst_node = dst->GetOwnerNodeBarePtr(); | ||
| 559 | - GELOGI("Insert node %s between %s->%s.", insert_node->GetName().c_str(), src_node->GetName().c_str(), | ||
| 560 | - dst_node->GetName().c_str()); | ||
| 561 | - if (src_node->GetOwnerComputeGraph() != dst_node->GetOwnerComputeGraph()) { | ||
| 562 | - REPORT_INNER_ERR_MSG("E18888", "src:%s and dst:%s does not exist in the same graph.", src_node->GetName().c_str(), | ||
| 563 | - dst_node->GetName().c_str()); | ||
| 564 | - GELOGE(af::GRAPH_FAILED, "[Check][Param] src:%s and dst:%s does not exist in the same graph.", | ||
| 565 | - src_node->GetName().c_str(), dst_node->GetName().c_str()); | ||
| 566 | - return af::GRAPH_FAILED; | ||
| 567 | - } | ||
| 568 | - | ||
| 569 | - (void)RemoveEdge(src, dst); | ||
| 570 | - if (AddEdge(insert_node->GetOutDataAnchor(static_cast<int32_t>(output_index)), dst) != af::GRAPH_SUCCESS) { | ||
| 571 | - REPORT_INNER_ERR_MSG("E18888", "ReplaceEdge from %s->%s to %s->%s failed.", src_node->GetName().c_str(), | ||
| 572 | - dst_node->GetName().c_str(), insert_node->GetName().c_str(), dst_node->GetName().c_str()); | ||
| 573 | - GELOGE(af::GRAPH_FAILED, "[Replace][Edge] from %s->%s to %s->%s failed.", src_node->GetName().c_str(), | ||
| 574 | - dst_node->GetName().c_str(), insert_node->GetName().c_str(), dst_node->GetName().c_str()); | ||
| 575 | - return af::GRAPH_FAILED; | ||
| 576 | - } | ||
| 577 | - | ||
| 578 | - if (!ctrl_edge_flag) { | ||
| 579 | - continue; | ||
| 580 | - } | ||
| 581 | - for (const InControlAnchorPtr &peer_in_ctrl_anchor : src_out_ctrl_anchor->GetPeerInControlAnchors()) { | ||
| 582 | - if ((RemoveEdge(src_out_ctrl_anchor, peer_in_ctrl_anchor) != af::GRAPH_SUCCESS) || | ||
| 583 | - (AddEdge(insert_node->GetOutControlAnchor(), peer_in_ctrl_anchor) != af::GRAPH_SUCCESS)) { | ||
| 584 | - REPORT_INNER_ERR_MSG("E18888", "ReplaceEdge from %s->%s to %s->%s failed.", src_node->GetName().c_str(), | ||
| 585 | - peer_in_ctrl_anchor->GetOwnerNode()->GetName().c_str(), insert_node->GetName().c_str(), | ||
| 586 | - peer_in_ctrl_anchor->GetOwnerNode()->GetName().c_str()); | ||
| 587 | - GELOGE(af::GRAPH_FAILED, "[Replace][Edge] from %s->%s to %s->%s failed.", src_node->GetName().c_str(), | ||
| 588 | - peer_in_ctrl_anchor->GetOwnerNode()->GetName().c_str(), insert_node->GetName().c_str(), | ||
| 589 | - peer_in_ctrl_anchor->GetOwnerNode()->GetName().c_str()); | ||
| 590 | - return af::GRAPH_FAILED; | ||
| 591 | - } | ||
| 592 | - } | ||
| 593 | - } | ||
| 594 | - GE_ASSERT_SUCCESS(InheritAttr(src_node->GetOpDesc(), insert_node->GetOpDesc())); | ||
| 595 | - | ||
| 596 | - return af::GRAPH_SUCCESS; | ||
| 597 | -} | ||
| 598 | - | ||
| 599 | -GE_FUNC_DEV_VISIBILITY GE_FUNC_HOST_VISIBILITY NodePtr GraphUtils::InsertNodeBefore(const InDataAnchorPtr &dst, | ||
| 600 | - const OpDescPtr &insert_op, | ||
| 601 | - const uint32_t input_index, | ||
| 602 | - const uint32_t output_index) { | ||
| 603 | - GE_ASSERT_NOTNULL(dst); | ||
| 604 | - const auto src_node_out_anchor = dst->GetPeerOutAnchor(); | ||
| 605 | - GE_ASSERT_NOTNULL(src_node_out_anchor); | ||
| 606 | - const auto src_node = src_node_out_anchor->GetOwnerNode(); | ||
| 607 | - GE_ASSERT_NOTNULL(src_node); | ||
| 608 | - auto compute_graph = src_node->GetOwnerComputeGraphBarePtr(); | ||
| 609 | - GE_ASSERT_NOTNULL(compute_graph); | ||
| 610 | - auto insert_node = compute_graph->InsertNode(src_node, insert_op); | ||
| 611 | - GE_ASSERT_GRAPH_SUCCESS(GraphUtils::InsertNodeBefore(dst, insert_node, input_index, output_index)); | ||
| 612 | - return insert_node; | ||
| 613 | -} | ||
| 614 | - | ||
| 615 | -graphStatus GraphUtils::InsertNodeBefore(const InDataAnchorPtr &dst, const NodePtr &insert_node, | ||
| 616 | - const uint32_t input_index, const uint32_t output_index) { | ||
| 617 | - GE_CHECK_NOTNULL(dst); | ||
| 618 | - GE_CHECK_NOTNULL(insert_node); | ||
| 619 | - const auto dst_node = dst->GetOwnerNodeBarePtr(); | ||
| 620 | - GE_CHECK_NOTNULL(dst_node); | ||
| 621 | - if (dst_node->GetOwnerComputeGraph() != insert_node->GetOwnerComputeGraph()) { | ||
| 622 | - GELOGE(af::GRAPH_FAILED, "[INSERT][NODE] dst:%s and insert_node:%s does not exist in the same graph.", | ||
| 623 | - dst_node->GetName().c_str(), insert_node->GetName().c_str()); | ||
| 624 | - return af::GRAPH_FAILED; | ||
| 625 | - } | ||
| 626 | - | ||
| 627 | - const auto src_node_out_anchor = dst->GetPeerOutAnchor(); | ||
| 628 | - GE_CHECK_NOTNULL(src_node_out_anchor); | ||
| 629 | - const auto src_node = src_node_out_anchor->GetOwnerNodeBarePtr(); | ||
| 630 | - GE_CHECK_NOTNULL(src_node); | ||
| 631 | - // insert node | ||
| 632 | - if ((RemoveEdge(src_node_out_anchor, dst) != af::GRAPH_SUCCESS) || | ||
| 633 | - (AddEdge(src_node_out_anchor, insert_node->GetInDataAnchor(static_cast<int32_t>(input_index))) != | ||
| 634 | - af::GRAPH_SUCCESS) || | ||
| 635 | - (AddEdge(insert_node->GetOutDataAnchor(static_cast<int32_t>(output_index)), dst) != af::GRAPH_SUCCESS)) { | ||
| 636 | - GELOGE(af::GRAPH_FAILED, "[INSERT][NODE] %s between %s->%s failed", insert_node->GetName().c_str(), | ||
| 637 | - src_node->GetName().c_str(), dst_node->GetName().c_str()); | ||
| 638 | - return af::GRAPH_FAILED; | ||
| 639 | - } | ||
| 640 | - GELOGI("[INSERT][NODE] %s between %s->%s", insert_node->GetName().c_str(), src_node->GetName().c_str(), | ||
| 641 | - dst_node->GetName().c_str()); | ||
| 642 | - | ||
| 643 | - // update control edges | ||
| 644 | - const auto in_ctrl_anchor = dst_node->GetInControlAnchor(); | ||
| 645 | - GE_CHECK_NOTNULL(in_ctrl_anchor); | ||
| 646 | - const auto insert_node_in_ctrl_anchor = insert_node->GetInControlAnchor(); | ||
| 647 | - for (const auto &peer_out_ctrl_anchor : in_ctrl_anchor->GetPeerOutControlAnchors()) { | ||
| 648 | - GE_CHECK_NOTNULL(peer_out_ctrl_anchor); | ||
| 649 | - const auto peer_node = peer_out_ctrl_anchor->GetOwnerNode(); | ||
| 650 | - if (NodeUtils::IsLikeAtomicClean(peer_node)) { | ||
| 651 | - continue; | ||
| 652 | - } | ||
| 653 | - if ((RemoveEdge(peer_out_ctrl_anchor, in_ctrl_anchor) != af::GRAPH_SUCCESS) || | ||
| 654 | - (AddEdge(peer_out_ctrl_anchor, insert_node_in_ctrl_anchor) != af::GRAPH_SUCCESS)) { | ||
| 655 | - GELOGE(af::GRAPH_FAILED, "[INSERT][NODE] replace control edge from %s->%s to %s->%s failed.", | ||
| 656 | - (peer_node != nullptr) ? peer_node->GetName().c_str() : "NULL", dst_node->GetName().c_str(), | ||
| 657 | - (peer_node != nullptr) ? peer_node->GetName().c_str() : "NULL", insert_node->GetName().c_str()); | ||
| 658 | - return af::GRAPH_FAILED; | ||
| 659 | - } | ||
| 660 | - } | ||
| 661 | - GE_ASSERT_SUCCESS(InheritAttr(dst_node->GetOpDesc(), insert_node->GetOpDesc())); | ||
| 662 | - | ||
| 663 | - return af::GRAPH_SUCCESS; | ||
| 664 | -} | ||
| 665 | - | ||
| 666 | GE_FUNC_DEV_VISIBILITY GE_FUNC_HOST_VISIBILITY graphStatus GraphUtils::RemoveJustNode(ComputeGraph &compute_graph, | 487 | GE_FUNC_DEV_VISIBILITY GE_FUNC_HOST_VISIBILITY graphStatus GraphUtils::RemoveJustNode(ComputeGraph &compute_graph, |
| 667 | const NodePtr &node) { | 488 | const NodePtr &node) { |
| 668 | if (node == nullptr) { | 489 | if (node == nullptr) { |
| @@ -19,6 +19,40 @@ namespace af { | |||
| 19 | class AscGraphUtils { | 19 | class AscGraphUtils { |
| 20 | public: | 20 | public: |
| 21 | static ComputeGraphPtr GetComputeGraph(const AscGraph &asc_graph); | 21 | static ComputeGraphPtr GetComputeGraph(const AscGraph &asc_graph); |
| 22 | + /** | ||
| 23 | + * 在源数据锚点和目标数据锚点之间插入节点,仅维护数据边;从源节点继承调度,并从 src 对应输出继承 | ||
| 24 | + * tensor 的 axis/repeats/strides 到插入节点指定输出。 | ||
| 25 | + * @note [Autofuse 完备适配] | ||
| 26 | + */ | ||
| 27 | + static graphStatus InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 28 | + const NodePtr &insert_node, const uint32_t input_index = 0U, | ||
| 29 | + const uint32_t output_index = 0U); | ||
| 30 | + /** | ||
| 31 | + * 在源数据锚点和其全部目标数据锚点之间插入节点,属性继承语义同指定 dsts 的重载。 | ||
| 32 | + * @note [Autofuse 完备适配] | ||
| 33 | + */ | ||
| 34 | + static graphStatus InsertNodeAfter(const OutDataAnchorPtr &src, const NodePtr &insert_node, | ||
| 35 | + const uint32_t input_index = 0U, const uint32_t output_index = 0U); | ||
| 36 | + /** | ||
| 37 | + * 通过 insert_op 创建节点后插入源数据锚点与目标数据锚点之间,属性继承语义同 NodePtr 重载。 | ||
| 38 | + * @note [Autofuse 完备适配] | ||
| 39 | + */ | ||
| 40 | + static NodePtr InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 41 | + const OpDescPtr &insert_op, const uint32_t input_index = 0U, | ||
| 42 | + const uint32_t output_index = 0U); | ||
| 43 | + /** | ||
| 44 | + * 在目标数据锚点前插入节点,仅维护数据边;从目标节点继承调度,并从原 src 对应输出继承 tensor 的 | ||
| 45 | + * axis/repeats/strides 到插入节点指定输出。 | ||
| 46 | + * @note [Autofuse 完备适配] | ||
| 47 | + */ | ||
| 48 | + static graphStatus InsertNodeBefore(const InDataAnchorPtr &dst, const NodePtr &insert_node, | ||
| 49 | + const uint32_t input_index = 0U, const uint32_t output_index = 0U); | ||
| 50 | + /** | ||
| 51 | + * 通过 insert_op 创建节点后插入目标数据锚点前,属性继承语义同 NodePtr 重载。 | ||
| 52 | + * @note [Autofuse 完备适配] | ||
| 53 | + */ | ||
| 54 | + static NodePtr InsertNodeBefore(const InDataAnchorPtr &dst, const OpDescPtr &insert_op, | ||
| 55 | + const uint32_t input_index = 0U, const uint32_t output_index = 0U); | ||
| 22 | static Status FromComputeGraph(const ComputeGraphPtr &compute_graph, AscGraph &graph); | 56 | static Status FromComputeGraph(const ComputeGraphPtr &compute_graph, AscGraph &graph); |
| 23 | /** | 57 | /** |
| 24 | * @param compute_graph的node对象是Node类型时候,接口内部转换为AscNode | 58 | * @param compute_graph的node对象是Node类型时候,接口内部转换为AscNode |
| @@ -311,69 +311,6 @@ class GraphUtils { | |||
| 311 | */ | 311 | */ |
| 312 | static OpDescPtr CopyOpDesc(const ConstOpDescPtr &org_op_desc, const AttrFilter &attr_filter); | 312 | static OpDescPtr CopyOpDesc(const ConstOpDescPtr &org_op_desc, const AttrFilter &attr_filter); |
| 313 | 313 | ||
| 314 | - /** | ||
| 315 | - * 接口行为是在数据`src`锚点所属的`src_node`节点和数据`dsts`锚点所属的`dst_node`节点们之间插入一个`insert_node`节点, | ||
| 316 | - * 默认是`insert_node`的`0`号数据输入锚点和`0`号输出数据锚点参与连边,`insert_node`插入之后, `src_node`和`insert_node` | ||
| 317 | - * 作为一个整体与原来的`src_node`具备等价的控制和数据关系 | ||
| 318 | - * `insert_node`继承`src_node`的用户属性 | ||
| 319 | - * @param src 源数据输出锚点 | ||
| 320 | - * @param dsts 源数据输出锚点连接的目的数据输入锚点,使用vector的原因是存在一个源锚点给到多个目的锚点的情况 | ||
| 321 | - * @param insert_node 表示要插入的节点 | ||
| 322 | - * @param input_index 表示插入节点的哪个数据输入锚点要跟src相连,如果不传递,默认取0 | ||
| 323 | - * @param output_index 表示插入节点的哪个数据输出锚点要跟dsts依次相连,如果不传递,默认取0 | ||
| 324 | - * @return 如果插入成功返回GRAPH_SUCCESS,失败返回GRAPH_FAILED | ||
| 325 | - */ | ||
| 326 | - static graphStatus InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 327 | - const NodePtr &insert_node, const uint32_t input_index = 0U, | ||
| 328 | - const uint32_t output_index = 0U); | ||
| 329 | - | ||
| 330 | - /** | ||
| 331 | - * 接口行为是通过insert_op在图中生成一个insert_node节点, | ||
| 332 | - * 在数据`src`锚点所属的`src_node`节点和数据`dsts`锚点所属的`dst_node`节点们之间插入一个`insert_node`节点, | ||
| 333 | - * 默认是`insert_node`的`0`号数据输入锚点和`0`号输出数据锚点参与连边 | ||
| 334 | - * `insert_node`插入之后, `src_node`和`insert_node` | ||
| 335 | - * 作为一个整体与原来的`src_node`具备等价的控制和数据关系 | ||
| 336 | - * `insert_node`继承`src_node`的用户属性 | ||
| 337 | - * @param src 源数据输出锚点 | ||
| 338 | - * @param dsts 源数据输出锚点连接的目的数据输入锚点,使用vector的原因是存在一个源锚点给到多个目的锚点的情况 | ||
| 339 | - * @param insert_op 表示要插入的opDesc,需要用其在src_node的图上生成一个node | ||
| 340 | - * @param input_index 表示插入节点的哪个数据输入锚点要跟src相连,如果不传递,默认取0 | ||
| 341 | - * @param output_index 表示插入节点的哪个数据输出锚点要跟dsts依次相连,如果不传递,默认取0 | ||
| 342 | - * @return 如果插入成功返回insert_node,失败返回nullptr | ||
| 343 | - */ | ||
| 344 | - static NodePtr InsertNodeAfter(const OutDataAnchorPtr &src, const std::vector<InDataAnchorPtr> &dsts, | ||
| 345 | - const OpDescPtr &insert_op, const uint32_t input_index = 0U, | ||
| 346 | - const uint32_t output_index = 0U); | ||
| 347 | - | ||
| 348 | - /** | ||
| 349 | - * 接口行为是在数据`dst`锚点所属的`dst_node`节点和其对端`src_node`节点之间插入一个`insert_node`节点, | ||
| 350 | - * 默认是`insert_node`的`0`号数据输入锚点和`0`号数据输出数据锚点参与连边,`insert_node`插入之后, | ||
| 351 | - * `dst_node`和`insert_node`作为一个整体与原来的`dst_node`具备等价的控制和数据关系 | ||
| 352 | - * `insert_node`继承`dst_node`的用户属性 | ||
| 353 | - * @param dst 目的数据输入锚点 | ||
| 354 | - * @param insert_node 表示要插入的节点 | ||
| 355 | - * @param input_index 表示插入节点的哪个数据输入锚点要跟dst的对端src锚点相连,如果不传递,默认取0 | ||
| 356 | - * @param output_index 表示插入节点的哪个数据输出锚点要跟dst相连,如果不传递,默认取0 | ||
| 357 | - * @return 如果插入成功返回GRAPH_SUCCESS,失败返回GRAPH_FAILED | ||
| 358 | - */ | ||
| 359 | - static graphStatus InsertNodeBefore(const InDataAnchorPtr &dst, const NodePtr &insert_node, | ||
| 360 | - const uint32_t input_index = 0U, const uint32_t output_index = 0U); | ||
| 361 | - | ||
| 362 | - /** | ||
| 363 | - * 接口行为是通过insert_op在图中生成一个insert_node节点, | ||
| 364 | - * 在数据`dst`锚点所属的`dst_node`节点和其对端`src_node`节点之间插入一个`insert_node`节点, | ||
| 365 | - * 默认是`insert_node`的`0`号数据输入锚点和`0`号数据输出数据锚点参与连边,`insert_node`插入之后, | ||
| 366 | - * `dst_node`和`insert_node`作为一个整体与原来的`dst_node`具备等价的控制和数据关系 | ||
| 367 | - * `insert_node`继承`dst_node`的用户属性 | ||
| 368 | - * @param dst 目的数据输入锚点 | ||
| 369 | - * @param insert_op 表示要插入的opDesc,需要用其在src_node的图上生成一个node | ||
| 370 | - * @param input_index 表示插入节点的哪个数据输入锚点要跟dst的对端src锚点相连,如果不传递,默认取0 | ||
| 371 | - * @param output_index 表示插入节点的哪个数据输出锚点要跟dst相连,如果不传递,默认取0 | ||
| 372 | - * @return 如果插入成功返回insert_node,失败返回nullptr | ||
| 373 | - */ | ||
| 374 | - static NodePtr InsertNodeBefore(const InDataAnchorPtr &dst, const OpDescPtr &insert_op, | ||
| 375 | - const uint32_t input_index = 0U, const uint32_t output_index = 0U); | ||
| 376 | - | ||
| 377 | /** | 314 | /** |
| 378 | * 从`compute_graph`智能指针管理的图对象的包含的nodes列表中删除`node`节点,仅仅是删除节点, | 315 | * 从`compute_graph`智能指针管理的图对象的包含的nodes列表中删除`node`节点,仅仅是删除节点, |
| 379 | * 不包含对node的断边和重新连边等操作 | 316 | * 不包含对node的断边和重新连边等操作 |
| @@ -13,6 +13,7 @@ | |||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | + | ||
| 16 | 17 | ||
| 17 | 18 | ||
| 18 | 19 | ||
| @@ -738,12 +739,7 @@ Status BufQueAllocator::ShortenVecinLifetime(af::AscGraph &graph, size_t max_que | |||
| 738 | 739 | ||
| 739 | auto load_out_anchor = top_cycle->node->GetOutDataAnchor(0); | 740 | auto load_out_anchor = top_cycle->node->GetOutDataAnchor(0); |
| 740 | GE_ASSERT_NOTNULL(load_out_anchor); | 741 | GE_ASSERT_NOTNULL(load_out_anchor); |
| 741 | - for (auto &peer_in_anchor : load_out_anchor->GetPeerInDataAnchors()) { | 742 | + GE_ASSERT_SUCCESS(af::AscGraphUtils::InsertNodeAfter(load_out_anchor, ub2ub_node)); |
| 742 | - GE_ASSERT_SUCCESS(af::GraphUtils::RemoveEdge(load_out_anchor, peer_in_anchor)); | ||
| 743 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(ub2ub_node->GetOutDataAnchor(0), peer_in_anchor)); | ||
| 744 | - } | ||
| 745 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(load_out_anchor, ub2ub_node->GetInDataAnchor(0))); | ||
| 746 | - ub2ub_node->attr.sched = top_cycle->node->attr.sched; | ||
| 747 | ub2ub_node->attr.api.compute_type = af::ComputeType::kComputeElewise; | 743 | ub2ub_node->attr.api.compute_type = af::ComputeType::kComputeElewise; |
| 748 | ub2ub_node->attr.api.type = af::ApiType::kAPITypeCompute; | 744 | ub2ub_node->attr.api.type = af::ApiType::kAPITypeCompute; |
| 749 | ub2ub_node->attr.api.unit = af::ComputeUnit::kUnitVector; | 745 | ub2ub_node->attr.api.unit = af::ComputeUnit::kUnitVector; |
| @@ -829,10 +825,7 @@ Status BufQueAllocator::ShortenVecoutLifetime(af::AscGraph &graph, size_t max_qu | |||
| 829 | af::AscNodePtr ub2ub_node = graph.AddNode(ub2ub); | 825 | af::AscNodePtr ub2ub_node = graph.AddNode(ub2ub); |
| 830 | GE_ASSERT_NOTNULL(ub2ub_node); | 826 | GE_ASSERT_NOTNULL(ub2ub_node); |
| 831 | 827 | ||
| 832 | - GE_ASSERT_SUCCESS(af::GraphUtils::RemoveEdge(out_data_anchor, peer_in_anchor)); | 828 | + GE_ASSERT_SUCCESS(af::AscGraphUtils::InsertNodeAfter(out_data_anchor, {peer_in_anchor}, ub2ub_node)); |
| 833 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(ub2ub_node->GetOutDataAnchor(0), peer_in_anchor)); | ||
| 834 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(out_data_anchor, ub2ub_node->GetInDataAnchor(0))); | ||
| 835 | - ub2ub_node->attr.sched = top_cycle->node->attr.sched; | ||
| 836 | ub2ub_node->attr.api.compute_type = af::ComputeType::kComputeElewise; | 829 | ub2ub_node->attr.api.compute_type = af::ComputeType::kComputeElewise; |
| 837 | ub2ub_node->attr.api.type = af::ApiType::kAPITypeCompute; | 830 | ub2ub_node->attr.api.type = af::ApiType::kAPITypeCompute; |
| 838 | ub2ub_node->attr.api.unit = af::ComputeUnit::kUnitVector; | 831 | ub2ub_node->attr.api.unit = af::ComputeUnit::kUnitVector; |
| @@ -12,6 +12,7 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | + | ||
| 15 | 16 | ||
| 16 | 17 | ||
| 17 | 18 | ||
| @@ -287,14 +288,10 @@ Status DtypeConsistency::InsertCastNode(af::AscGraph &graph, const af::AscNodePt | |||
| 287 | Cast cast_node(cast_name.c_str()); | 288 | Cast cast_node(cast_name.c_str()); |
| 288 | auto cast_node_ptr = graph.AddNode(cast_node); | 289 | auto cast_node_ptr = graph.AddNode(cast_node); |
| 289 | GE_ASSERT_NOTNULL(cast_node_ptr); | 290 | GE_ASSERT_NOTNULL(cast_node_ptr); |
| 290 | - cast_node_ptr->attr.sched = dst_node->attr.sched; | ||
| 291 | cast_node_ptr->outputs[0].attr = dst_node->inputs[input_idx].attr; | 291 | cast_node_ptr->outputs[0].attr = dst_node->inputs[input_idx].attr; |
| 292 | cast_node_ptr->outputs[0].attr.dtype = target_dtype; | 292 | cast_node_ptr->outputs[0].attr.dtype = target_dtype; |
| 293 | 293 | ||
| 294 | - // Reconnect edges | 294 | + GE_ASSERT_SUCCESS(af::AscGraphUtils::InsertNodeBefore(in_anchor, cast_node_ptr)); |
| 295 | - GE_ASSERT_SUCCESS(af::GraphUtils::RemoveEdge(src_out_anchor, in_anchor)); | ||
| 296 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(src_out_anchor, cast_node_ptr->GetInDataAnchor(0))); | ||
| 297 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(cast_node_ptr->GetOutDataAnchor(0), in_anchor)); | ||
| 298 | 295 | ||
| 299 | return af::SUCCESS; | 296 | return af::SUCCESS; |
| 300 | } | 297 | } |
| @@ -10,6 +10,7 @@ | |||
| 10 | 10 | ||
| 11 | 11 | ||
| 12 | 12 | ||
| 13 | + | ||
| 13 | 14 | ||
| 14 | 15 | ||
| 15 | 16 | ||
| @@ -345,10 +346,7 @@ af::Status BaseAlignmentStrategy::AddPadForAlignmentConflictOneNode(ascir::ImplG | |||
| 345 | tensor_to_align_type_[&pad_node->outputs[0].attr] = {AlignmentType::kAligned}; | 346 | tensor_to_align_type_[&pad_node->outputs[0].attr] = {AlignmentType::kAligned}; |
| 346 | auto out_anchor = node->GetOutDataAnchor(static_cast<int32_t>(i)); | 347 | auto out_anchor = node->GetOutDataAnchor(static_cast<int32_t>(i)); |
| 347 | GE_ASSERT_NOTNULL(out_anchor); | 348 | GE_ASSERT_NOTNULL(out_anchor); |
| 348 | - for (auto &in_anchor : out_anchor->GetPeerInDataAnchors()) { | 349 | + GE_ASSERT_SUCCESS(af::AscGraphUtils::InsertNodeAfter(out_anchor, pad_node)); |
| 349 | - GE_ASSERT_SUCCESS(af::GraphUtils::ReplaceEdgeSrc(out_anchor, in_anchor, pad_node->GetOutDataAnchor(0))); | ||
| 350 | - } | ||
| 351 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(out_anchor, pad_node->GetInDataAnchor(0))); | ||
| 352 | } | 350 | } |
| 353 | return af::SUCCESS; | 351 | return af::SUCCESS; |
| 354 | } | 352 | } |
| @@ -423,29 +423,16 @@ Status InsertCastToIncreasePrecision(AscGraph &asc_graph, const NodePtr &load_no | |||
| 423 | GE_ASSERT_SUCCESS(UpdateTopoId(asc_graph, load_node, 1)); | 423 | GE_ASSERT_SUCCESS(UpdateTopoId(asc_graph, load_node, 1)); |
| 424 | auto c_node = BuildCastNode(asc_graph, load_node); | 424 | auto c_node = BuildCastNode(asc_graph, load_node); |
| 425 | GE_ASSERT_NOTNULL(c_node); | 425 | GE_ASSERT_NOTNULL(c_node); |
| 426 | - GE_ASSERT_GRAPH_SUCCESS(GraphUtils::ReplaceNodeDataAnchors(c_node, load_node, {}, {0})); | 426 | + const auto load_out_anchor = load_node->GetOutDataAnchor(0); |
| 427 | - GE_ASSERT_GRAPH_SUCCESS(GraphUtils::AddEdge(load_node->GetOutDataAnchor(0), c_node->GetInDataAnchor(0))); | 427 | + GE_ASSERT_NOTNULL(load_out_anchor); |
| 428 | + GE_ASSERT_GRAPH_SUCCESS(AscGraphUtils::InsertNodeAfter(load_out_anchor, c_node)); | ||
| 428 | const auto c_opdesc = c_node->GetOpDesc(); | 429 | const auto c_opdesc = c_node->GetOpDesc(); |
| 429 | GE_ASSERT_NOTNULL(c_opdesc); | 430 | GE_ASSERT_NOTNULL(c_opdesc); |
| 430 | const auto c_output_tensor_desc = c_opdesc->MutableOutputDesc(0); | 431 | const auto c_output_tensor_desc = c_opdesc->MutableOutputDesc(0); |
| 431 | GE_ASSERT_NOTNULL(c_output_tensor_desc); | 432 | GE_ASSERT_NOTNULL(c_output_tensor_desc); |
| 432 | c_output_tensor_desc->SetDataType(DT_FLOAT); | 433 | c_output_tensor_desc->SetDataType(DT_FLOAT); |
| 433 | - const auto c_o_attr = c_output_tensor_desc->GetOrCreateAttrsGroup<AscTensorAttr>(); | ||
| 434 | - GE_ASSERT_NOTNULL(c_o_attr); | ||
| 435 | const auto load_opdesc = load_node->GetOpDesc(); | 434 | const auto load_opdesc = load_node->GetOpDesc(); |
| 436 | GE_ASSERT_NOTNULL(load_opdesc); | 435 | GE_ASSERT_NOTNULL(load_opdesc); |
| 437 | - const auto load_output_tensor_desc = load_opdesc->MutableOutputDesc(0); | ||
| 438 | - GE_ASSERT_NOTNULL(load_output_tensor_desc); | ||
| 439 | - const auto load_attr = load_output_tensor_desc->GetAttrsGroup<AscTensorAttr>(); | ||
| 440 | - GE_ASSERT_NOTNULL(load_attr); | ||
| 441 | - c_o_attr->axis = load_attr->axis; | ||
| 442 | - c_o_attr->repeats = load_attr->repeats; | ||
| 443 | - c_o_attr->strides = load_attr->strides; | ||
| 444 | - const auto c_node_attr = c_opdesc->GetOrCreateAttrsGroup<AscNodeAttr>(); | ||
| 445 | - GE_ASSERT_NOTNULL(c_node_attr); | ||
| 446 | - const auto load_node_attr = load_opdesc->GetAttrsGroup<AscNodeAttr>(); | ||
| 447 | - GE_ASSERT_NOTNULL(load_node_attr); | ||
| 448 | - c_node_attr->sched.axis = load_node_attr->sched.axis; | ||
| 449 | c_opdesc->SetId(load_opdesc->GetId() + 1); | 436 | c_opdesc->SetId(load_opdesc->GetId() + 1); |
| 450 | return af::SUCCESS; | 437 | return af::SUCCESS; |
| 451 | } | 438 | } |
| @@ -156,26 +156,22 @@ Status CastOptimizationPass::DoOptimize(AscGraph &graph, const AscNodePtr &node, | |||
| 156 | GELOGD("input index = %d, source dtype already matches dst_dtype, skip adding Cast", concat_in_anchor->GetIdx()); | 156 | GELOGD("input index = %d, source dtype already matches dst_dtype, skip adding Cast", concat_in_anchor->GetIdx()); |
| 157 | continue; | 157 | continue; |
| 158 | } | 158 | } |
| 159 | - GE_ASSERT_SUCCESS(GraphUtils::RemoveEdge(src_out_anchor, concat_in_anchor)); | ||
| 160 | const auto it = out_anchor_to_cast_node.find(src_out_anchor.get()); | 159 | const auto it = out_anchor_to_cast_node.find(src_out_anchor.get()); |
| 161 | if (it != out_anchor_to_cast_node.cend()) { | 160 | if (it != out_anchor_to_cast_node.cend()) { |
| 161 | + GE_ASSERT_SUCCESS(GraphUtils::RemoveEdge(src_out_anchor, concat_in_anchor)); | ||
| 162 | GE_ASSERT_SUCCESS(GraphUtils::AddEdge(it->second->GetOutDataAnchor(0), concat_in_anchor)); | 162 | GE_ASSERT_SUCCESS(GraphUtils::AddEdge(it->second->GetOutDataAnchor(0), concat_in_anchor)); |
| 163 | GELOGD("input index = %d, reuse existing Cast node for shared source", concat_in_anchor->GetIdx()); | 163 | GELOGD("input index = %d, reuse existing Cast node for shared source", concat_in_anchor->GetIdx()); |
| 164 | continue; | 164 | continue; |
| 165 | } | 165 | } |
| 166 | ascir_op::Cast cast_op((src_node->GetName() + "_cast_optimization_pass").c_str()); | 166 | ascir_op::Cast cast_op((src_node->GetName() + "_cast_optimization_pass").c_str()); |
| 167 | cast_op.attr = out_cast_node->attr; | 167 | cast_op.attr = out_cast_node->attr; |
| 168 | - cast_op.attr.sched = src_node->attr.sched; | ||
| 169 | auto &src_node_output_tensor_attr = src_node->outputs[0].attr; | 168 | auto &src_node_output_tensor_attr = src_node->outputs[0].attr; |
| 170 | - *cast_op.y.axis = src_node_output_tensor_attr.axis; | ||
| 171 | cast_op.y.dtype = dst_dtype; | 169 | cast_op.y.dtype = dst_dtype; |
| 172 | - *cast_op.y.repeats = src_node_output_tensor_attr.repeats; | ||
| 173 | - ::optimize::ScheduleUtils::GenerateStrides(src_node_output_tensor_attr.repeats, *cast_op.y.strides); | ||
| 174 | const auto cast_node = graph.AddNode(cast_op); | 170 | const auto cast_node = graph.AddNode(cast_op); |
| 175 | GE_ASSERT_NOTNULL(cast_node); | 171 | GE_ASSERT_NOTNULL(cast_node); |
| 176 | out_anchor_to_cast_node[src_out_anchor.get()] = cast_node; | 172 | out_anchor_to_cast_node[src_out_anchor.get()] = cast_node; |
| 177 | - GE_ASSERT_SUCCESS(GraphUtils::AddEdge(src_out_anchor, cast_node->GetInDataAnchor(0))); | 173 | + GE_ASSERT_SUCCESS(AscGraphUtils::InsertNodeAfter(src_out_anchor, {concat_in_anchor}, cast_node)); |
| 178 | - GE_ASSERT_SUCCESS(GraphUtils::AddEdge(cast_node->GetOutDataAnchor(0), concat_in_anchor)); | 174 | + ::optimize::ScheduleUtils::GenerateStrides(src_node_output_tensor_attr.repeats, cast_node->outputs[0].attr.strides); |
| 179 | GELOGD("input index = %d, new Cast node was added", concat_in_anchor->GetIdx()); | 175 | GELOGD("input index = %d, new Cast node was added", concat_in_anchor->GetIdx()); |
| 180 | } | 176 | } |
| 181 | GE_ASSERT_SUCCESS(UpdateDtype(node, dst_dtype)); | 177 | GE_ASSERT_SUCCESS(UpdateDtype(node, dst_dtype)); |
| @@ -11,6 +11,7 @@ | |||
| 11 | 11 | ||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | + | ||
| 14 | 15 | ||
| 15 | 16 | ||
| 16 | 17 | ||
| @@ -139,7 +140,6 @@ Status ConcatInputUnificationPass::DoOptimize(ascir::ImplGraph &graph, const af: | |||
| 139 | af::ascir_op::Ub2ub ub2ub(ub_name.c_str()); | 140 | af::ascir_op::Ub2ub ub2ub(ub_name.c_str()); |
| 140 | af::AscNodePtr ub2ub_node = graph.AddNode(ub2ub); | 141 | af::AscNodePtr ub2ub_node = graph.AddNode(ub2ub); |
| 141 | GE_ASSERT_NOTNULL(ub2ub_node); | 142 | GE_ASSERT_NOTNULL(ub2ub_node); |
| 142 | - ub2ub_node->attr.sched = asc_node->attr.sched; | ||
| 143 | ub2ub_node->attr.api.compute_type = af::ComputeType::kComputeElewise; | 143 | ub2ub_node->attr.api.compute_type = af::ComputeType::kComputeElewise; |
| 144 | ub2ub_node->attr.api.type = af::ApiType::kAPITypeCompute; | 144 | ub2ub_node->attr.api.type = af::ApiType::kAPITypeCompute; |
| 145 | ub2ub_node->attr.api.unit = af::ComputeUnit::kUnitVector; | 145 | ub2ub_node->attr.api.unit = af::ComputeUnit::kUnitVector; |
| @@ -147,9 +147,7 @@ Status ConcatInputUnificationPass::DoOptimize(ascir::ImplGraph &graph, const af: | |||
| 147 | ub2ub_node->outputs[0].attr.buf = {}; | 147 | ub2ub_node->outputs[0].attr.buf = {}; |
| 148 | ub2ub_node->outputs[0].attr.que = {}; | 148 | ub2ub_node->outputs[0].attr.que = {}; |
| 149 | 149 | ||
| 150 | - GE_ASSERT_SUCCESS(af::GraphUtils::RemoveEdge(out_anchor, in_anchor)); | 150 | + GE_ASSERT_SUCCESS(af::AscGraphUtils::InsertNodeAfter(out_anchor, {in_anchor}, ub2ub_node)); |
| 151 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(ub2ub_node->GetOutDataAnchor(0), in_anchor)); | ||
| 152 | - GE_ASSERT_SUCCESS(af::GraphUtils::AddEdge(out_anchor, ub2ub_node->GetInDataAnchor(0))); | ||
| 153 | GELOGD("Ub2ub node: %s added", ub2ub_node->GetNamePtr()); | 151 | GELOGD("Ub2ub node: %s added", ub2ub_node->GetNamePtr()); |
| 154 | } | 152 | } |
| 155 | return af::SUCCESS; | 153 | return af::SUCCESS; |
| @@ -36,6 +36,7 @@ add_library(test_ascir_ut OBJECT | |||
| 36 | reg_func/test_reg_func_reduce_max.cpp | 36 | reg_func/test_reg_func_reduce_max.cpp |
| 37 | code_dumper_unittest.cc | 37 | code_dumper_unittest.cc |
| 38 | ascir_utils_unittest.cc | 38 | ascir_utils_unittest.cc |
| 39 | + test_asc_graph_utils.cpp | ||
| 39 | ) | 40 | ) |
| 40 | target_include_directories(test_ascir_ut PRIVATE | 41 | target_include_directories(test_ascir_ut PRIVATE |
| 41 | ${ASCEND_ROOT}/x86_64-linux/include | 42 | ${ASCEND_ROOT}/x86_64-linux/include |
| @@ -0,0 +1,185 @@ | |||
| 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 af { | ||
| 19 | +namespace { | ||
| 20 | +class GraphBuilder { | ||
| 21 | + public: | ||
| 22 | + explicit GraphBuilder(const std::string &name) : graph_(std::make_shared<ComputeGraph>(name)) {} | ||
| 23 | + | ||
| 24 | + NodePtr AddNode(const std::string &name, int32_t input_num, int32_t output_num) { | ||
| 25 | + auto op_desc = std::make_shared<OpDesc>(name, name); | ||
| 26 | + for (int32_t i = 0; i < input_num; ++i) { | ||
| 27 | + op_desc->AddInputDesc(GeTensorDesc()); | ||
| 28 | + } | ||
| 29 | + for (int32_t i = 0; i < output_num; ++i) { | ||
| 30 | + op_desc->AddOutputDesc(GeTensorDesc()); | ||
| 31 | + } | ||
| 32 | + return graph_->AddNode(op_desc); | ||
| 33 | + } | ||
| 34 | + | ||
| 35 | + void AddEdge(const NodePtr &src, const NodePtr &dst) { | ||
| 36 | + ASSERT_EQ(GraphUtils::AddEdge(src->GetOutDataAnchor(0), dst->GetInDataAnchor(0)), GRAPH_SUCCESS); | ||
| 37 | + } | ||
| 38 | + | ||
| 39 | + private: | ||
| 40 | + ComputeGraphPtr graph_; | ||
| 41 | +}; | ||
| 42 | + | ||
| 43 | +AscNodeAttr *GetNodeAttr(const NodePtr &node) { | ||
| 44 | + return node->GetOpDesc()->GetOrCreateAttrsGroup<AscNodeAttr>(); | ||
| 45 | +} | ||
| 46 | + | ||
| 47 | +AscTensorAttr *GetTensorAttr(const NodePtr &node, size_t index = 0U) { | ||
| 48 | + return node->GetOpDesc()->MutableOutputDesc(index)->GetOrCreateAttrsGroup<AscTensorAttr>(); | ||
| 49 | +} | ||
| 50 | + | ||
| 51 | +void SetReferenceAttrs(const NodePtr &node) { | ||
| 52 | + auto node_attr = GetNodeAttr(node); | ||
| 53 | + node_attr->sched.exec_order = 7; | ||
| 54 | + node_attr->sched.axis = {1, 2}; | ||
| 55 | + node_attr->sched.loop_axis = 2; | ||
| 56 | + node_attr->sched.exec_condition = ExecuteCondition::kCacheBlockSplitFusedBroadcastAxis; | ||
| 57 | + node_attr->api.type = ApiType::kAPITypeCompute; | ||
| 58 | + | ||
| 59 | + auto tensor_attr = GetTensorAttr(node); | ||
| 60 | + tensor_attr->axis = {1, 2}; | ||
| 61 | + tensor_attr->repeats = {Symbol(8), Symbol(16)}; | ||
| 62 | + tensor_attr->strides = {Symbol(16), Symbol(1)}; | ||
| 63 | +} | ||
| 64 | + | ||
| 65 | +void SetInsertedAttrs(const NodePtr &node) { | ||
| 66 | + auto node_attr = GetNodeAttr(node); | ||
| 67 | + node_attr->sched.axis = {9}; | ||
| 68 | + node_attr->api.type = ApiType::kAPITypeBuffer; | ||
| 69 | + | ||
| 70 | + auto tensor_attr = GetTensorAttr(node); | ||
| 71 | + node->GetOpDesc()->MutableOutputDesc(0)->SetDataType(DT_FLOAT16); | ||
| 72 | + tensor_attr->axis = {9}; | ||
| 73 | + tensor_attr->repeats = {Symbol(4)}; | ||
| 74 | + tensor_attr->strides = {Symbol(1)}; | ||
| 75 | + tensor_attr->vectorized_axis = {9}; | ||
| 76 | + tensor_attr->vectorized_strides = {Symbol(3)}; | ||
| 77 | + tensor_attr->mem.alloc_type = AllocType::kAllocTypeQueue; | ||
| 78 | + tensor_attr->que.id = 11; | ||
| 79 | + tensor_attr->buf.id = 12; | ||
| 80 | + tensor_attr->opt.reuse_id = 13; | ||
| 81 | +} | ||
| 82 | + | ||
| 83 | +void ExpectInheritedAttrs(const NodePtr &node_reference, const NodePtr &tensor_reference, const NodePtr &inserted) { | ||
| 84 | + const auto reference_node_attr = GetNodeAttr(node_reference); | ||
| 85 | + const auto inserted_node_attr = GetNodeAttr(inserted); | ||
| 86 | + EXPECT_EQ(inserted_node_attr->sched.exec_order, reference_node_attr->sched.exec_order); | ||
| 87 | + EXPECT_EQ(inserted_node_attr->sched.axis, reference_node_attr->sched.axis); | ||
| 88 | + EXPECT_EQ(inserted_node_attr->sched.loop_axis, reference_node_attr->sched.loop_axis); | ||
| 89 | + EXPECT_EQ(inserted_node_attr->sched.exec_condition, reference_node_attr->sched.exec_condition); | ||
| 90 | + EXPECT_EQ(inserted_node_attr->api.type, ApiType::kAPITypeBuffer); | ||
| 91 | + | ||
| 92 | + const auto reference_tensor_attr = GetTensorAttr(tensor_reference); | ||
| 93 | + const auto inserted_tensor_attr = GetTensorAttr(inserted); | ||
| 94 | + EXPECT_EQ(inserted_tensor_attr->axis, reference_tensor_attr->axis); | ||
| 95 | + EXPECT_EQ(inserted_tensor_attr->repeats, reference_tensor_attr->repeats); | ||
| 96 | + EXPECT_EQ(inserted_tensor_attr->strides, reference_tensor_attr->strides); | ||
| 97 | + EXPECT_EQ(inserted->GetOpDesc()->MutableOutputDesc(0)->GetDataType(), DT_FLOAT16); | ||
| 98 | + EXPECT_EQ(inserted_tensor_attr->vectorized_axis, std::vector<int64_t>{9}); | ||
| 99 | + EXPECT_EQ(inserted_tensor_attr->vectorized_strides, std::vector<Expression>{Symbol(3)}); | ||
| 100 | + EXPECT_EQ(inserted_tensor_attr->mem.alloc_type, AllocType::kAllocTypeQueue); | ||
| 101 | + EXPECT_EQ(inserted_tensor_attr->que.id, 11); | ||
| 102 | + EXPECT_EQ(inserted_tensor_attr->buf.id, 12); | ||
| 103 | + EXPECT_EQ(inserted_tensor_attr->opt.reuse_id, 13); | ||
| 104 | +} | ||
| 105 | +} // namespace | ||
| 106 | + | ||
| 107 | +TEST(AscGraphUtilsInsertNodeTest, InsertBeforeInheritsConsumerSchedAndProducerTensorAttrs) { | ||
| 108 | + GraphBuilder builder("before"); | ||
| 109 | + auto producer = builder.AddNode("producer", 0, 1); | ||
| 110 | + auto consumer = builder.AddNode("consumer", 1, 1); | ||
| 111 | + auto inserted = builder.AddNode("inserted", 1, 1); | ||
| 112 | + builder.AddEdge(producer, consumer); | ||
| 113 | + SetReferenceAttrs(producer); | ||
| 114 | + SetReferenceAttrs(consumer); | ||
| 115 | + GetTensorAttr(consumer)->axis = {3}; | ||
| 116 | + GetTensorAttr(consumer)->repeats = {Symbol(32)}; | ||
| 117 | + GetTensorAttr(consumer)->strides = {Symbol(4)}; | ||
| 118 | + SetInsertedAttrs(inserted); | ||
| 119 | + | ||
| 120 | + ASSERT_EQ(AscGraphUtils::InsertNodeBefore(consumer->GetInDataAnchor(0), inserted), GRAPH_SUCCESS); | ||
| 121 | + | ||
| 122 | + EXPECT_EQ(producer->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 1U); | ||
| 123 | + EXPECT_EQ(inserted->GetInDataAnchor(0)->GetPeerOutAnchor(), producer->GetOutDataAnchor(0)); | ||
| 124 | + EXPECT_EQ(inserted->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 1U); | ||
| 125 | + EXPECT_EQ(consumer->GetInDataAnchor(0)->GetPeerOutAnchor(), inserted->GetOutDataAnchor(0)); | ||
| 126 | + ExpectInheritedAttrs(consumer, producer, inserted); | ||
| 127 | +} | ||
| 128 | + | ||
| 129 | +TEST(AscGraphUtilsInsertNodeTest, InsertAfterInheritsProducerAutofuseAttrs) { | ||
| 130 | + GraphBuilder builder("after"); | ||
| 131 | + auto producer = builder.AddNode("producer", 0, 1); | ||
| 132 | + auto consumer1 = builder.AddNode("consumer1", 1, 1); | ||
| 133 | + auto consumer2 = builder.AddNode("consumer2", 1, 1); | ||
| 134 | + auto inserted = builder.AddNode("inserted", 1, 1); | ||
| 135 | + builder.AddEdge(producer, consumer1); | ||
| 136 | + builder.AddEdge(producer, consumer2); | ||
| 137 | + SetReferenceAttrs(producer); | ||
| 138 | + SetInsertedAttrs(inserted); | ||
| 139 | + | ||
| 140 | + ASSERT_EQ(AscGraphUtils::InsertNodeAfter(producer->GetOutDataAnchor(0), inserted), GRAPH_SUCCESS); | ||
| 141 | + | ||
| 142 | + EXPECT_EQ(producer->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 1U); | ||
| 143 | + EXPECT_EQ(inserted->GetInDataAnchor(0)->GetPeerOutAnchor(), producer->GetOutDataAnchor(0)); | ||
| 144 | + EXPECT_EQ(inserted->GetOutDataAnchor(0)->GetPeerInDataAnchors().size(), 2U); | ||
| 145 | + EXPECT_EQ(consumer1->GetInDataAnchor(0)->GetPeerOutAnchor(), inserted->GetOutDataAnchor(0)); | ||
| 146 | + EXPECT_EQ(consumer2->GetInDataAnchor(0)->GetPeerOutAnchor(), inserted->GetOutDataAnchor(0)); | ||
| 147 | + ExpectInheritedAttrs(producer, producer, inserted); | ||
| 148 | +} | ||
| 149 | + | ||
| 150 | +TEST(AscGraphUtilsInsertNodeTest, InsertAfterInheritsSelectedSourceOutputToSelectedInsertedOutput) { | ||
| 151 | + GraphBuilder builder("after_selected_output"); | ||
| 152 | + auto producer = builder.AddNode("producer", 0, 2); | ||
| 153 | + auto consumer = builder.AddNode("consumer", 1, 1); | ||
| 154 | + auto inserted = builder.AddNode("inserted", 1, 2); | ||
| 155 | + ASSERT_EQ(GraphUtils::AddEdge(producer->GetOutDataAnchor(1), consumer->GetInDataAnchor(0)), GRAPH_SUCCESS); | ||
| 156 | + | ||
| 157 | + GetTensorAttr(producer, 0)->axis = {10}; | ||
| 158 | + GetTensorAttr(producer, 1)->axis = {20}; | ||
| 159 | + GetTensorAttr(producer, 1)->repeats = {Symbol(8)}; | ||
| 160 | + GetTensorAttr(producer, 1)->strides = {Symbol(2)}; | ||
| 161 | + GetTensorAttr(inserted, 0)->axis = {30}; | ||
| 162 | + | ||
| 163 | + ASSERT_EQ( | ||
| 164 | + AscGraphUtils::InsertNodeAfter(producer->GetOutDataAnchor(1), {consumer->GetInDataAnchor(0)}, inserted, 0, 1), | ||
| 165 | + GRAPH_SUCCESS); | ||
| 166 | + | ||
| 167 | + EXPECT_EQ(GetTensorAttr(inserted, 0)->axis, std::vector<int64_t>{30}); | ||
| 168 | + EXPECT_EQ(GetTensorAttr(inserted, 1)->axis, std::vector<int64_t>{20}); | ||
| 169 | + EXPECT_EQ(GetTensorAttr(inserted, 1)->repeats, std::vector<Expression>{Symbol(8)}); | ||
| 170 | + EXPECT_EQ(GetTensorAttr(inserted, 1)->strides, std::vector<Expression>{Symbol(2)}); | ||
| 171 | +} | ||
| 172 | + | ||
| 173 | +TEST(AscGraphUtilsInsertNodeTest, RejectsInsertNodeFromAnotherGraph) { | ||
| 174 | + GraphBuilder first_builder("first"); | ||
| 175 | + auto producer = first_builder.AddNode("producer", 0, 1); | ||
| 176 | + auto consumer = first_builder.AddNode("consumer", 1, 1); | ||
| 177 | + first_builder.AddEdge(producer, consumer); | ||
| 178 | + GraphBuilder second_builder("second"); | ||
| 179 | + auto inserted = second_builder.AddNode("inserted", 1, 1); | ||
| 180 | + | ||
| 181 | + EXPECT_NE(AscGraphUtils::InsertNodeBefore(consumer->GetInDataAnchor(0), inserted), GRAPH_SUCCESS); | ||
| 182 | + EXPECT_EQ(consumer->GetInDataAnchor(0)->GetPeerOutAnchor(), producer->GetOutDataAnchor(0)); | ||
| 183 | +} | ||
| 184 | + | ||
| 185 | +} // namespace af | ||
| @@ -97,6 +97,27 @@ static ascir::FusedScheduledResult MakeFusedScheduledResultWithGraphs(std::vecto | |||
| 97 | return fused_result; | 97 | return fused_result; |
| 98 | } | 98 | } |
| 99 | 99 | ||
| 100 | +TEST_F(BufQueAllocatorUT, ShortenVecoutLifetimeInsertsUb2ubBeforeStore) { | ||
| 101 | + auto graph = MakeStaticLoadStoreGraph("shorten_vecout", 32); | ||
| 102 | + ASSERT_EQ(ScheduleUtils::TopologicalSorting(graph), af::GRAPH_SUCCESS); | ||
| 103 | + | ||
| 104 | + ASSERT_EQ(BufQueAllocator::ShortenVecoutLifetime(graph, 0), af::SUCCESS); | ||
| 105 | + | ||
| 106 | + const auto load = graph.FindNode("load"); | ||
| 107 | + const auto ub2ub = graph.FindNode("ub_cpy_load_0"); | ||
| 108 | + const auto store = graph.FindNode("store"); | ||
| 109 | + ASSERT_NE(load, nullptr); | ||
| 110 | + ASSERT_NE(ub2ub, nullptr); | ||
| 111 | + ASSERT_NE(store, nullptr); | ||
| 112 | + EXPECT_EQ(ub2ub->GetInDataAnchor(0)->GetPeerOutAnchor(), load->GetOutDataAnchor(0)); | ||
| 113 | + EXPECT_EQ(store->GetInDataAnchor(0)->GetPeerOutAnchor(), ub2ub->GetOutDataAnchor(0)); | ||
| 114 | + EXPECT_EQ(ub2ub->attr.sched.axis, load->attr.sched.axis); | ||
| 115 | + EXPECT_EQ(ub2ub->attr.sched.loop_axis, load->attr.sched.loop_axis); | ||
| 116 | + EXPECT_EQ(ub2ub->outputs[0].attr.axis, load->outputs[0].attr.axis); | ||
| 117 | + EXPECT_EQ(ub2ub->outputs[0].attr.repeats, load->outputs[0].attr.repeats); | ||
| 118 | + EXPECT_EQ(ub2ub->outputs[0].attr.strides, load->outputs[0].attr.strides); | ||
| 119 | +} | ||
| 120 | + | ||
| 100 | static af::AscGraph MakeSimtInlineTransformGraph() { | 121 | static af::AscGraph MakeSimtInlineTransformGraph() { |
| 101 | af::AscGraph graph("simt_inline_transform"); | 122 | af::AscGraph graph("simt_inline_transform"); |
| 102 | const auto size = graph.CreateSizeVar(8); | 123 | const auto size = graph.CreateSizeVar(8); |