已合并
perf: 缓存batch_all_memory_size引用消除重复查找 #4609
tangqunzhang创建于 8月28日
perf: 缓存batch_all_memory_size引用消除重复查找 #4609
已合并
共 6 个文件变更+246-111
| @@ -48,6 +48,43 @@ const std::string kOffline = "offline"; | |||
| 48 | const int32_t kReuseMaxOpNum = 10; | 48 | const int32_t kReuseMaxOpNum = 10; |
| 49 | const int32_t kReuseMaxCharNum = 2000; | 49 | const 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 | + | ||
| 51 | std::string FormatStreamEdgeName(const char *src_name, const char *dst_name) { | 88 | std::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 | } // namespace | 123 | } // namespace |
| 124 | + | ||
| 87 | namespace ge { | 125 | namespace 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. |
| 89 | bool SizeIndependentOfBatch(const std::string &node_type) { | 127 | bool 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,c | 713 | /// 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,f | 715 | /// 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 node | 736 | + 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,所有输入使用这一个block | 850 | * 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 memory | 385 | /// @brief Cascade memory scenarios to obtain the actual life time begin of continuous input memory |
| 386 | /// @return void | 386 | /// @return void |
| 387 | /// @author | 387 | /// @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 | - | ||
| 74 | Status MemReuseUtils::GetOutputNoAlignSize(const ge::OpDesc &desc, uint32_t index, size_t &size) { | 65 | Status 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 ge | 2162 | } // 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 | - | ||
| 373 | TEST_F(UtestReuseChecker, MemReuseUtils_GetOutputNoAlignSize_Success) { | 345 | TEST_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)); |