Try not to modify base name (#3638)

This commit is contained in:
Wei-Sheng Chin 2020-04-22 22:24:43 -07:00 committed by GitHub
parent ffe19ae49b
commit d9641f292d
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 47 additions and 1 deletions

View file

@ -1110,6 +1110,14 @@ class Graph {
// Graph value_info.
std::vector<const NodeArg*> value_info_;
// Strings which have been used as node names.
// New node name should not conflict with this set.
std::unordered_set<std::string> 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<std::string> generated_node_arg_names_;
// All node args owned by <*this> graph. Key is node arg name.
std::unordered_map<std::string, std::unique_ptr<NodeArg>> node_args_;

View file

@ -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<Node>& 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<Node>& 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;
}