已合并
perf: 缓存batch_all_memory_size引用消除重复查找 #4609
perf: 缓存batch_all_memory_size引用消除重复查找 #4609
已合并
tangqunzhang创建于 8月28日
共 6 个文件变更+246-111
@@ -48,6 +48,43 @@ const std::string kOffline = "offline";
48const int32_t kReuseMaxOpNum = 10;48const int32_t kReuseMaxOpNum = 10;
49const int32_t kReuseMaxCharNum = 2000;49const int32_t kReuseMaxCharNum = 2000;
50 50 
51+struct ContinuousNodeLifeTimeItem {
52+ const ge::Node *node;
53+ int32_t in_index; // 本节点经由上级连续节点的哪个输入 anchor 到达(首输入 0 为 block 分配者)
54+ int32_t out_index; // 进入本节点所经的本节点输出 anchor 索引(org 自身为 -1)
55+};
56+ 
57+/// ref 节点(output 复用 input 透传)穿透:沿进入本节点的输出边 out_index 找到其
58+/// 复用的 input 的真实上游生产者及其输出边索引。仅该输出边对应的 ref 关系才属于
59+/// 当前遍历链,节点的其他输出(复用其他 input)属于别的链,不参与本链生命期计算。
60+/// 返回是否发生穿透。ref 判断优先于连续输入:透传时节点自身不拥有内存,
61+/// 其 input 上游才是内存归属者,避免生命期记为 ref 节点而非生产者导致踩踏。
62+bool TraverseRefOnOutEdge(const ge::Node *cur_node, const int32_t out_index, const ge::Node *&producer,
63+ int32_t &producer_out_index) {
64+ producer = nullptr;
65+ producer_out_index = -1;
66+ if ((out_index < 0) || (out_index >= static_cast<int32_t>(cur_node->GetAllOutDataAnchorsSize()))) {
67+ return false;
68+ }
69+ const auto out_anchor = cur_node->GetOutDataAnchor(out_index);
70+ if (out_anchor == nullptr) {
71+ return false;
72+ }
73+ int32_t reuse_in_index = -1;
74+ if (!ge::GraphUtils::IsRefFromInput(out_anchor, reuse_in_index)) {
75+ return false;
76+ }
77+ const auto ref_in_anchor = cur_node->GetInDataAnchor(reuse_in_index);
78+ if (ref_in_anchor != nullptr) {
79+ const auto peer_out_anchor = ref_in_anchor->GetPeerOutAnchor();
80+ if ((peer_out_anchor != nullptr) && (peer_out_anchor->GetOwnerNodeBarePtr() != nullptr)) {
81+ producer = peer_out_anchor->GetOwnerNodeBarePtr();
82+ producer_out_index = peer_out_anchor->GetIdx();
83+ }
84+ }
85+ return true; // 是 ref 边但上游无效时同样跳过本节点自身的记录
86+}
87+ 
51std::string FormatStreamEdgeName(const char *src_name, const char *dst_name) {88std::string FormatStreamEdgeName(const char *src_name, const char *dst_name) {
52 if ((src_name != nullptr) && (dst_name != nullptr)) {89 if ((src_name != nullptr) && (dst_name != nullptr)) {
53 return "[" + std::string(dst_name) + "<-" + std::string(src_name) + "] ";90 return "[" + std::string(dst_name) + "<-" + std::string(src_name) + "] ";
@@ -84,6 +121,7 @@ bool NotMatchNoReuseType(const std::set<std::string> &no_reuse_types, const std:
84}121}
85 122 
86} // namespace123} // namespace
124+ 
87namespace ge {125namespace ge {
88// Memory size is fixed and has nothing to do with different batches.126// Memory size is fixed and has nothing to do with different batches.
89bool SizeIndependentOfBatch(const std::string &node_type) {127bool SizeIndependentOfBatch(const std::string &node_type) {
@@ -503,7 +541,7 @@ Status BlockMemAssigner::GetOutAndWorkSpaceMem(std::vector<int64_t> &all_memory_
503 std::set<int64_t> exclude_merge_streams = GetStreamMergeAndOutStreams(compute_graph_);541 std::set<int64_t> exclude_merge_streams = GetStreamMergeAndOutStreams(compute_graph_);
504 for (const NodePtr &n : compute_graph_->GetAllNodes()) {542 for (const NodePtr &n : compute_graph_->GetAllNodes()) {
505 GetDiffStreamEdgeLife(n, exclude_merge_streams);543 GetDiffStreamEdgeLife(n, exclude_merge_streams);
506- GetContinuousNodeLifeTimeBegin(n.get(), n.get(), 0, 0U);544+ GetContinuousNodeLifeTimeBegin(n.get(), 0);
507 545 
508 auto node_op_desc = n->GetOpDescBarePtr();546 auto node_op_desc = n->GetOpDescBarePtr();
509 GE_ASSERT_NOTNULL(node_op_desc);547 GE_ASSERT_NOTNULL(node_op_desc);
@@ -514,6 +552,7 @@ Status BlockMemAssigner::GetOutAndWorkSpaceMem(std::vector<int64_t> &all_memory_
514 552 
515 std::string batch_label;553 std::string batch_label;
516 (void)ge::AttrUtils::GetStr(node_op_desc, ATTR_NAME_BATCH_LABEL, batch_label);554 (void)ge::AttrUtils::GetStr(node_op_desc, ATTR_NAME_BATCH_LABEL, batch_label);
555+ auto &batch_mem = batch_all_memory_size[batch_label];
517 556 
518 if (NodeUtils::IsLikeAtomicClean(n)) {557 if (NodeUtils::IsLikeAtomicClean(n)) {
519 atomic_addr_clean_id_ = node_op_desc->GetId();558 atomic_addr_clean_id_ = node_op_desc->GetId();
@@ -531,7 +570,7 @@ Status BlockMemAssigner::GetOutAndWorkSpaceMem(std::vector<int64_t> &all_memory_
531 " is invalid, "570 " is invalid, "
532 "maybe it is unknown shape node, Node_name:%s",571 "maybe it is unknown shape node, Node_name:%s",
533 size, node_op_desc->GetNamePtr());572 size, node_op_desc->GetNamePtr());
534- batch_all_memory_size[batch_label].emplace_back(size);573+ batch_mem.emplace_back(size);
535 batch_total_size[batch_label] += size;574 batch_total_size[batch_label] += size;
536 575 
537 if (!anchor_to_symbol_.empty()) {576 if (!anchor_to_symbol_.empty()) {
@@ -550,7 +589,7 @@ Status BlockMemAssigner::GetOutAndWorkSpaceMem(std::vector<int64_t> &all_memory_
550 }589 }
551 temp.clear();590 temp.clear();
552 GetNodeWorkSpaceSize(n, temp, batch_total_size[batch_label]);591 GetNodeWorkSpaceSize(n, temp, batch_total_size[batch_label]);
553- batch_all_memory_size[batch_label].insert(batch_all_memory_size[batch_label].cend(), temp.cbegin(), temp.cend());592+ batch_mem.insert(batch_mem.cend(), temp.cbegin(), temp.cend());
554 }593 }
555 HandleInStreamRedundantDependence(in_stream_edges_);594 HandleInStreamRedundantDependence(in_stream_edges_);
556 InsertStreamOutEdge();595 InsertStreamOutEdge();
@@ -674,85 +713,139 @@ void BlockMemAssigner::GetRefContinuousInputNodeAndFixedAddrPriorFlag(const std:
674/// h and j are nopading continuous input, g can't reuse with a,b,c713/// h and j are nopading continuous input, g can't reuse with a,b,c
675/// because their(d,e,f) memory will be replaced by g's memory (cascade continuous input)714/// because their(d,e,f) memory will be replaced by g's memory (cascade continuous input)
676/// so g's real life time begin is min of d,e,f715/// so g's real life time begin is min of d,e,f
677-void BlockMemAssigner::GetContinuousNodeLifeTimeBegin(const Node *const org_node, const Node *const node,716+void BlockMemAssigner::GetContinuousNodeLifeTimeBegin(const Node *const node, const int32_t in_index) {
678- const int32_t index, uint32_t depth) {717+ const auto org_node_desc = node->GetOpDescBarePtr();
679- ++depth;718+ if (org_node_desc == nullptr) {
680- GE_IF_BOOL_EXEC((depth > kMaxDepthNum), return);
681- 
682- bool is_nopading_input_continuous = false;
683- const auto node_op_desc = node->GetOpDescBarePtr();
684- GE_CHECK_NOTNULL_EXEC(node_op_desc, return);
685- const auto &org_node_desc = org_node->GetOpDescBarePtr();
686- GE_CHECK_NOTNULL_EXEC(org_node_desc, return);
687- (void)ge::AttrUtils::GetBool(node_op_desc, ATTR_NAME_NOPADDING_CONTINUOUS_INPUT, is_nopading_input_continuous);
688- if (is_nopading_input_continuous) {
689- for (const auto in_anchor : node->GetAllInDataAnchorsPtr()) {
690- const bool invalid_node = ((in_anchor == nullptr) || (in_anchor->GetPeerOutAnchor() == nullptr) ||
691- (in_anchor->GetPeerOutAnchor()->GetOwnerNodeBarePtr() == nullptr));
692- GE_IF_BOOL_EXEC(invalid_node, continue);
693- GetContinuousNodeLifeTimeBegin(org_node, in_anchor->GetPeerOutAnchor()->GetOwnerNodeBarePtr(),
694- in_anchor->GetIdx(), depth);
695- }
696- 
697- if (org_node == node) {
698- SetContinuousNodeLifeTimeBegin(node, node, 0U);
699- }
700- } else {
701- // 2 means has continuous input
702- GE_IF_BOOL_EXEC((depth < 2U), return);
703- auto it = cascade_min_life_time_.find(org_node_desc->GetNamePtr());
704- if (it == cascade_min_life_time_.end()) {
705- cascade_min_life_time_[org_node_desc->GetNamePtr()] = node_op_desc->GetId();
706- } else {
707- if (static_cast<size_t>(node_op_desc->GetId()) < it->second) {
708- it->second = node_op_desc->GetId();
709- }
710- }
711- // only set first node, continuous first input need alloc memory
712- if (index == 0) {
713- cascade_min_life_time_[node_op_desc->GetNamePtr()] = node_op_desc->GetId();
714- }
715- GELOGD("Find node:%s life time begin:%" PRId64 " by ref node:%s index:%d.", node_op_desc->GetNamePtr(),
716- node_op_desc->GetId(), org_node_desc->GetNamePtr(), index);
717- }
718- return;
719-}
720- 
721-void BlockMemAssigner::SetContinuousNodeLifeTimeBegin(const Node *const org_node, const Node *const node,
722- uint32_t depth) {
723- ++depth;
724- if (depth > kMaxDepthNum) {
725 return;719 return;
726 }720 }
727 721 
728- const auto node_op_desc = node->GetOpDescBarePtr();722+ // 迭代式遍历 nopadding 连续输入链
729- GE_CHECK_NOTNULL_EXEC(node_op_desc, return);723+ std::stack<ContinuousNodeLifeTimeItem> node_stack;
730- bool is_nopading_input_continuous = false;724+ std::unordered_set<const Node *> visited;
731- (void)ge::AttrUtils::GetBool(node_op_desc, ATTR_NAME_NOPADDING_CONTINUOUS_INPUT, is_nopading_input_continuous);725+ node_stack.push({node, in_index, -1});
732- if (is_nopading_input_continuous) {726+ bool is_continuous = false;
733- for (const auto in_anchor : node->GetAllInDataAnchorsPtr()) {727+ while (!node_stack.empty()) {
734- const bool invalid_node = (in_anchor == nullptr) || (in_anchor->GetPeerOutAnchor() == nullptr);728+ const auto cur_item = node_stack.top();
735- GE_IF_BOOL_EXEC(invalid_node, continue);729+ node_stack.pop();
736- const auto peer_in_node = in_anchor->GetPeerOutAnchor()->GetOwnerNodeBarePtr();730+ const auto cur_node = cur_item.node;
737- GE_CHECK_NOTNULL_EXEC(peer_in_node, continue);731+ const auto cur_op_desc = cur_node->GetOpDescBarePtr();
738- SetContinuousNodeLifeTimeBegin(org_node, peer_in_node, depth);732+ if (cur_op_desc == nullptr) {
733+ continue;
739 }734 }
740- } else {735+ 
741- // set min life time, only set first node736+ const bool cur_node_is_continuous = MemLayoutConflictUtil::IsNoPaddingContinuousInput(cur_node);
742- auto it = cascade_min_life_time_.find(node_op_desc->GetNamePtr());737+ if (cur_node == node) {
743- if (it != cascade_min_life_time_.end()) {738+ is_continuous = cur_node_is_continuous;
744- const auto org_node_desc = org_node->GetOpDescBarePtr();739+ }
745- GE_CHECK_NOTNULL_EXEC(org_node_desc, return);740+ 
741+ const Node *ref_producer = nullptr;
742+ int32_t producer_out_index = -1;
743+ if (TraverseRefOnOutEdge(cur_node, cur_item.out_index, ref_producer, producer_out_index)) {
744+ if (ref_producer != nullptr) {
745+ node_stack.push({ref_producer, cur_item.in_index, producer_out_index});
746+ }
747+ continue;
748+ }
749+ 
750+ if (cur_node_is_continuous) {
751+ if (!visited.insert(cur_node).second) {
752+ continue;
753+ }
754+ for (const auto in_anchor : cur_node->GetAllInDataAnchorsPtr()) {
755+ if (in_anchor == nullptr) {
756+ continue;
757+ }
758+ const auto peer_out_anchor = in_anchor->GetPeerOutAnchor();
759+ const bool invalid_node = ((peer_out_anchor == nullptr) || (peer_out_anchor->GetOwnerNodeBarePtr() == nullptr));
760+ if (!invalid_node) {
761+ node_stack.push({peer_out_anchor->GetOwnerNodeBarePtr(), in_anchor->GetIdx(), peer_out_anchor->GetIdx()});
762+ }
763+ }
764+ } else {
765+ // 起始节点自身不记录(连续输入链至少展开一层后到达的节点才记录生命期)
766+ if (cur_node == node) {
767+ continue;
768+ }
769+ auto it = cascade_min_life_time_.find(org_node_desc->GetNamePtr());
770+ if (it == cascade_min_life_time_.end()) {
771+ cascade_min_life_time_[org_node_desc->GetNamePtr()] = cur_op_desc->GetId();
772+ } else {
773+ if (static_cast<size_t>(cur_op_desc->GetId()) < it->second) {
774+ it->second = cur_op_desc->GetId();
775+ }
776+ }
777+ // only set first node, continuous first input need alloc memory
778+ if (cur_item.in_index == 0) {
779+ cascade_min_life_time_[cur_op_desc->GetNamePtr()] = cur_op_desc->GetId();
780+ }
781+ GELOGD("Find node:%s life time begin:%" PRId64 " by ref node:%s index:%d.", cur_op_desc->GetNamePtr(),
782+ cur_op_desc->GetId(), org_node_desc->GetNamePtr(), cur_item.in_index);
783+ }
784+ }
785+ 
786+ // 仅连续输入节点需要传播生命期,非连续节点跳过避免无谓的栈操作和 map 查找
787+ if (is_continuous) {
788+ SetContinuousNodeLifeTimeBegin(node);
789+ }
790+}
791+ 
792+void BlockMemAssigner::SetContinuousNodeLifeTimeBegin(const Node *const node) {
793+ const auto org_node_desc = node->GetOpDescBarePtr();
794+ if (org_node_desc == nullptr) {
795+ return;
796+ }
797+ 
798+ // 迭代式遍历 nopadding 连续输入链,visited 防环,替代递归深度限制
799+ std::stack<std::pair<const Node *, int32_t>> node_stack;
800+ std::unordered_set<const Node *> visited;
801+ node_stack.push({node, -1});
802+ while (!node_stack.empty()) {
803+ const auto [cur_node, cur_out_index] = node_stack.top();
804+ node_stack.pop();
805+ const auto cur_op_desc = cur_node->GetOpDescBarePtr();
806+ if (cur_op_desc == nullptr) {
807+ continue;
808+ }
809+ if (!visited.insert(cur_node).second) {
810+ continue;
811+ }
812+ 
813+ const Node *ref_producer = nullptr;
814+ int32_t producer_out_index = -1;
815+ if (TraverseRefOnOutEdge(cur_node, cur_out_index, ref_producer, producer_out_index)) {
816+ if (ref_producer != nullptr) {
817+ node_stack.push({ref_producer, producer_out_index});
818+ }
819+ continue;
820+ }
821+ 
822+ if (MemLayoutConflictUtil::IsNoPaddingContinuousInput(cur_node)) {
823+ for (const auto in_anchor : cur_node->GetAllInDataAnchorsPtr()) {
824+ if (in_anchor == nullptr) {
825+ continue;
826+ }
827+ const auto peer_out_anchor = in_anchor->GetPeerOutAnchor();
828+ const bool invalid_node = ((peer_out_anchor == nullptr) || (peer_out_anchor->GetOwnerNodeBarePtr() == nullptr));
829+ if (!invalid_node) {
830+ node_stack.push({peer_out_anchor->GetOwnerNodeBarePtr(), peer_out_anchor->GetIdx()});
831+ }
832+ }
833+ } else {
834+ // set min life time, only set first node
835+ auto it = cascade_min_life_time_.find(cur_op_desc->GetNamePtr());
836+ if (it == cascade_min_life_time_.end()) {
837+ continue;
838+ }
746 const auto it_org = cascade_min_life_time_.find(org_node_desc->GetNamePtr());839 const auto it_org = cascade_min_life_time_.find(org_node_desc->GetNamePtr());
747 if (it_org != cascade_min_life_time_.cend()) {840 if (it_org != cascade_min_life_time_.cend()) {
748- GELOGI("Node:%s set min life time begin from %zu to %zu by ref node:%s.", node->GetNamePtr(), it->second,841+ GELOGI("Node:%s set min life time begin from %zu to %zu by ref node:%s.", cur_node->GetNamePtr(), it->second,
749 it_org->second, org_node_desc->GetNamePtr());842 it_org->second, org_node_desc->GetNamePtr());
750 it->second = it_org->second;843 it->second = it_org->second;
751 }844 }
752 }845 }
753 }846 }
754- return;
755}847}
848+ 
756/*849/*
757 * 1. NoPadding连续输入,仅首个输入分配一个block,所有输入使用这一个block850 * 1. NoPadding连续输入,仅首个输入分配一个block,所有输入使用这一个block
758 * 2. 带Padding连续输入,每个输入有自己的block,连续在一起。851 * 2. 带Padding连续输入,每个输入有自己的block,连续在一起。
@@ -385,10 +385,9 @@ class BlockMemAssigner : public MemAssigner {
385 /// @brief Cascade memory scenarios to obtain the actual life time begin of continuous input memory385 /// @brief Cascade memory scenarios to obtain the actual life time begin of continuous input memory
386 /// @return void386 /// @return void
387 /// @author387 /// @author
388- void GetContinuousNodeLifeTimeBegin(const Node *const org_node, const Node *const node, const int32_t index,388+ void GetContinuousNodeLifeTimeBegin(const Node *const node, const int32_t in_index);
389- uint32_t depth);
390 389 
391- void SetContinuousNodeLifeTimeBegin(const Node *const org_node, const Node *const node, uint32_t depth);390+ void SetContinuousNodeLifeTimeBegin(const Node *const node);
392 391 
393 void GetRefContinuousInputNodeAndFixedAddrPriorFlag(const std::string &symbol, const std::list<NodeIndexIO> &anchors);392 void GetRefContinuousInputNodeAndFixedAddrPriorFlag(const std::string &symbol, const std::list<NodeIndexIO> &anchors);
394 393 
@@ -62,15 +62,6 @@ bool MemReuseUtils::IsMergeNode(const Node *node) {
62 return (node->GetType() == STREAMMERGE) || (node->GetType() == MERGE);62 return (node->GetType() == STREAMMERGE) || (node->GetType() == MERGE);
63}63}
64 64 
65-bool MemReuseUtils::IsNoPaddingContinuousInput(const Node *node) {
66- GE_ASSERT_NOTNULL(node);
67- const auto op_desc = node->GetOpDescBarePtr();
68- GE_ASSERT_NOTNULL(op_desc);
69- bool is_nopading_input_continuous = false;
70- (void)ge::AttrUtils::GetBool(op_desc, ATTR_NAME_NOPADDING_CONTINUOUS_INPUT, is_nopading_input_continuous);
71- return is_nopading_input_continuous && (node->GetAllInDataAnchorsSize() > 1U);
72-}
73- 
74Status MemReuseUtils::GetOutputNoAlignSize(const ge::OpDesc &desc, uint32_t index, size_t &size) {65Status MemReuseUtils::GetOutputNoAlignSize(const ge::OpDesc &desc, uint32_t index, size_t &size) {
75 const auto tensor_desc = desc.GetOutputDesc(index);66 const auto tensor_desc = desc.GetOutputDesc(index);
76 GE_ASSERT_SUCCESS(GetNoAlignSize(tensor_desc, size), "node: %s, output index: %u", desc.GetNamePtr(), index);67 GE_ASSERT_SUCCESS(GetNoAlignSize(tensor_desc, size), "node: %s, output index: %u", desc.GetNamePtr(), index);
@@ -31,7 +31,6 @@ class MemReuseUtils {
31 static void SetStreamId(ge::OpDesc *const desc, int64_t stream_id);31 static void SetStreamId(ge::OpDesc *const desc, int64_t stream_id);
32 static bool IsMergeNode(const NodePtr &node);32 static bool IsMergeNode(const NodePtr &node);
33 static bool IsMergeNode(const Node *node);33 static bool IsMergeNode(const Node *node);
34- static bool IsNoPaddingContinuousInput(const Node *node);
35 static Status GetNoAlignSize(const GeTensorDesc &tensor, size_t &size);34 static Status GetNoAlignSize(const GeTensorDesc &tensor, size_t &size);
36 static Status GetOutputNoAlignSize(const ge::OpDesc &desc, uint32_t index, size_t &size);35 static Status GetOutputNoAlignSize(const ge::OpDesc &desc, uint32_t index, size_t &size);
37 static bool IsAllOutRefAllInput(const NodePtr &node);36 static bool IsAllOutRefAllInput(const NodePtr &node);
@@ -2078,4 +2078,85 @@ TEST_F(UtestBlockMemAssigner, AssignWorkSpaceMemoryWithReuse_BasicOp) {
2078 EXPECT_EQ(assigner.AssignWorkSpaceMemoryWithReuse(node, ranges), SUCCESS);2078 EXPECT_EQ(assigner.AssignWorkSpaceMemoryWithReuse(node, ranges), SUCCESS);
2079}2079}
2080 2080 
2081+/*
2082+ * 级联+ref 场景(topo序: a ref b pc1 d pc2):
2083+ *
2084+ * a b
2085+ * | |
2086+ * ref |
2087+ * | |
2088+ * └────┘
2089+ * pc1(连续输入)
2090+ * |
2091+ * d |
2092+ * | |
2093+ * └────┘
2094+ * pc2(连续输入)
2095+ *
2096+ * a 经 ref(REF透传) 接 pc1.in0,pc1.in1 <- b;pc2.in0 <- d,pc2.in1 <- pc1,
2097+ * pc1/pc2 均为 nopadding 连续输入节点。
2098+ *
2099+ * pc2 为级联连续输入节点,其 output 走 ApplyMemory 消费 cascade[pc2]。
2100+ * 遍历 pc2 输入链:d(非连续)、pc1(连续)→ref(REF)→a(非连续)、b(非连续),
2101+ * 穿透 ref 后 cascade[pc2] = min(d, a, b) = a_id;不穿透则记为 ref_id(>a_id)。
2102+ * Set 传播后 cascade[d] = cascade[pc2] = a_id,d 的 life_time_begin_ 提前到 a_id,
2103+ * 验证级联展开与 ref 穿透的正确性,防止 [a, ref) 窗口内误复用导致 a 的数据被踩踏。
2104+ */
2105+TEST_F(UtestBlockMemAssigner, NoPaddingContinuousInputThroughRefNodeLifeTimeBegin) {
2106+ const auto refnode_cfg = OP_CFG(ASSIGNADD).Attr(ATTR_NAME_REFERENCE, true).InNames({"ref"}).OutNames({"ref"});
2107+ DEF_GRAPH(g1) {
2108+ CHAIN(NODE("a", RELU)->NODE("ref", refnode_cfg)->NODE("pc1", PHONYCONCAT)->EDGE(0, 1)->NODE("pc2", PHONYCONCAT));
2109+ CHAIN(NODE("b", RELU)->EDGE(0, 1)->NODE("pc1", PHONYCONCAT));
2110+ CHAIN(NODE("d", RELU)->EDGE(0, 0)->NODE("pc2", PHONYCONCAT));
2111+ };
2112+ auto graph = ToComputeGraph(g1);
2113+ MemConflictShareGraph::SetNoPaddingContinuousInput(graph, "pc1");
2114+ MemConflictShareGraph::SetNoPaddingContinuousInput(graph, "pc2");
2115+ MemConflictShareGraph::SetSizeForAllNodes(graph);
2116+ AttrUtils::SetBool(graph, ATTR_NAME_NO_NEED_DYNAMIC_SHAPE_PARTITION, true);
2117+ // ref 的 output 标记复用 input0,与 ref 属性配合建立 input/output 同 symbol
2118+ const auto ref_node = graph->FindNode("ref");
2119+ ASSERT_NE(ref_node, nullptr);
2120+ auto ref_out_tensor = ref_node->GetOpDesc()->MutableOutputDesc(0);
2121+ ASSERT_NE(ref_out_tensor, nullptr);
2122+ TensorUtils::SetReuseInput(*ref_out_tensor, true);
2123+ TensorUtils::SetReuseInputIndex(*ref_out_tensor, 0U);
2124+ (void)graph->TopologicalSorting();
2125+ 
2126+ MemAssistInfo mem_assist_info;
2127+ mem_assist_info.compute_graph = graph;
2128+ auto ret = GraphUtils::GetRefMapping(graph, mem_assist_info.symbol_to_anchors, mem_assist_info.anchor_to_symbol);
2129+ EXPECT_EQ(ret, SUCCESS);
2130+ BlockMemAssigner::PreparationForAssign(mem_assist_info);
2131+ 
2132+ std::vector<int64_t> ranges;
2133+ BinaryBlockMemAssigner assigner(mem_assist_info);
2134+ assigner.SetReuseStrategy(ReuseStrategy{false, false, false, true});
2135+ assigner.GetMemoryRanges(ranges);
2136+ EXPECT_EQ(assigner.AssignMemoryWithReuse(ranges), SUCCESS);
2137+ assigner.SetOpMemOffset(false);
2138+ 
2139+ const auto a = graph->FindNode("a");
2140+ const auto d = graph->FindNode("d");
2141+ ASSERT_NE(a, nullptr);
2142+ ASSERT_NE(d, nullptr);
2143+ const auto expected_begin = static_cast<size_t>(a->GetOpDesc()->GetId());
2144+ const auto blocks = assigner.GetMemoryBlocks();
2145+ bool has_checked = false;
2146+ for (const auto block : blocks) {
2147+ if (block == nullptr) {
2148+ continue;
2149+ }
2150+ for (const auto &node : block->node_type_index_list_) {
2151+ if ((node.mem_type_ == OpMemoryType::kOutput) && (node.node_ != nullptr) && (node.node_->GetName() == "d")) {
2152+ // 穿透 ref 后 cascade[pc2] = min(d, a, b) = a_id,Set 传播 cascade[d] = a_id,
2153+ // d 的 life_time_begin_ 应精确等于 a 的 topo id(而非 ref 的 topo id)
2154+ EXPECT_EQ(node.life_time_begin_, expected_begin);
2155+ has_checked = true;
2156+ }
2157+ }
2158+ }
2159+ EXPECT_TRUE(has_checked);
2160+}
2161+ 
2081} // namespace ge2162} // namespace ge
@@ -342,34 +342,6 @@ TEST_F(UtestReuseChecker, ModifyPhonyConcatLastInputMemSize) {
342 EXPECT_EQ(memory_assigner.AssignMemory(mem_offset, zero_memory_size), GRAPH_SUCCESS);342 EXPECT_EQ(memory_assigner.AssignMemory(mem_offset, zero_memory_size), GRAPH_SUCCESS);
343}343}
344 344 
345-TEST_F(UtestReuseChecker, MemReuseUtils_IsNoPaddingContinuousInput_NotSet) {
346- auto compute_graph = std::make_shared<ComputeGraph>("test_nopadding");
347- auto op_desc = std::make_shared<OpDesc>("test_op", "Add");
348- op_desc->AddInputDesc(GeTensorDesc());
349- op_desc->AddInputDesc(GeTensorDesc());
350- auto node = compute_graph->AddNode(op_desc);
351- EXPECT_FALSE(MemReuseUtils::IsNoPaddingContinuousInput(node.get()));
352-}
353- 
354-TEST_F(UtestReuseChecker, MemReuseUtils_IsNoPaddingContinuousInput_SetButSingleInput) {
355- auto compute_graph = std::make_shared<ComputeGraph>("test_nopadding2");
356- auto op_desc = std::make_shared<OpDesc>("test_op", "Add");
357- op_desc->AddInputDesc(GeTensorDesc());
358- AttrUtils::SetBool(op_desc, ATTR_NAME_NOPADDING_CONTINUOUS_INPUT, true);
359- auto node = compute_graph->AddNode(op_desc);
360- EXPECT_FALSE(MemReuseUtils::IsNoPaddingContinuousInput(node.get()));
361-}
362- 
363-TEST_F(UtestReuseChecker, MemReuseUtils_IsNoPaddingContinuousInput_SetMultiInput) {
364- auto compute_graph = std::make_shared<ComputeGraph>("test_nopadding3");
365- auto op_desc = std::make_shared<OpDesc>("test_op", "Add");
366- op_desc->AddInputDesc(GeTensorDesc());
367- op_desc->AddInputDesc(GeTensorDesc());
368- AttrUtils::SetBool(op_desc, ATTR_NAME_NOPADDING_CONTINUOUS_INPUT, true);
369- auto node = compute_graph->AddNode(op_desc);
370- EXPECT_TRUE(MemReuseUtils::IsNoPaddingContinuousInput(node.get()));
371-}
372- 
373TEST_F(UtestReuseChecker, MemReuseUtils_GetOutputNoAlignSize_Success) {345TEST_F(UtestReuseChecker, MemReuseUtils_GetOutputNoAlignSize_Success) {
374 auto op_desc = std::make_shared<OpDesc>("test_op", "Add");346 auto op_desc = std::make_shared<OpDesc>("test_op", "Add");
375 op_desc->AddOutputDesc(GeTensorDesc(GeShape({2, 3}), FORMAT_ND, DT_FLOAT));347 op_desc->AddOutputDesc(GeTensorDesc(GeShape({2, 3}), FORMAT_ND, DT_FLOAT));