已合并
perf: 优化stream间内存复用性能 #4569
tangqunzhang创建于 8月26日
perf: 优化stream间内存复用性能 #4569
已合并
共 17 个文件变更+106-142
| @@ -925,7 +925,7 @@ Status AtomicMemoryAssigner::AssignFusionAtomicWorkspaceMemory( | |||
| 925 | return SUCCESS; | 925 | return SUCCESS; |
| 926 | } | 926 | } |
| 927 | 927 | ||
| 928 | -void AtomicMemoryAssigner::AlignMemOffset(const int64_t &mem_align_size, int64_t memory_type) { | 928 | +void AtomicMemoryAssigner::AlignMemOffset(int64_t mem_align_size, int64_t memory_type) { |
| 929 | if (mem_align_size <= 0) { | 929 | if (mem_align_size <= 0) { |
| 930 | return; | 930 | return; |
| 931 | } | 931 | } |
| @@ -193,7 +193,7 @@ class AtomicMemoryAssigner { | |||
| 193 | std::map<int64_t, std::vector<int64_t>> &mem_type_to_real_atomic_sizes); | 193 | std::map<int64_t, std::vector<int64_t>> &mem_type_to_real_atomic_sizes); |
| 194 | Status AppendAddrSizeToMemSetOp(const NodePtr &memset_node, const MemsetNodeAddrAndAttr &addr_type) const; | 194 | Status AppendAddrSizeToMemSetOp(const NodePtr &memset_node, const MemsetNodeAddrAndAttr &addr_type) const; |
| 195 | Status AppendAttrsToMemSetOp(const NodePtr &memset_node, const MemsetNodeAddrAndAttr &addr_type) const; | 195 | Status AppendAttrsToMemSetOp(const NodePtr &memset_node, const MemsetNodeAddrAndAttr &addr_type) const; |
| 196 | - void AlignMemOffset(const int64_t &mem_align_size, int64_t memory_type); | 196 | + void AlignMemOffset(int64_t mem_align_size, int64_t memory_type); |
| 197 | Status UpdateParentNodeOutputOffset(const ge::NodePtr &node, int64_t output_index, int64_t offset) const; | 197 | Status UpdateParentNodeOutputOffset(const ge::NodePtr &node, int64_t output_index, int64_t offset) const; |
| 198 | Status GetMemoryAssignmentStatus(const ge::NodePtr &node, int64_t output_index, bool &is_mem_assigned) const; | 198 | Status GetMemoryAssignmentStatus(const ge::NodePtr &node, int64_t output_index, bool &is_mem_assigned) const; |
| 199 | 199 | ||
| @@ -48,6 +48,13 @@ 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 | +std::string FormatStreamEdgeName(const char *src_name, const char *dst_name) { | ||
| 52 | + if ((src_name != nullptr) && (dst_name != nullptr)) { | ||
| 53 | + return "[" + std::string(dst_name) + "<-" + std::string(src_name) + "] "; | ||
| 54 | + } | ||
| 55 | + return ""; | ||
| 56 | +} | ||
| 57 | + | ||
| 51 | int64_t GetStreamId(const ge::OpDesc *const desc) { | 58 | int64_t GetStreamId(const ge::OpDesc *const desc) { |
| 52 | return ge::MemReuseUtils::GetStreamId(desc); | 59 | return ge::MemReuseUtils::GetStreamId(desc); |
| 53 | } | 60 | } |
| @@ -247,52 +254,32 @@ void BlockMemAssigner::InsertStreamOutEdge() { | |||
| 247 | * 3<-2 保留 | 254 | * 3<-2 保留 |
| 248 | * 可以简单记为:id差越小越好 | 255 | * 可以简单记为:id差越小越好 |
| 249 | */ | 256 | */ |
| 250 | -void BlockMemAssigner::InsertStreamInEdge(const EdgeLife &new_in_edge, const int64_t src_stream_id, | 257 | +void BlockMemAssigner::InsertStreamInEdge(std::set<EdgeLife, CompareEdgeLife> &in_edge_set, const EdgeLife &new_in_edge, |
| 251 | - const int64_t dst_stream_id, const char *src_name, const char *dst_name) { | 258 | + const int64_t src_stream_id, const int64_t dst_stream_id, |
| 252 | - auto &in_edge_set = in_stream_edges_[dst_stream_id][src_stream_id]; | 259 | + const std::pair<const char *, const char *> &node_names) { |
| 253 | const auto old_in_edge_iter = in_edge_set.find(new_in_edge); | 260 | const auto old_in_edge_iter = in_edge_set.find(new_in_edge); |
| 254 | if (old_in_edge_iter != in_edge_set.end()) { | 261 | if (old_in_edge_iter != in_edge_set.end()) { |
| 255 | if (old_in_edge_iter->peer_node_id < new_in_edge.peer_node_id) { | 262 | if (old_in_edge_iter->peer_node_id < new_in_edge.peer_node_id) { |
| 256 | const auto old_peer_node_id = old_in_edge_iter->peer_node_id; | 263 | const auto old_peer_node_id = old_in_edge_iter->peer_node_id; |
| 257 | in_edge_set.erase(old_in_edge_iter); // after erase, cannot use old_peer_node_id below | 264 | in_edge_set.erase(old_in_edge_iter); // after erase, cannot use old_peer_node_id below |
| 258 | in_edge_set.insert(new_in_edge); | 265 | in_edge_set.insert(new_in_edge); |
| 259 | - if ((src_name != nullptr) && (dst_name != nullptr)) { | 266 | + GELOGI("[StreamEdge]In depend Node: %sstream_id:[%" PRId64 "<-%" PRId64 |
| 260 | - GELOGI("[StreamEdge]In depend Node: [%s<-%s] stream_id:[%" PRId64 "<-%" PRId64 | 267 | + "] life_time:[%zu<-%zu], erase and insert," |
| 261 | - "] life_time:[%zu<-%zu], erase and insert," | 268 | + " old_peer_node_id[%zu].", |
| 262 | - " old_peer_node_id[%zu].", | 269 | + FormatStreamEdgeName(node_names.first, node_names.second).c_str(), dst_stream_id, src_stream_id, |
| 263 | - dst_name, src_name, dst_stream_id, src_stream_id, new_in_edge.node_id, new_in_edge.peer_node_id, | 270 | + new_in_edge.node_id, new_in_edge.peer_node_id, old_peer_node_id); |
| 264 | - old_peer_node_id); | ||
| 265 | - } else { | ||
| 266 | - GELOGI("[StreamEdge]In depend Node: stream_id:[%" PRId64 "<-%" PRId64 | ||
| 267 | - "] life_time:[%zu<-%zu], erase and insert," | ||
| 268 | - " old_peer_node_id[%zu].", | ||
| 269 | - dst_stream_id, src_stream_id, new_in_edge.node_id, new_in_edge.peer_node_id, old_peer_node_id); | ||
| 270 | - } | ||
| 271 | } else { | 271 | } else { |
| 272 | - if ((src_name != nullptr) && (dst_name != nullptr)) { | 272 | + GELOGI("[StreamEdge]In depend Node: %sstream_id:[%" PRId64 "<-%" PRId64 |
| 273 | - GELOGI("[StreamEdge]In depend Node: [%s<-%s] stream_id:[%" PRId64 "<-%" PRId64 | 273 | + "] life_time:[%zu<-%zu], not erase," |
| 274 | - "] life_time:[%zu<-%zu], not erase, not insert, " | 274 | + " not insert, old_peer_node_id[%zu] >= new_peer_node_id[%zu].", |
| 275 | - "old_peer_node_id[%zu] >= new_peer_node_id[%zu].", | 275 | + FormatStreamEdgeName(node_names.first, node_names.second).c_str(), dst_stream_id, src_stream_id, |
| 276 | - dst_name, src_name, dst_stream_id, src_stream_id, new_in_edge.node_id, new_in_edge.peer_node_id, | 276 | + new_in_edge.node_id, new_in_edge.peer_node_id, old_in_edge_iter->peer_node_id, new_in_edge.peer_node_id); |
| 277 | - old_in_edge_iter->peer_node_id, new_in_edge.peer_node_id); | ||
| 278 | - } else { | ||
| 279 | - GELOGI("[StreamEdge]In depend Node: stream_id:[%" PRId64 "<-%" PRId64 | ||
| 280 | - "] life_time:[%zu<-%zu], not erase, not insert, " | ||
| 281 | - "old_peer_node_id[%zu] >= new_peer_node_id[%zu].", | ||
| 282 | - dst_stream_id, src_stream_id, new_in_edge.node_id, new_in_edge.peer_node_id, | ||
| 283 | - old_in_edge_iter->peer_node_id, new_in_edge.peer_node_id); | ||
| 284 | - } | ||
| 285 | } | 277 | } |
| 286 | } else { | 278 | } else { |
| 287 | in_edge_set.insert(new_in_edge); | 279 | in_edge_set.insert(new_in_edge); |
| 288 | - if ((src_name != nullptr) && (dst_name != nullptr)) { | 280 | + GELOGI("[StreamEdge]In depend Node: %sstream_id:[%" PRId64 "<-%" PRId64 "] life_time:[%zu<-%zu], only insert.", |
| 289 | - GELOGI("[StreamEdge]In depend Node: [%s<-%s] stream_id:[%" PRId64 "<-%" PRId64 | 281 | + FormatStreamEdgeName(node_names.first, node_names.second).c_str(), dst_stream_id, src_stream_id, |
| 290 | - "] life_time:[%zu<-%zu], only insert.", | 282 | + new_in_edge.node_id, new_in_edge.peer_node_id); |
| 291 | - dst_name, src_name, dst_stream_id, src_stream_id, new_in_edge.node_id, new_in_edge.peer_node_id); | ||
| 292 | - } else { | ||
| 293 | - GELOGI("[StreamEdge]In depend Node: stream_id:[%" PRId64 "<-%" PRId64 "] life_time:[%zu<-%zu], only insert.", | ||
| 294 | - dst_stream_id, src_stream_id, new_in_edge.node_id, new_in_edge.peer_node_id); | ||
| 295 | - } | ||
| 296 | } | 283 | } |
| 297 | } | 284 | } |
| 298 | 285 | ||
| @@ -356,7 +343,7 @@ void BlockMemAssigner::AddInStreamEdge(const ge::OpDesc *const node_desc, const | |||
| 356 | if (old_edge_it != in_edge_set.end()) { | 343 | if (old_edge_it != in_edge_set.end()) { |
| 357 | EraseIntersectedEdge(in_edge_set, *old_edge_it, new_in_edge, third_stream_id, stream_id); | 344 | EraseIntersectedEdge(in_edge_set, *old_edge_it, new_in_edge, third_stream_id, stream_id); |
| 358 | } | 345 | } |
| 359 | - InsertStreamInEdge(new_in_edge, third_stream_id, stream_id); | 346 | + InsertStreamInEdge(in_edge_set, new_in_edge, third_stream_id, stream_id); |
| 360 | } | 347 | } |
| 361 | } | 348 | } |
| 362 | 349 | ||
| @@ -371,6 +358,7 @@ void BlockMemAssigner::GetDiffStreamEdgeLife(const NodePtr &node, const std::set | |||
| 371 | if (NodeUtils::IsLikeAtomicClean(node) || (node_desc->GetOpKernelLibName() == kEngineNameGeLocal)) { | 358 | if (NodeUtils::IsLikeAtomicClean(node) || (node_desc->GetOpKernelLibName() == kEngineNameGeLocal)) { |
| 372 | return; | 359 | return; |
| 373 | } | 360 | } |
| 361 | + const auto stream_id = GetStreamId(node_desc); | ||
| 374 | for (const auto &out_anchor : node->GetAllOutAnchors()) { | 362 | for (const auto &out_anchor : node->GetAllOutAnchors()) { |
| 375 | GE_CHECK_NOTNULL_JUST_RETURN(out_anchor); | 363 | GE_CHECK_NOTNULL_JUST_RETURN(out_anchor); |
| 376 | for (auto const peer_in_anchor : out_anchor->GetPeerAnchorsPtr()) { | 364 | for (auto const peer_in_anchor : out_anchor->GetPeerAnchorsPtr()) { |
| @@ -379,7 +367,6 @@ void BlockMemAssigner::GetDiffStreamEdgeLife(const NodePtr &node, const std::set | |||
| 379 | GE_CHECK_NOTNULL_JUST_RETURN(peer_node); | 367 | GE_CHECK_NOTNULL_JUST_RETURN(peer_node); |
| 380 | const auto peer_in_node_desc = peer_node->GetOpDescBarePtr(); | 368 | const auto peer_in_node_desc = peer_node->GetOpDescBarePtr(); |
| 381 | GE_CHECK_NOTNULL_JUST_RETURN(peer_in_node_desc); | 369 | GE_CHECK_NOTNULL_JUST_RETURN(peer_in_node_desc); |
| 382 | - const auto stream_id = GetStreamId(node_desc); | ||
| 383 | const auto peer_in_stream_id = GetStreamId(peer_in_node_desc); | 370 | const auto peer_in_stream_id = GetStreamId(peer_in_node_desc); |
| 384 | if (stream_id == peer_in_stream_id) { | 371 | if (stream_id == peer_in_stream_id) { |
| 385 | continue; | 372 | continue; |
| @@ -395,8 +382,9 @@ void BlockMemAssigner::GetDiffStreamEdgeLife(const NodePtr &node, const std::set | |||
| 395 | const auto node_id = static_cast<size_t>(node_desc->GetId()); | 382 | const auto node_id = static_cast<size_t>(node_desc->GetId()); |
| 396 | const auto peer_node_id = static_cast<size_t>(peer_in_node_desc->GetId()); | 383 | const auto peer_node_id = static_cast<size_t>(peer_in_node_desc->GetId()); |
| 397 | const EdgeLife new_in_edge{peer_node_id, node_id}; // 从peer_node看,由node连接进来的边称为入边 | 384 | const EdgeLife new_in_edge{peer_node_id, node_id}; // 从peer_node看,由node连接进来的边称为入边 |
| 398 | - InsertStreamInEdge(new_in_edge, stream_id, peer_in_stream_id, node_desc->GetNamePtr(), | 385 | + auto &in_edge_set = in_stream_edges_[peer_in_stream_id][stream_id]; |
| 399 | - peer_in_node_desc->GetNamePtr()); | 386 | + InsertStreamInEdge(in_edge_set, new_in_edge, stream_id, peer_in_stream_id, |
| 387 | + {node_desc->GetNamePtr(), peer_in_node_desc->GetNamePtr()}); | ||
| 400 | AddInStreamEdge(peer_in_node_desc, node_desc); | 388 | AddInStreamEdge(peer_in_node_desc, node_desc); |
| 401 | } | 389 | } |
| 402 | } | 390 | } |
| @@ -544,11 +532,7 @@ Status BlockMemAssigner::GetOutAndWorkSpaceMem(std::vector<int64_t> &all_memory_ | |||
| 544 | "maybe it is unknown shape node, Node_name:%s", | 532 | "maybe it is unknown shape node, Node_name:%s", |
| 545 | size, node_op_desc->GetNamePtr()); | 533 | size, node_op_desc->GetNamePtr()); |
| 546 | batch_all_memory_size[batch_label].emplace_back(size); | 534 | batch_all_memory_size[batch_label].emplace_back(size); |
| 547 | - if (batch_total_size.find(batch_label) == batch_total_size.end()) { | 535 | + batch_total_size[batch_label] += size; |
| 548 | - batch_total_size[batch_label] = size; | ||
| 549 | - } else { | ||
| 550 | - batch_total_size[batch_label] += size; | ||
| 551 | - } | ||
| 552 | 536 | ||
| 553 | if (!anchor_to_symbol_.empty()) { | 537 | if (!anchor_to_symbol_.empty()) { |
| 554 | auto iter1 = anchor_to_symbol_.find(NodeIndexIO(n.get(), out_anchor->GetIdx(), kOut).ToString()); | 538 | auto iter1 = anchor_to_symbol_.find(NodeIndexIO(n.get(), out_anchor->GetIdx(), kOut).ToString()); |
| @@ -378,8 +378,9 @@ class BlockMemAssigner : public MemAssigner { | |||
| 378 | void GetDiffStreamEdgeLife(const NodePtr &node, const std::set<int64_t> &exclude_merge_streams); | 378 | void GetDiffStreamEdgeLife(const NodePtr &node, const std::set<int64_t> &exclude_merge_streams); |
| 379 | void AddInStreamEdge(const ge::OpDesc *const node_desc, const ge::OpDesc *const in_node_desc); | 379 | void AddInStreamEdge(const ge::OpDesc *const node_desc, const ge::OpDesc *const in_node_desc); |
| 380 | void InsertStreamOutEdge(); | 380 | void InsertStreamOutEdge(); |
| 381 | - void InsertStreamInEdge(const EdgeLife &new_in_edge, const int64_t src_stream_id, const int64_t dst_stream_id, | 381 | + void InsertStreamInEdge(std::set<EdgeLife, CompareEdgeLife> &in_edge_set, const EdgeLife &new_in_edge, |
| 382 | - const char *src_name = nullptr, const char *dst_name = nullptr); | 382 | + const int64_t src_stream_id, const int64_t dst_stream_id, |
| 383 | + const std::pair<const char *, const char *> &node_names = {nullptr, nullptr}); | ||
| 383 | /// @ingroup GE | 384 | /// @ingroup GE |
| 384 | /// @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 |
| 385 | /// @return void | 386 | /// @return void |
| @@ -13,10 +13,6 @@ | |||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | 16 | ||
| 21 | 17 | ||
| 22 | namespace ge { | 18 | namespace ge { |
| @@ -118,13 +114,14 @@ void GetDiffStreamMaxLifeTime(const Node *const node, const int64_t stream_id, | |||
| 118 | max_life_time = kMaxLifeTime; | 114 | max_life_time = kMaxLifeTime; |
| 119 | auto node_op_desc = node->GetOpDescBarePtr(); | 115 | auto node_op_desc = node->GetOpDescBarePtr(); |
| 120 | GE_CHECK_NOTNULL_JUST_RETURN(node_op_desc); | 116 | GE_CHECK_NOTNULL_JUST_RETURN(node_op_desc); |
| 117 | + const auto node_stream_id = MemReuseUtils::GetStreamId(node_op_desc); | ||
| 121 | GELOGD("Out depend node:[%s] life begin:%" PRId64 " stream_id:[%" PRId64 "->%" PRId64 "]", node_op_desc->GetNamePtr(), | 118 | GELOGD("Out depend node:[%s] life begin:%" PRId64 " stream_id:[%" PRId64 "->%" PRId64 "]", node_op_desc->GetNamePtr(), |
| 122 | - node_op_desc->GetId(), MemReuseUtils::GetStreamId(node_op_desc), stream_id); | 119 | + node_op_desc->GetId(), node_stream_id, stream_id); |
| 123 | - if (MemReuseUtils::GetStreamId(node_op_desc) == stream_id) { | 120 | + if (node_stream_id == stream_id) { |
| 124 | max_life_time = node_op_desc->GetId(); | 121 | max_life_time = node_op_desc->GetId(); |
| 125 | return; | 122 | return; |
| 126 | } | 123 | } |
| 127 | - const auto it = diff_stream_edge_life.find(MemReuseUtils::GetStreamId(node_op_desc)); | 124 | + const auto it = diff_stream_edge_life.find(node_stream_id); |
| 128 | if (it == diff_stream_edge_life.cend()) { | 125 | if (it == diff_stream_edge_life.cend()) { |
| 129 | return; | 126 | return; |
| 130 | } | 127 | } |
| @@ -137,8 +134,8 @@ void GetDiffStreamMaxLifeTime(const Node *const node, const int64_t stream_id, | |||
| 137 | return; | 134 | return; |
| 138 | } | 135 | } |
| 139 | GELOGD("Node:[%s] life begin:%" PRId64 " stream_id:[%" PRId64 "->%" PRId64 "] life_time:[%" PRId64 "->%" PRId64 "]", | 136 | GELOGD("Node:[%s] life begin:%" PRId64 " stream_id:[%" PRId64 "->%" PRId64 "] life_time:[%" PRId64 "->%" PRId64 "]", |
| 140 | - node_op_desc->GetNamePtr(), node_op_desc->GetId(), MemReuseUtils::GetStreamId(node_op_desc), stream_id, | 137 | + node_op_desc->GetNamePtr(), node_op_desc->GetId(), node_stream_id, stream_id, (*edge_it).node_id, |
| 141 | - (*edge_it).node_id, (*edge_it).peer_node_id); | 138 | + (*edge_it).peer_node_id); |
| 142 | max_life_time = (*edge_it).peer_node_id; | 139 | max_life_time = (*edge_it).peer_node_id; |
| 143 | } | 140 | } |
| 144 | 141 | ||
| @@ -177,7 +177,7 @@ bool IsKnownSubgraphData(const Node *node) { | |||
| 177 | } | 177 | } |
| 178 | 178 | ||
| 179 | void SetReleaseBlockLifeEnd(MemoryBlock *to_release, int64_t stream_id) { | 179 | void SetReleaseBlockLifeEnd(MemoryBlock *to_release, int64_t stream_id) { |
| 180 | - const auto to_release_out_stream_life_time = to_release->NodeTypeIndexList().back().out_stream_life_time_; | 180 | + const auto &to_release_out_stream_life_time = to_release->NodeTypeIndexList().back().out_stream_life_time_; |
| 181 | if (to_release_out_stream_life_time.size() == 1) { | 181 | if (to_release_out_stream_life_time.size() == 1) { |
| 182 | for (const auto &item : to_release_out_stream_life_time) { | 182 | for (const auto &item : to_release_out_stream_life_time) { |
| 183 | size_t end_life_time = item.second.second; | 183 | size_t end_life_time = item.second.second; |
| @@ -17,13 +17,12 @@ | |||
| 17 | namespace ge { | 17 | namespace ge { |
| 18 | namespace { | 18 | namespace { |
| 19 | constexpr int64_t kAtomicCleanAllInput = -1; | 19 | constexpr int64_t kAtomicCleanAllInput = -1; |
| 20 | -std::vector<std::string> kNodeMemAttrStrs{"data", "concentrate_atomic"}; | 20 | +const std::vector<std::string> kNodeMemAttrStrs{"data", "concentrate_atomic"}; |
| 21 | -std::string GetNodeAttrStr(const NodeMemAttr &attr) { | 21 | +const char *GetNodeAttrStr(const NodeMemAttr &attr) { |
| 22 | if (static_cast<size_t>(attr) < kNodeMemAttrStrs.size()) { | 22 | if (static_cast<size_t>(attr) < kNodeMemAttrStrs.size()) { |
| 23 | - return kNodeMemAttrStrs.at(static_cast<size_t>(attr)); | 23 | + return kNodeMemAttrStrs[static_cast<size_t>(attr)].c_str(); |
| 24 | - } else { | ||
| 25 | - return "unknown"; | ||
| 26 | } | 24 | } |
| 25 | + return "unknown"; | ||
| 27 | } | 26 | } |
| 28 | 27 | ||
| 29 | bool IsNextNodeCleanInput(const Node *const node, const int32_t out_index) { | 28 | bool IsNextNodeCleanInput(const Node *const node, const int32_t out_index) { |
| @@ -226,7 +226,6 @@ bool ContinuousMemMng::IsTargetScenario(const Node *const node, ContinuousMemSce | |||
| 226 | OutDataAnchor *continuous_out_anchor = nullptr; | 226 | OutDataAnchor *continuous_out_anchor = nullptr; |
| 227 | if ((peer_anchor != nullptr) && | 227 | if ((peer_anchor != nullptr) && |
| 228 | (MemLayoutConflictUtil::IsContinuousOutputThroughRefNode(peer_anchor.get(), false, continuous_out_anchor))) { | 228 | (MemLayoutConflictUtil::IsContinuousOutputThroughRefNode(peer_anchor.get(), false, continuous_out_anchor))) { |
| 229 | - (void)continuous_out_anchor; | ||
| 230 | scenario = ContinuousMemScenario::kContinuousMemScenarioOutIn; | 229 | scenario = ContinuousMemScenario::kContinuousMemScenarioOutIn; |
| 231 | return true; | 230 | return true; |
| 232 | } | 231 | } |
| @@ -21,15 +21,8 @@ bool DynamicBatchBlockReuse(MemoryBlock &block) { | |||
| 21 | } | 21 | } |
| 22 | 22 | ||
| 23 | struct CompareSize { | 23 | struct CompareSize { |
| 24 | - explicit CompareSize() {} | ||
| 25 | - | ||
| 26 | bool operator()(const MemoryBlock *const left, const MemoryBlock *const right) const { | 24 | bool operator()(const MemoryBlock *const left, const MemoryBlock *const right) const { |
| 27 | - if ((left != nullptr) && (right != nullptr)) { | 25 | + return (left != nullptr) && (right != nullptr) && (left->Size() > right->Size()); |
| 28 | - auto left_size = left->Size(); | ||
| 29 | - auto right_size = right->Size(); | ||
| 30 | - return (left_size > right_size); | ||
| 31 | - } | ||
| 32 | - return false; | ||
| 33 | } | 26 | } |
| 34 | }; | 27 | }; |
| 35 | 28 | ||
| @@ -225,8 +218,8 @@ void DynamicBatchMemAssigner::DoResizeDynamicBatchBlocks( | |||
| 225 | } | 218 | } |
| 226 | 219 | ||
| 227 | void GetMaxBatchAllMemorySize(std::map<std::string, std::vector<int64_t>> &batch_all_memory_size, | 220 | void GetMaxBatchAllMemorySize(std::map<std::string, std::vector<int64_t>> &batch_all_memory_size, |
| 228 | - std::map<std::string, int64_t> batch_total_size, std::vector<int64_t> &all_memory_size, | 221 | + const std::map<std::string, int64_t> &batch_total_size, |
| 229 | - std::string &max_batch_label) { | 222 | + std::vector<int64_t> &all_memory_size, std::string &max_batch_label) { |
| 230 | // use max batch all memory size for reuse range | 223 | // use max batch all memory size for reuse range |
| 231 | int64_t max_batch_size = 0; | 224 | int64_t max_batch_size = 0; |
| 232 | for (const auto &it : batch_total_size) { | 225 | for (const auto &it : batch_total_size) { |
| @@ -51,8 +51,8 @@ class DynamicBatchMemAssigner { | |||
| 51 | }; | 51 | }; |
| 52 | 52 | ||
| 53 | void GetMaxBatchAllMemorySize(std::map<std::string, std::vector<int64_t>> &batch_all_memory_size, | 53 | void GetMaxBatchAllMemorySize(std::map<std::string, std::vector<int64_t>> &batch_all_memory_size, |
| 54 | - std::map<std::string, int64_t> batch_total_size, std::vector<int64_t> &all_memory_size, | 54 | + const std::map<std::string, int64_t> &batch_total_size, |
| 55 | - std::string &max_batch_label); | 55 | + std::vector<int64_t> &all_memory_size, std::string &max_batch_label); |
| 56 | 56 | ||
| 57 | } // namespace ge | 57 | } // namespace ge |
| 58 | 58 | ||
| @@ -2265,10 +2265,9 @@ void GraphMemoryAssigner::UpdateCurNodeInputDesc(const NodePtr &cur_node, int64_ | |||
| 2265 | 2265 | ||
| 2266 | void GraphMemoryAssigner::CheckNeedCalcDistAndUpdateVisitInfo( | 2266 | void GraphMemoryAssigner::CheckNeedCalcDistAndUpdateVisitInfo( |
| 2267 | const NodePtr &peer_out_node, const OutDataAnchorPtr &peer_out_anchor, size_t matched_mem_offset, | 2267 | const NodePtr &peer_out_node, const OutDataAnchorPtr &peer_out_anchor, size_t matched_mem_offset, |
| 2268 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, | 2268 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, |
| 2269 | bool &is_need_calc_distance) const { | 2269 | bool &is_need_calc_distance) const { |
| 2270 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>>::const_iterator iter = | 2270 | + const auto iter = mem_block_visit_info.find(matched_mem_offset); |
| 2271 | - mem_block_visit_info.find(matched_mem_offset); | ||
| 2272 | // cannot find visit info, peer_out_node must be a producer and this data is the first time to be visited. | 2271 | // cannot find visit info, peer_out_node must be a producer and this data is the first time to be visited. |
| 2273 | if (iter == mem_block_visit_info.end()) { | 2272 | if (iter == mem_block_visit_info.end()) { |
| 2274 | if (IsOutputVisitedByMultiStream(peer_out_node, peer_out_anchor->GetIdx())) { | 2273 | if (IsOutputVisitedByMultiStream(peer_out_node, peer_out_anchor->GetIdx())) { |
| @@ -2284,13 +2283,12 @@ void GraphMemoryAssigner::CheckNeedCalcDistAndUpdateVisitInfo( | |||
| 2284 | return; | 2283 | return; |
| 2285 | } | 2284 | } |
| 2286 | } else { | 2285 | } else { |
| 2287 | - if (mem_block_visit_info[matched_mem_offset].first == nullptr) { | 2286 | + if (iter->second.first == nullptr) { |
| 2288 | // multi-stream visit, no need to calculate | 2287 | // multi-stream visit, no need to calculate |
| 2289 | is_need_calc_distance = false; | 2288 | is_need_calc_distance = false; |
| 2290 | return; | 2289 | return; |
| 2291 | } | 2290 | } |
| 2292 | - if (peer_out_node->GetOpDesc()->GetStreamId() != | 2291 | + if (peer_out_node->GetOpDesc()->GetStreamId() != iter->second.first->GetOpDesc()->GetStreamId()) { |
| 2293 | - mem_block_visit_info[matched_mem_offset].first->GetOpDesc()->GetStreamId()) { | ||
| 2294 | // cur node and peer_out_node not in the same stream, no need to calculate | 2292 | // cur node and peer_out_node not in the same stream, no need to calculate |
| 2295 | is_need_calc_distance = false; | 2293 | is_need_calc_distance = false; |
| 2296 | return; | 2294 | return; |
| @@ -2302,12 +2300,14 @@ void GraphMemoryAssigner::CheckNeedCalcDistAndUpdateVisitInfo( | |||
| 2302 | 2300 | ||
| 2303 | // calculate distance, update visit info, update prev_node input desc, update cur node input desc | 2301 | // calculate distance, update visit info, update prev_node input desc, update cur node input desc |
| 2304 | void GraphMemoryAssigner::CalcDistanceAndUpdateDesc( | 2302 | void GraphMemoryAssigner::CalcDistanceAndUpdateDesc( |
| 2305 | - const std::map<std::string, int64_t> &node_index_in_stream, const InDataAnchorPtr &in_data_anchor, | 2303 | + const std::unordered_map<std::string, int64_t> &node_index_in_stream, const InDataAnchorPtr &in_data_anchor, |
| 2306 | size_t matched_mem_offset, const NodePtr &node, | 2304 | size_t matched_mem_offset, const NodePtr &node, |
| 2307 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, bool &is_need_skip) const { | 2305 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, |
| 2306 | + bool &is_need_skip) const { | ||
| 2308 | int64_t distance = -1; | 2307 | int64_t distance = -1; |
| 2309 | - auto prev_node = mem_block_visit_info[matched_mem_offset].first; | 2308 | + auto &visit_info = mem_block_visit_info[matched_mem_offset]; |
| 2310 | - auto prev_node_input_index_vec = mem_block_visit_info[matched_mem_offset].second; | 2309 | + auto prev_node = visit_info.first; |
| 2310 | + auto prev_node_input_index_vec = visit_info.second; | ||
| 2311 | GE_IF_BOOL_EXEC(prev_node == nullptr, is_need_skip = true; return); | 2311 | GE_IF_BOOL_EXEC(prev_node == nullptr, is_need_skip = true; return); |
| 2312 | if (prev_node_input_index_vec.size() == 1 && prev_node_input_index_vec[0] == -1) { | 2312 | if (prev_node_input_index_vec.size() == 1 && prev_node_input_index_vec[0] == -1) { |
| 2313 | // prev_node is producer and the data is just be produced(not visited by other node) | 2313 | // prev_node is producer and the data is just be produced(not visited by other node) |
| @@ -2322,9 +2322,9 @@ void GraphMemoryAssigner::CalcDistanceAndUpdateDesc( | |||
| 2322 | distance = node_index_in_stream.at(node->GetName()) - iter->second - 1; | 2322 | distance = node_index_in_stream.at(node->GetName()) - iter->second - 1; |
| 2323 | } | 2323 | } |
| 2324 | } | 2324 | } |
| 2325 | - mem_block_visit_info[matched_mem_offset].first = node; | 2325 | + visit_info.first = node; |
| 2326 | - mem_block_visit_info[matched_mem_offset].second.clear(); | 2326 | + visit_info.second.clear(); |
| 2327 | - mem_block_visit_info[matched_mem_offset].second.push_back(in_data_anchor->GetIdx()); | 2327 | + visit_info.second.push_back(in_data_anchor->GetIdx()); |
| 2328 | } else { // the data is visit by other customer just before. | 2328 | } else { // the data is visit by other customer just before. |
| 2329 | if (prev_node_input_index_vec.empty()) { | 2329 | if (prev_node_input_index_vec.empty()) { |
| 2330 | GELOGW("Missing prev node[%s] input index.", prev_node->GetName().c_str()); | 2330 | GELOGW("Missing prev node[%s] input index.", prev_node->GetName().c_str()); |
| @@ -2347,13 +2347,13 @@ void GraphMemoryAssigner::CalcDistanceAndUpdateDesc( | |||
| 2347 | } else { | 2347 | } else { |
| 2348 | distance = prev_next_distances[0]; // use the same prev_distance as previous anchor | 2348 | distance = prev_next_distances[0]; // use the same prev_distance as previous anchor |
| 2349 | } | 2349 | } |
| 2350 | - mem_block_visit_info[matched_mem_offset].second.push_back(in_data_anchor->GetIdx()); | 2350 | + visit_info.second.push_back(in_data_anchor->GetIdx()); |
| 2351 | } else { | 2351 | } else { |
| 2352 | distance = node_index_in_stream.at(node->GetName()) - node_index_in_stream.at(prev_node->GetName()) - 1; | 2352 | distance = node_index_in_stream.at(node->GetName()) - node_index_in_stream.at(prev_node->GetName()) - 1; |
| 2353 | UpdatePrevNodeInputDesc(prev_node, prev_node_input_index_vec, distance); | 2353 | UpdatePrevNodeInputDesc(prev_node, prev_node_input_index_vec, distance); |
| 2354 | - mem_block_visit_info[matched_mem_offset].first = node; | 2354 | + visit_info.first = node; |
| 2355 | - mem_block_visit_info[matched_mem_offset].second.clear(); | 2355 | + visit_info.second.clear(); |
| 2356 | - mem_block_visit_info[matched_mem_offset].second.push_back(in_data_anchor->GetIdx()); | 2356 | + visit_info.second.push_back(in_data_anchor->GetIdx()); |
| 2357 | } | 2357 | } |
| 2358 | } | 2358 | } |
| 2359 | UpdateCurNodeInputDesc(node, in_data_anchor->GetIdx(), distance); | 2359 | UpdateCurNodeInputDesc(node, in_data_anchor->GetIdx(), distance); |
| @@ -2361,7 +2361,7 @@ void GraphMemoryAssigner::CalcDistanceAndUpdateDesc( | |||
| 2361 | 2361 | ||
| 2362 | void GraphMemoryAssigner::DeleteVisitInfoWhenLifecycleEnded( | 2362 | void GraphMemoryAssigner::DeleteVisitInfoWhenLifecycleEnded( |
| 2363 | const NodePtr &node, const InDataAnchorPtr &in_data_anchor, size_t matched_mem_offset, | 2363 | const NodePtr &node, const InDataAnchorPtr &in_data_anchor, size_t matched_mem_offset, |
| 2364 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info) const { | 2364 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info) const { |
| 2365 | GE_IF_BOOL_EXEC(node->GetOpDesc() == nullptr, return); | 2365 | GE_IF_BOOL_EXEC(node->GetOpDesc() == nullptr, return); |
| 2366 | auto input_desc = node->GetOpDesc()->GetInputDesc(in_data_anchor->GetIdx()); | 2366 | auto input_desc = node->GetOpDesc()->GetInputDesc(in_data_anchor->GetIdx()); |
| 2367 | bool is_end_of_inputmem_lifecycle = false; | 2367 | bool is_end_of_inputmem_lifecycle = false; |
| @@ -2371,8 +2371,7 @@ void GraphMemoryAssigner::DeleteVisitInfoWhenLifecycleEnded( | |||
| 2371 | is_end_of_inputmem_lifecycle) { | 2371 | is_end_of_inputmem_lifecycle) { |
| 2372 | GELOGD("ATTR_NAME_IS_END_OF_INPUTMEM_LIFECYCLE is true, node name is [%s], in_data_anchor index is [%d]", | 2372 | GELOGD("ATTR_NAME_IS_END_OF_INPUTMEM_LIFECYCLE is true, node name is [%s], in_data_anchor index is [%d]", |
| 2373 | node->GetName().c_str(), in_data_anchor->GetIdx()); | 2373 | node->GetName().c_str(), in_data_anchor->GetIdx()); |
| 2374 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>>::const_iterator iter = | 2374 | + const auto iter = mem_block_visit_info.find(matched_mem_offset); |
| 2375 | - mem_block_visit_info.find(matched_mem_offset); | ||
| 2376 | if (iter != mem_block_visit_info.cend()) { | 2375 | if (iter != mem_block_visit_info.cend()) { |
| 2377 | mem_block_visit_info.erase(iter); | 2376 | mem_block_visit_info.erase(iter); |
| 2378 | } | 2377 | } |
| @@ -2380,8 +2379,8 @@ void GraphMemoryAssigner::DeleteVisitInfoWhenLifecycleEnded( | |||
| 2380 | } | 2379 | } |
| 2381 | 2380 | ||
| 2382 | void GraphMemoryAssigner::MarkNodeDistanceAttr( | 2381 | void GraphMemoryAssigner::MarkNodeDistanceAttr( |
| 2383 | - const NodePtr &node, std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, | 2382 | + const NodePtr &node, std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, |
| 2384 | - const std::map<std::string, int64_t> &node_index_in_stream) { | 2383 | + const std::unordered_map<std::string, int64_t> &node_index_in_stream) { |
| 2385 | GELOGD("Begin to mark node distance attr, node name is [%s]", node->GetName().c_str()); | 2384 | GELOGD("Begin to mark node distance attr, node name is [%s]", node->GetName().c_str()); |
| 2386 | for (const auto &in_data_anchor : node->GetAllInDataAnchors()) { | 2385 | for (const auto &in_data_anchor : node->GetAllInDataAnchors()) { |
| 2387 | auto peer_out_anchor = in_data_anchor->GetPeerOutAnchor(); | 2386 | auto peer_out_anchor = in_data_anchor->GetPeerOutAnchor(); |
| @@ -2412,11 +2411,11 @@ void GraphMemoryAssigner::MarkNodeDistanceAttr( | |||
| 2412 | 2411 | ||
| 2413 | void GraphMemoryAssigner::MarkDistanceAttr() { | 2412 | void GraphMemoryAssigner::MarkDistanceAttr() { |
| 2414 | // key: mem_offset of the memory which we visited. value: node we visited and input index of this node | 2413 | // key: mem_offset of the memory which we visited. value: node we visited and input index of this node |
| 2415 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> mem_block_visit_info; | 2414 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> mem_block_visit_info; |
| 2416 | // key: node name, value: topo order of node in it's belonged stream(exclude ge_local_op) | 2415 | // key: node name, value: topo order of node in it's belonged stream(exclude ge_local_op) |
| 2417 | - std::map<std::string, int64_t> node_index_in_stream; | 2416 | + std::unordered_map<std::string, int64_t> node_index_in_stream; |
| 2418 | // key: stream id, value: cur nodes num in that stream | 2417 | // key: stream id, value: cur nodes num in that stream |
| 2419 | - std::map<int64_t, int64_t> stream_nodes_num; | 2418 | + std::unordered_map<int64_t, int64_t> stream_nodes_num; |
| 2420 | 2419 | ||
| 2421 | for (auto &node : compute_graph_->GetAllNodes()) { | 2420 | for (auto &node : compute_graph_->GetAllNodes()) { |
| 2422 | auto node_op_desc = node->GetOpDesc(); | 2421 | auto node_op_desc = node->GetOpDesc(); |
| @@ -2424,13 +2423,10 @@ void GraphMemoryAssigner::MarkDistanceAttr() { | |||
| 2424 | // Only sinking computing nodes need to be calculated, excluding the nodes which don't have task | 2423 | // Only sinking computing nodes need to be calculated, excluding the nodes which don't have task |
| 2425 | if ((node_op_desc->GetOpKernelLibName() != kEngineNameGeLocal) && (!node_op_desc->HasAttr(ATTR_NAME_NOTASK))) { | 2424 | if ((node_op_desc->GetOpKernelLibName() != kEngineNameGeLocal) && (!node_op_desc->HasAttr(ATTR_NAME_NOTASK))) { |
| 2426 | int64_t stream_id = node_op_desc->GetStreamId(); | 2425 | int64_t stream_id = node_op_desc->GetStreamId(); |
| 2427 | - if (stream_nodes_num.find(stream_id) == stream_nodes_num.end()) { | 2426 | + auto stream_count = stream_nodes_num.emplace(stream_id, 0).first; |
| 2428 | - stream_nodes_num.insert(std::make_pair(stream_id, 1)); | 2427 | + ++stream_count->second; |
| 2429 | - } else { | 2428 | + node_index_in_stream.emplace(node->GetName(), stream_count->second - 1); |
| 2430 | - ++stream_nodes_num[stream_id]; | 2429 | + (void)AttrUtils::SetInt(node_op_desc, ATTR_NAME_OP_READ_WRITE_INDEX, stream_count->second - 1); |
| 2431 | - } | ||
| 2432 | - node_index_in_stream.insert(std::make_pair(node->GetName(), stream_nodes_num[stream_id] - 1)); | ||
| 2433 | - (void)AttrUtils::SetInt(node->GetOpDesc(), ATTR_NAME_OP_READ_WRITE_INDEX, stream_nodes_num[stream_id] - 1); | ||
| 2434 | 2430 | ||
| 2435 | MarkNodeDistanceAttr(node, mem_block_visit_info, node_index_in_stream); | 2431 | MarkNodeDistanceAttr(node, mem_block_visit_info, node_index_in_stream); |
| 2436 | } else { | 2432 | } else { |
| @@ -173,21 +173,22 @@ class GraphMemoryAssigner { | |||
| 173 | 173 | ||
| 174 | void CheckNeedCalcDistAndUpdateVisitInfo( | 174 | void CheckNeedCalcDistAndUpdateVisitInfo( |
| 175 | const NodePtr &peer_out_node, const OutDataAnchorPtr &peer_out_anchor, size_t matched_mem_offset, | 175 | const NodePtr &peer_out_node, const OutDataAnchorPtr &peer_out_anchor, size_t matched_mem_offset, |
| 176 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, | 176 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, |
| 177 | bool &is_need_calc_distance) const; | 177 | bool &is_need_calc_distance) const; |
| 178 | 178 | ||
| 179 | - void CalcDistanceAndUpdateDesc(const std::map<std::string, int64_t> &node_index_in_stream, | 179 | + void CalcDistanceAndUpdateDesc( |
| 180 | - const InDataAnchorPtr &in_data_anchor, size_t matched_mem_offset, const NodePtr &node, | 180 | + const std::unordered_map<std::string, int64_t> &node_index_in_stream, const InDataAnchorPtr &in_data_anchor, |
| 181 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, | 181 | + size_t matched_mem_offset, const NodePtr &node, |
| 182 | - bool &is_need_skip) const; | 182 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, |
| 183 | + bool &is_need_skip) const; | ||
| 183 | 184 | ||
| 184 | void DeleteVisitInfoWhenLifecycleEnded( | 185 | void DeleteVisitInfoWhenLifecycleEnded( |
| 185 | const NodePtr &node, const InDataAnchorPtr &in_data_anchor, size_t matched_mem_offset, | 186 | const NodePtr &node, const InDataAnchorPtr &in_data_anchor, size_t matched_mem_offset, |
| 186 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info) const; | 187 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info) const; |
| 187 | 188 | ||
| 188 | void MarkNodeDistanceAttr(const NodePtr &node, | 189 | void MarkNodeDistanceAttr(const NodePtr &node, |
| 189 | - std::map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, | 190 | + std::unordered_map<size_t, std::pair<NodePtr, std::vector<int64_t>>> &mem_block_visit_info, |
| 190 | - const std::map<std::string, int64_t> &node_index_in_stream); | 191 | + const std::unordered_map<std::string, int64_t> &node_index_in_stream); |
| 191 | 192 | ||
| 192 | MemoryOffsetMap memory_offset_; | 193 | MemoryOffsetMap memory_offset_; |
| 193 | ComputeGraphPtr compute_graph_; | 194 | ComputeGraphPtr compute_graph_; |
| @@ -44,7 +44,7 @@ Status HybridMemAssigner::AssignMemory(BlockMemAssigner *block_assigner, size_t | |||
| 44 | GE_ASSERT_SUCCESS(block_assigner->AssignMemoryWithReuse(ranges)); | 44 | GE_ASSERT_SUCCESS(block_assigner->AssignMemoryWithReuse(ranges)); |
| 45 | 45 | ||
| 46 | // total size | 46 | // total size |
| 47 | - for (auto it : block_assigner->GetMemOffsets()) { | 47 | + for (const auto &it : block_assigner->GetMemOffsets()) { |
| 48 | mem_size += it.second; | 48 | mem_size += it.second; |
| 49 | } | 49 | } |
| 50 | return SUCCESS; | 50 | return SUCCESS; |
| @@ -137,7 +137,7 @@ Status HybridMemAssigner::Assign() { | |||
| 137 | for (const auto &memory_assigner : memory_assigners) { | 137 | for (const auto &memory_assigner : memory_assigners) { |
| 138 | GELOGI("%s memory assigner memory size:%zu", memory_assigner.first.c_str(), memory_assigner.second.second); | 138 | GELOGI("%s memory assigner memory size:%zu", memory_assigner.first.c_str(), memory_assigner.second.second); |
| 139 | } | 139 | } |
| 140 | - if ((!vector_future.empty()) && (!memory_assigners.empty())) { | 140 | + if (!vector_future.empty()) { |
| 141 | memory_assigners[0].second.first->SetOpMemOffset(false); | 141 | memory_assigners[0].second.first->SetOpMemOffset(false); |
| 142 | mem_offsets_ = memory_assigners[0].second.first->GetMemOffsets(); | 142 | mem_offsets_ = memory_assigners[0].second.first->GetMemOffsets(); |
| 143 | memory_stat_ = memory_assigners[0].second.first->GetMemoryStat(); | 143 | memory_stat_ = memory_assigners[0].second.first->GetMemoryStat(); |
| @@ -66,8 +66,9 @@ Status GetReadOnlySymbol(const MemAssistInfo &mem_assist_info, std::set<std::str | |||
| 66 | } | 66 | } |
| 67 | for (const auto &out_anchor : node->GetAllOutDataAnchors()) { | 67 | for (const auto &out_anchor : node->GetAllOutDataAnchors()) { |
| 68 | const NodeIndexIO output_info(node, out_anchor->GetIdx(), kOut); | 68 | const NodeIndexIO output_info(node, out_anchor->GetIdx(), kOut); |
| 69 | - const auto output_symbol = mem_assist_info.anchor_to_symbol.find(output_info.ToString())->second; | 69 | + const auto symbol_it = mem_assist_info.anchor_to_symbol.find(output_info.ToString()); |
| 70 | - read_only_symbols.insert(output_symbol); | 70 | + GE_ASSERT_TRUE(symbol_it != mem_assist_info.anchor_to_symbol.end()); |
| 71 | + read_only_symbols.insert(symbol_it->second); | ||
| 71 | } | 72 | } |
| 72 | } | 73 | } |
| 73 | return SUCCESS; | 74 | return SUCCESS; |
| @@ -175,8 +176,12 @@ Status RemoveSymbolConflicts(const MemAssistInfo &mem_assist_info, const NodePtr | |||
| 175 | GE_ASSERT_TRUE(ge::IntegerChecker<uint32_t>::Compat(input_index)); | 176 | GE_ASSERT_TRUE(ge::IntegerChecker<uint32_t>::Compat(input_index)); |
| 176 | const NodeIndexIO cur_node_input_info(node, static_cast<uint32_t>(input_index), kIn); | 177 | const NodeIndexIO cur_node_input_info(node, static_cast<uint32_t>(input_index), kIn); |
| 177 | const NodeIndexIO cur_node_output_info(node, static_cast<uint32_t>(output_index), kOut); | 178 | const NodeIndexIO cur_node_output_info(node, static_cast<uint32_t>(output_index), kOut); |
| 178 | - const auto &input_symbol = mem_assist_info.anchor_to_symbol.find(cur_node_input_info.ToString())->second; | 179 | + const auto input_it = mem_assist_info.anchor_to_symbol.find(cur_node_input_info.ToString()); |
| 179 | - const auto &output_symbol = mem_assist_info.anchor_to_symbol.find(cur_node_output_info.ToString())->second; | 180 | + GE_ASSERT_TRUE(input_it != mem_assist_info.anchor_to_symbol.end()); |
| 181 | + const auto output_it = mem_assist_info.anchor_to_symbol.find(cur_node_output_info.ToString()); | ||
| 182 | + GE_ASSERT_TRUE(output_it != mem_assist_info.anchor_to_symbol.end()); | ||
| 183 | + const auto &input_symbol = input_it->second; | ||
| 184 | + const auto &output_symbol = output_it->second; | ||
| 180 | if (input_symbol == output_symbol) { // 输入符号和输出符号相同,不需要合并和设置复用关系 | 185 | if (input_symbol == output_symbol) { // 输入符号和输出符号相同,不需要合并和设置复用关系 |
| 181 | GELOGD("Node %s input symbol[%s] is equal to output symbol[%s], skip inplace.", node->GetName().c_str(), | 186 | GELOGD("Node %s input symbol[%s] is equal to output symbol[%s], skip inplace.", node->GetName().c_str(), |
| 182 | input_symbol.c_str(), output_symbol.c_str()); | 187 | input_symbol.c_str(), output_symbol.c_str()); |
| @@ -376,13 +376,15 @@ bool MemReuseUtils::IsContinuousOutput(const ge::NodePtr &n) { | |||
| 376 | } | 376 | } |
| 377 | 377 | ||
| 378 | bool MemReuseUtils::IsNoReleaseNodeOutBlock(const ge::Node *const node) { | 378 | bool MemReuseUtils::IsNoReleaseNodeOutBlock(const ge::Node *const node) { |
| 379 | - for (const auto &input_desc : node->GetOpDesc()->GetAllInputsDescPtr()) { | 379 | + const auto op_desc = node->GetOpDescBarePtr(); |
| 380 | + GE_ASSERT_NOTNULL(op_desc); | ||
| 381 | + for (const auto &input_desc : op_desc->GetAllInputsDescPtr()) { | ||
| 380 | if ((input_desc != nullptr) && | 382 | if ((input_desc != nullptr) && |
| 381 | (kNotPostReuseDataType.find(input_desc->GetDataType()) != kNotPostReuseDataType.cend())) { | 383 | (kNotPostReuseDataType.find(input_desc->GetDataType()) != kNotPostReuseDataType.cend())) { |
| 382 | return true; | 384 | return true; |
| 383 | } | 385 | } |
| 384 | } | 386 | } |
| 385 | - for (const auto &output_desc : node->GetOpDesc()->GetAllOutputsDescPtr()) { | 387 | + for (const auto &output_desc : op_desc->GetAllOutputsDescPtr()) { |
| 386 | if ((output_desc != nullptr) && | 388 | if ((output_desc != nullptr) && |
| 387 | (kNotPostReuseDataType.find(output_desc->GetDataType()) != kNotPostReuseDataType.cend())) { | 389 | (kNotPostReuseDataType.find(output_desc->GetDataType()) != kNotPostReuseDataType.cend())) { |
| 388 | return true; | 390 | return true; |
| @@ -454,8 +456,8 @@ bool MemReuseUtils::IsAtomicWorkSpace(const int64_t index, | |||
| 454 | if (it.second.empty()) { | 456 | if (it.second.empty()) { |
| 455 | continue; | 457 | continue; |
| 456 | } | 458 | } |
| 457 | - for (const auto &workspae_info : it.second) { | 459 | + for (const auto &workspace_info : it.second) { |
| 458 | - if (workspae_info.first == index) { | 460 | + if (workspace_info.first == index) { |
| 459 | GELOGD("Node:%s's workspace:%" PRId64 " is atomic.", it.first.c_str(), index); | 461 | GELOGD("Node:%s's workspace:%" PRId64 " is atomic.", it.first.c_str(), index); |
| 460 | return true; | 462 | return true; |
| 461 | } | 463 | } |
| @@ -13,13 +13,8 @@ | |||
| 13 | 13 | ||
| 14 | 14 | ||
| 15 | 15 | ||
| 16 | - | ||
| 17 | 16 | ||
| 18 | - | ||
| 19 | - | ||
| 20 | 17 | ||
| 21 | - | ||
| 22 | - | ||
| 23 | 18 | ||
| 24 | 19 | ||
| 25 | 20 | ||
| @@ -871,13 +866,7 @@ std::string MemoryBlock::String() const { | |||
| 871 | 866 | ||
| 872 | // ascending order | 867 | // ascending order |
| 873 | bool CompareBlockIndex(const MemoryBlock *const left, const MemoryBlock *const right) { | 868 | bool CompareBlockIndex(const MemoryBlock *const left, const MemoryBlock *const right) { |
| 874 | - if (left == nullptr || right == nullptr) { | 869 | + return (left != nullptr) && (right != nullptr) && (left->input_index_ < right->input_index_); |
| 875 | - return false; | ||
| 876 | - } | ||
| 877 | - if (left->input_index_ < right->input_index_) { | ||
| 878 | - return true; | ||
| 879 | - } | ||
| 880 | - return false; | ||
| 881 | } | 870 | } |
| 882 | 871 | ||
| 883 | bool CompareLifeInterval::operator()(MemoryBlock *const left, MemoryBlock *const right) const { | 872 | bool CompareLifeInterval::operator()(MemoryBlock *const left, MemoryBlock *const right) const { |
| @@ -170,10 +170,8 @@ struct NodeTypeIndex { | |||
| 170 | auto life_begin = GetLifeBegin(); | 170 | auto life_begin = GetLifeBegin(); |
| 171 | if (life_begin != node_id_) { | 171 | if (life_begin != node_id_) { |
| 172 | return std::to_string(life_begin) + "--" + std::to_string(node_id_); | 172 | return std::to_string(life_begin) + "--" + std::to_string(node_id_); |
| 173 | - } else { | ||
| 174 | - return std::to_string(life_begin); | ||
| 175 | } | 173 | } |
| 176 | - return ""; | 174 | + return std::to_string(life_begin); |
| 177 | } | 175 | } |
| 178 | 176 | ||
| 179 | std::vector<size_t> GetLifeEnd() const { | 177 | std::vector<size_t> GetLifeEnd() const { |