From d9641f292d474aaf8f428cce4547c7915ce60f5c Mon Sep 17 00:00:00 2001 From: Wei-Sheng Chin Date: Wed, 22 Apr 2020 22:24:43 -0700 Subject: [PATCH] Try not to modify base name (#3638) --- include/onnxruntime/core/graph/graph.h | 8 ++++++ onnxruntime/core/graph/graph.cc | 40 +++++++++++++++++++++++++- 2 files changed, 47 insertions(+), 1 deletion(-) diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index 766c6d96b7..5bf8ca9275 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -1110,6 +1110,14 @@ class Graph { // Graph value_info. std::vector value_info_; + // Strings which have been used as node names. + // New node name should not conflict with this set. + std::unordered_set generated_node_names_; + + // Strings which have been used as node_arg names. + // New node_arg name should not conflict this this set. + std::unordered_set generated_node_arg_names_; + // All node args owned by <*this> graph. Key is node arg name. std::unordered_map> node_args_; diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index 16e2a504a2..40dfa89a7e 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -2376,16 +2376,47 @@ Node& Graph::AddNode(const NodeProto& node_proto, } std::string Graph::GenerateNodeArgName(const std::string& base_name) { + // Check if base_name has been used in as any of node_args_' names. + bool found = node_args_.find(base_name) != node_args_.end(); + // Check if base_name has been generated by this function. + // If not, add base_name into name set and return the base_name + // as the generated name. + found |= generated_node_arg_names_.find(base_name) != generated_node_arg_names_.end(); + if (!found) { + generated_node_arg_names_.insert(base_name); + return base_name; + } + + // base_name has been used by another node. Because two node_arg's cannot have + // the sam name, we are going to generate another string. std::string new_name; do { std::ostringstream str; str << base_name << "_" << name_generator_++; new_name = str.str(); - } while (node_args_.find(new_name) != node_args_.end()); + // If node_args_ or generated_node_arg_names_ contains new_name, we go to the next iteration. + } while (node_args_.find(new_name) != node_args_.end() || + generated_node_arg_names_.find(new_name) != generated_node_arg_names_.end()); + + // Now new_name is different than any of existing node_arg names. + // Register new_name so that it won't be used again. + generated_node_arg_names_.insert(new_name); + return new_name; } std::string Graph::GenerateNodeName(const std::string& base_name) { + // Check if base_name has been used in as any of nodes_' names. + bool found = std::find_if(nodes_.cbegin(), nodes_.cend(), [&base_name](const std::unique_ptr& n) { + return (n != nullptr) && (n->Name() == base_name);}) != nodes_.end(); + // Check if base_name has been generated by this function. + found |= generated_node_names_.find(base_name) != generated_node_names_.end(); + if (!found) { + // Register base_name so that it won't be used again. + generated_node_names_.insert(base_name); + return base_name; + } + std::string new_name; bool keep_going = true; @@ -2394,11 +2425,18 @@ std::string Graph::GenerateNodeName(const std::string& base_name) { str << base_name << "_" << name_generator_++; new_name = str.str(); + // Check if new_name has been used in as any of nodes_' names. keep_going = std::find_if(nodes_.cbegin(), nodes_.cend(), [&new_name](const std::unique_ptr& n) { return (n != nullptr) && (n->Name() == new_name); }) != nodes_.end(); + // Check if new_name has been generated by this function. + keep_going |= generated_node_names_.find(new_name) != generated_node_names_.end(); } while (keep_going); + // Now new_name is different than any of existing node names. + // Register new_name so that it won't be used again. + generated_node_names_.insert(new_name); + return new_name; }