From 8d830adf249024d60d58ced349831c507781b007 Mon Sep 17 00:00:00 2001 From: Scott McKay Date: Fri, 5 Aug 2022 13:20:25 +1000 Subject: [PATCH] Rework parts of Graph::Resolve to reduce memory usage (#12176) * Rework some aspects of Graph::Resolve to reduce memory usage. --- include/onnxruntime/core/graph/graph.h | 47 +++++---- onnxruntime/core/graph/graph.cc | 131 +++++++++---------------- 2 files changed, 78 insertions(+), 100 deletions(-) diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index 73069ce7c3..5359edbbb6 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -690,8 +690,8 @@ class Graph { #if !defined(DISABLE_EXTERNAL_INITIALIZERS) /** This function takes externally provided data for initializers with external data - * and replaces graph initializers with its content. - */ + * and replaces graph initializers with its content. + */ common::Status InjectExternalInitializedTensors(const InlinedHashMap& external_initializers); #endif // !defined(DISABLE_EXTERNAL_INITIALIZERS) @@ -1273,7 +1273,7 @@ class Graph { return Resolve(default_options); } - const std::unordered_set& GetOuterScopeNodeArgNames() const noexcept{ + const std::unordered_set& GetOuterScopeNodeArgNames() const noexcept { return outer_scope_node_arg_names_; } @@ -1305,11 +1305,11 @@ class Graph { Graph(Graph& parent_graph, const Node& parent_node, ONNX_NAMESPACE::GraphProto& subgraph_proto); Graph(const Model& owning_model, - IOnnxRuntimeOpSchemaCollectionPtr schema_registry, - ONNX_NAMESPACE::GraphProto& subgraph_proto, - const std::unordered_map& domain_version_map, - const logging::Logger& logger, - bool strict_shape_type_inference); + IOnnxRuntimeOpSchemaCollectionPtr schema_registry, + ONNX_NAMESPACE::GraphProto& subgraph_proto, + const std::unordered_map& domain_version_map, + const logging::Logger& logger, + bool strict_shape_type_inference); #endif virtual ~Graph(); @@ -1423,23 +1423,33 @@ class Graph { // or a subgraph, so that the various operations that are part of the Resolve can work iteratively or // recursively as needed. struct ResolveContext { - ResolveContext() = default; + ResolveContext(const Graph& owning_graph) : graph{owning_graph} { + } - std::unordered_map> output_args; - std::unordered_set inputs_and_initializers; - std::unordered_set outer_scope_node_args; - std::unordered_map node_name_to_index; + std::unordered_map> output_args; + std::unordered_set inputs_and_initializers; + std::unordered_map node_name_to_index; std::unordered_set nodes_with_subgraphs; + // check if the provided name is an input/initialize/node output of this Graph instance during Graph::Resolve. + // Graph::node_args_ can have stale entries so we can't rely on that. + bool IsLocalValue(const std::string& name) const; + + // check if an ancestor graph has a valid value with the provided name during Graph::Resolve. + // Once Graph::Resolve completes Graph::IsOuterScopeValue can be used and is more efficient. + bool IsOuterScopeValue(const std::string& name) const; + void Clear() { output_args.clear(); inputs_and_initializers.clear(); - outer_scope_node_args.clear(); node_name_to_index.clear(); nodes_with_subgraphs.clear(); } private: + bool IsInputInitializerOrOutput(const std::string& name, bool check_ancestors) const; + + const Graph& graph; ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(ResolveContext); }; @@ -1486,6 +1496,7 @@ class Graph { // Infer and set type information across <*this> graph if needed, and verify type/attribute // information matches between node and op. + common::Status VerifyNodeAndOpMatch(const ResolveOptions& options); // Set graph inputs/outputs when resolving a graph.. @@ -1585,9 +1596,9 @@ class Graph { #if !defined(ORT_MINIMAL_BUILD) IOnnxRuntimeOpSchemaCollectionPtr schema_registry_; - //Currently to make the ORT in-memory graph work, we have to create a temporary op schema - //for the fused kernel. I really don't like it. but for short-term solution, let's host - //those schemas here. + // Currently to make the ORT in-memory graph work, we have to create a temporary op schema + // for the fused kernel. I really don't like it. but for short-term solution, let's host + // those schemas here. InlinedVector> fused_schemas_containers_; #endif // !defined(ORT_MINIMAL_BUILD) @@ -1656,7 +1667,7 @@ class Graph { // Model IR version. Version ir_version_{ONNX_NAMESPACE::Version::IR_VERSION}; - ResolveContext resolve_context_; + ResolveContext resolve_context_{*this}; // the parent graph if this is a subgraph. Graph* parent_graph_; diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index 5def27f160..b847e0264d 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -1411,39 +1411,6 @@ Status Graph::VerifyNoDuplicateName() { return Status::OK(); } -// Recurse into any subgraphs to update the list of NodeArg values in outer scope. -// This information is needed to resolve any dependencies on outer scope values. -common::Status Graph::SetOuterScopeNodeArgs(const std::unordered_set& outer_scope_node_args) { - resolve_context_.outer_scope_node_args = outer_scope_node_args; - - if (!resolve_context_.nodes_with_subgraphs.empty()) { - // Build the list of NodeArg's that are valid for a subgraph of this GraphBase instance: - // - outer scope for this graph - // - any inputs/initializers from this graph - // - any outputs from nodes in this graph - // - // We provide outputs from all nodes in this graph at this stage. - // BuildConnections will link the node with the subgraph to any outer scope Node/NodeArgs it consumes. - // PerformTopologicalSortAndCheckIsAcyclic will validate these links. - std::unordered_set node_args_in_scope_for_subgraph = outer_scope_node_args; - - node_args_in_scope_for_subgraph.insert(resolve_context_.inputs_and_initializers.cbegin(), - resolve_context_.inputs_and_initializers.cend()); - - std::transform(resolve_context_.output_args.cbegin(), resolve_context_.output_args.cend(), - std::inserter(node_args_in_scope_for_subgraph, node_args_in_scope_for_subgraph.end()), - [](const std::pair>& entry) { return entry.first; }); - - for (auto* node : resolve_context_.nodes_with_subgraphs) { - for (auto& subgraph : node->MutableSubgraphs()) { - auto status = subgraph->SetOuterScopeNodeArgs(node_args_in_scope_for_subgraph); - ORT_RETURN_IF_ERROR(status); - } - } - } - - return Status::OK(); -} #endif // !defined(ORT_MINIMAL_BUILD) #if !defined(ORT_MINIMAL_BUILD) || defined(ORT_EXTENDED_MINIMAL_BUILD) @@ -1536,10 +1503,6 @@ void Graph::RemoveEdge(NodeIndex src_node_index, NodeIndex dst_node_index, int s #if !defined(ORT_MINIMAL_BUILD) GSL_SUPPRESS(es .84) // ignoring return value from unordered_map::insert causes noisy complaint Status Graph::BuildConnections(std::unordered_set& outer_scope_node_args_consumed) { - const std::unordered_set& outer_scope_node_args = resolve_context_.outer_scope_node_args; - - std::unordered_set node_args_consumed_by_subgraphs; - // recurse into subgraphs first so we can update any nodes in this graph that are used by those subgraphs if (!resolve_context_.nodes_with_subgraphs.empty()) { for (auto* node : resolve_context_.nodes_with_subgraphs) { @@ -1574,11 +1537,12 @@ Status Graph::BuildConnections(std::unordered_set& outer_scope_node " Graph may not conform to the ONNX spec and contain initializers that are not graph inputs."); } } else { - // this value may be produced by this graph, or it could still be coming from a parent graph if it - // is also directly consumed at this level as we create a NodeArg for all Node inputs in this graph. - // due to that we need to check the outputs from this level to determine if it is an outer scope value. - // we don't have that info yet so store and check before returning from BuildConnections - ORT_IGNORE_RETURN_VALUE(node_args_consumed_by_subgraphs.insert(node_arg_name)); + // as we create a NodeArg instance for all Node inputs the value could be produced by this graph, + // or could be coming from outer scope. check the valid values for just this Graph using resolve_context_. + // if none are found, it's an outer scope value. + if (resolve_context_.IsLocalValue(node_arg_name) == false) { + ORT_IGNORE_RETURN_VALUE(outer_scope_node_args_consumed.insert(node_arg_name)); + } } // add it to the Node's list of implicit inputs @@ -1644,7 +1608,7 @@ Status Graph::BuildConnections(std::unordered_set& outer_scope_node // If it is present in the outer scope it will be 'fed' by the execution frame // providing access to the OrtValue from the outer scope. Pass the name back up so nodes can // be linked correctly at that level. - if (outer_scope_node_args.find(input_arg_name) != outer_scope_node_args.cend()) { + if (resolve_context_.IsOuterScopeValue(input_arg_name)) { ORT_IGNORE_RETURN_VALUE(outer_scope_node_args_consumed.insert(input_arg_name)); } } @@ -1679,15 +1643,6 @@ Status Graph::BuildConnections(std::unordered_set& outer_scope_node ORT_RETURN_IF_ERROR(PopulateNodeArgToProducerConsumerLookupsFromNodes()); - // finally check any node args consumed by subgraphs to see if they're available locally. - // if not we add them to the list of outer scope values consumed. - for (const auto& name : node_args_consumed_by_subgraphs) { - if (node_arg_to_producer_node_.find(name) == node_arg_to_producer_node_.cend() && - resolve_context_.inputs_and_initializers.find(name) == resolve_context_.inputs_and_initializers.cend()) { - ORT_IGNORE_RETURN_VALUE(outer_scope_node_args_consumed.insert(name)); - } - } - return Status::OK(); } @@ -2483,21 +2438,28 @@ Status Graph::VerifyNodeAndOpMatch(const ResolveOptions& options) { // Set the parent directory of model path to load external tensors if exist ctx.set_model_dir(ToUTF8String(ModelPath().ParentPath().ToPathString())); - LexicalScopeContext lsc; - lsc.output_names.insert(resolve_context_.inputs_and_initializers.cbegin(), - resolve_context_.inputs_and_initializers.cend()); + LexicalScopeContext parent; + if (parent_node_) { + // add outer scope values. these are set as implicit inputs to the node containing the subgraph + // in BuildConnections, which happens prior to this being called during Graph::Resolve + const auto& outer_scope_values = parent_node_->ImplicitInputDefs(); + parent.output_names.reserve(outer_scope_values.size()); - // technically we could add values from Node.GetDefinitions().implicit_input_defs on a per-node basis inside - // the below loop so that we only check against the specific outer dependencies of the node. - // doing that requires lots of copies of LexicalScopeContext.output_names to clear out the per-Node values - // after each loop. instead add all the outer scope values upfront so we can just accumulate new inner scope values - // during each loop iteration. - lsc.output_names.insert(resolve_context_.outer_scope_node_args.cbegin(), - resolve_context_.outer_scope_node_args.cend()); + for (const auto* implicit_inputs : outer_scope_values) { + parent.output_names.insert(implicit_inputs->Name()); + } + } else { + // we may have some locally defined outer scope args if we're in the middle of constructing a subgraph + // and need to call Resolve. parent_node_ would be null in this case + parent.output_names.insert(outer_scope_node_arg_names_.cbegin(), outer_scope_node_arg_names_.cend()); + } - // we may have some locally defined outer scope args if we're in the middle of constructing a subgraph - // and need to call Resolve - lsc.output_names.insert(outer_scope_node_arg_names_.cbegin(), outer_scope_node_arg_names_.cend()); + LexicalScopeContext lsc{parent}; + lsc.output_names.reserve(resolve_context_.inputs_and_initializers.size() + resolve_context_.output_args.size()); + + for (const std::string_view& input : resolve_context_.inputs_and_initializers) { + lsc.output_names.insert(std::string(input)); + } for (auto node_index : nodes_in_topological_order_) { // Node verification. @@ -2588,22 +2550,16 @@ Status Graph::VerifyNodeAndOpMatch(const ResolveOptions& options) { } } - // verify subgraphs - for (auto node_index : nodes_in_topological_order_) { - auto& node = *GetNode(node_index); - for (auto& entry : node.GetAttributeNameToMutableSubgraphMap()) { - Graph* subgraph = entry.second; - ORT_RETURN_IF_ERROR(subgraph->VerifyNodeAndOpMatch(options)); - } - } - return Status::OK(); } Status Graph::VerifyInputAndInitializerNames() { - std::unordered_set& inputs_and_initializers = resolve_context_.inputs_and_initializers; + std::unordered_set& inputs_and_initializers = resolve_context_.inputs_and_initializers; - for (auto* input : GetInputs()) { + const auto& graph_inputs = GetInputs(); + inputs_and_initializers.reserve(graph_inputs.size() + name_to_initial_tensor_.size()); + + for (auto* input : graph_inputs) { auto result = inputs_and_initializers.insert(input->Name()); if (!result.second) { Status status(ONNXRUNTIME, onnxruntime::common::StatusCode::FAIL, @@ -2623,8 +2579,6 @@ Status Graph::VerifyInputAndInitializerNames() { } Status Graph::InitInputsInitializersOutputs() { - resolve_context_.Clear(); - // clear the previous relationships, as we re-create them when resolving. // same applies to the implicit input defs as they are built from any subgraphs within this graph. for (auto& node : Nodes()) { @@ -2708,13 +2662,10 @@ Status Graph::Resolve(const ResolveOptions& options) { return Status::OK(); } - // init all graph/subgraphs. non-recursive. + // init all graph/subgraphs. non-recursive so call via ForThisAndAllSubgraphs. auto init_func = [](Graph& graph) { return graph.InitInputsInitializersOutputs(); }; ORT_RETURN_IF_ERROR(ForThisAndAllSubgraphs(all_subgraphs, init_func)); - // recursively set the outer scope node args. - ORT_RETURN_IF_ERROR(SetOuterScopeNodeArgs(resolve_context_.outer_scope_node_args)); - std::unordered_set outer_scope_node_args_consumed; // recursively build connections between nodes in this graph and all subgraphs @@ -2742,6 +2693,8 @@ Status Graph::Resolve(const ResolveOptions& options) { graph.GraphProtoSyncNeeded(false); } + graph.resolve_context_.Clear(); + return Status::OK(); }; ORT_RETURN_IF_ERROR(ForThisAndAllSubgraphs(all_subgraphs, finalize_func)); @@ -2759,6 +2712,20 @@ void Graph::SetDescription(const std::string& description) { graph_proto_->set_doc_string(description); } +bool Graph::ResolveContext::IsLocalValue(const std::string& name) const { + return output_args.find(name) != output_args.cend() || + inputs_and_initializers.find(name) != inputs_and_initializers.cend(); +} + +bool Graph::ResolveContext::IsInputInitializerOrOutput(const std::string& name, bool check_ancestors) const { + return IsLocalValue(name) || + (check_ancestors && graph.parent_graph_ && + graph.parent_graph_->resolve_context_.IsInputInitializerOrOutput(name, check_ancestors)); +} + +bool Graph::ResolveContext::IsOuterScopeValue(const std::string& name) const { + return graph.parent_graph_ && graph.parent_graph_->resolve_context_.IsInputInitializerOrOutput(name, true); +} #endif // !defined(ORT_MINIMAL_BUILD) #if !defined(ORT_MINIMAL_BUILD) || defined(ORT_EXTENDED_MINIMAL_BUILD)