已合并
【PR】: 重构,使用InsertNode收编节点插入行为 #1677
【PR】: 重构,使用InsertNode收编节点插入行为 #1677
已合并
shengnan创建于 8月7日
共 13 个文件变更+359-288
@@ -17,6 +17,15 @@
17 17 
18namespace af {18namespace af {
19namespace {19namespace {
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+ 
20graphStatus EstablishAscNodeAndEdges(const ascendc_ir::proto::AscGraphDef &asc_graph_def, AscGraph &out_asc_graph) {29graphStatus 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 AscGraph31 // 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} // namespace121} // 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+ 
113ComputeGraphPtr AscGraphUtils::GetComputeGraph(const AscGraph &asc_graph) {216ComputeGraphPtr 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;
78const std::set<std::string> kMergeInputSkipTypes{STREAMACTIVE, STREAMSWITCH, CONSTANT, CONSTANTOP};78const std::set<std::string> kMergeInputSkipTypes{STREAMACTIVE, STREAMSWITCH, CONSTANT, CONSTANTOP};
79constexpr int32_t kInvalidStream = -1;79constexpr int32_t kInvalidStream = -1;
80constexpr size_t kNoOpOptimizeThreshold = 1000UL;80constexpr 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 
99graphStatus ReLinkInputDataEdge(const NodePtr &input_node, const NodePtr &target_node) {82graphStatus 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- 
666GE_FUNC_DEV_VISIBILITY GE_FUNC_HOST_VISIBILITY graphStatus GraphUtils::RemoveJustNode(ComputeGraph &compute_graph,487GE_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 {
19class AscGraphUtils {19class 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类型时候,接口内部转换为AscNode58 * @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#include "ascir_ops.h"13#include "ascir_ops.h"
14#include "ascgen_log.h"14#include "ascgen_log.h"
15#include "ascir_ops_utils.h"15#include "ascir_ops_utils.h"
16+#include "graph/ascendc_ir/utils/asc_graph_utils.h"
16#include "schedule_utils.h"17#include "schedule_utils.h"
17#include "graph_utils.h"18#include "graph_utils.h"
18#include "common_utils.h"19#include "common_utils.h"
@@ -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#include "ascir_ops.h"12#include "ascir_ops.h"
13#include "ascir_ops_utils.h"13#include "ascir_ops_utils.h"
14#include "common_utils.h"14#include "common_utils.h"
15+#include "graph/ascendc_ir/utils/asc_graph_utils.h"
15#include "graph_utils.h"16#include "graph_utils.h"
16#include "node_utils.h"17#include "node_utils.h"
17#include "schedule_utils.h"18#include "schedule_utils.h"
@@ -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 edges294+ 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#include "base_alignment_strategy.h"11#include "base_alignment_strategy.h"
12#include "common_utils.h"12#include "common_utils.h"
13+#include "graph/ascendc_ir/utils/asc_graph_utils.h"
13#include "graph/symbolizer/symbolic_utils.h"14#include "graph/symbolizer/symbolic_utils.h"
14#include "indirect_load_utils.h"15#include "indirect_load_utils.h"
15#include "platform/platform_factory.h"16#include "platform/platform_factory.h"
@@ -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#include "optimize/task_generator/concat_inputs_unification_pass.h"11#include "optimize/task_generator/concat_inputs_unification_pass.h"
12 12 
13#include "ascir_utils.h"13#include "ascir_utils.h"
14+#include "graph/ascendc_ir/utils/asc_graph_utils.h"
14#include "graph_utils.h"15#include "graph_utils.h"
15#include "schedule_utils.h"16#include "schedule_utils.h"
16#include "buffer_allocate/tensor_mem_defs.h"17#include "buffer_allocate/tensor_mem_defs.h"
@@ -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.cpp36 reg_func/test_reg_func_reduce_max.cpp
37 code_dumper_unittest.cc37 code_dumper_unittest.cc
38 ascir_utils_unittest.cc38 ascir_utils_unittest.cc
39+ test_asc_graph_utils.cpp
39)40)
40target_include_directories(test_ascir_ut PRIVATE41target_include_directories(test_ascir_ut PRIVATE
41 ${ASCEND_ROOT}/x86_64-linux/include42 ${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+#include <gtest/gtest.h>
12+ 
13+#include "graph/ascendc_ir/ascendc_ir_core/ascendc_ir_def.h"
14+#include "graph/ascendc_ir/utils/asc_graph_utils.h"
15+#include "graph/compute_graph.h"
16+#include "graph/utils/graph_utils.h"
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+ 
100static af::AscGraph MakeSimtInlineTransformGraph() {121static 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);