diff --git a/include/onnxruntime/core/graph/graph_viewer.h b/include/onnxruntime/core/graph/graph_viewer.h index 7e2a0364ed0db..0656db845b31e 100644 --- a/include/onnxruntime/core/graph/graph_viewer.h +++ b/include/onnxruntime/core/graph/graph_viewer.h @@ -105,6 +105,8 @@ class GraphViewer { /** Check if this is a Subgraph */ bool IsSubgraph() const; + ONNX_NAMESPACE::GraphProto ToGraphProto() const; + private: ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(GraphViewer); diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index 2ac2a15303a11..66a8f54ac4f1e 100644 --- a/onnxruntime/core/graph/graph_utils.cc +++ b/onnxruntime/core/graph/graph_utils.cc @@ -3,6 +3,7 @@ #include "core/graph/graph_utils.h" #include "core/graph/graph.h" +#include "core/graph/model.h" #include "core/framework/tensorprotoutils.h" #include "core/common/logging/logging.h" @@ -445,6 +446,22 @@ size_t RemoveNodeOutputEdges(Graph& graph, Node& node) { return output_edges.size(); } +ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::GraphViewer& graph) { + onnxruntime::Model model(graph.Name(), true, ModelMetaData(), IOnnxRuntimeOpSchemaRegistryList(), graph.DomainToVersionMap()); + ONNX_NAMESPACE::ModelProto model_proto = model.ToProto(); + *model_proto.mutable_graph() = graph.ToGraphProto(); + return model_proto; +} + +ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::Function& func_body) { + // Reconstruct graph proto from fused node's function body + const Graph& graph_body = func_body.Body(); + onnxruntime::Model model(graph_body.Name(), true, ModelMetaData(), IOnnxRuntimeOpSchemaRegistryList(), graph_body.DomainToVersionMap()); + ONNX_NAMESPACE::ModelProto model_proto = model.ToProto(); + *model_proto.mutable_graph() = graph_body.ToGraphProto(); + return model_proto; +} + } // namespace graph_utils } // namespace onnxruntime diff --git a/onnxruntime/core/graph/graph_utils.h b/onnxruntime/core/graph/graph_utils.h index 9d60e70fbdc98..7106b48cf04cc 100644 --- a/onnxruntime/core/graph/graph_utils.h +++ b/onnxruntime/core/graph/graph_utils.h @@ -4,6 +4,7 @@ #pragma once #include "core/graph/onnx_protobuf.h" +#include "core/graph/graph_viewer.h" #include "core/graph/graph.h" namespace onnxruntime { @@ -100,6 +101,10 @@ bool RemoveNode(Graph& graph, Node& node); This should probably be elevated to the Graph API eventually. */ size_t RemoveNodeOutputEdges(Graph& graph, Node& node); +ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::GraphViewer& graph); + +ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::Function& func_body); + } // namespace graph_utils } // namespace onnxruntime diff --git a/onnxruntime/core/graph/graph_viewer.cc b/onnxruntime/core/graph/graph_viewer.cc index 262a2591ddb08..1b567dbd3fc54 100644 --- a/onnxruntime/core/graph/graph_viewer.cc +++ b/onnxruntime/core/graph/graph_viewer.cc @@ -7,6 +7,7 @@ #endif #include "core/graph/graph_viewer.h" +#include "core/graph/model.h" namespace onnxruntime { @@ -109,4 +110,8 @@ bool GraphViewer::IsSubgraph() const { return graph_->IsSubgraph(); } +ONNX_NAMESPACE::GraphProto GraphViewer::ToGraphProto() const { + return graph_->ToGraphProto(); +} + } // namespace onnxruntime diff --git a/onnxruntime/core/providers/ngraph/ngraph_execution_provider.cc b/onnxruntime/core/providers/ngraph/ngraph_execution_provider.cc index 459deae2c81f9..4ed3d252e51dd 100644 --- a/onnxruntime/core/providers/ngraph/ngraph_execution_provider.cc +++ b/onnxruntime/core/providers/ngraph/ngraph_execution_provider.cc @@ -6,6 +6,7 @@ #include "core/framework/compute_capability.h" #include "core/framework/allocatormgr.h" #include "core/framework/kernel_registry.h" +#include "core/graph/graph_utils.h" #include "core/graph/graph_viewer.h" #include "core/graph/model.h" #include "ngraph_execution_provider.h" @@ -476,37 +477,7 @@ NGRAPHExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_vie } static ONNX_NAMESPACE::ModelProto GetModelProtoFromFusedNode(const onnxruntime::Node* fused_node) { - const auto& attributes = fused_node->GetAttributes(); - const auto& initializers = attributes.at("initializers").tensors(); - - ONNX_NAMESPACE::ModelProto model_proto; - auto graph_proto = model_proto.mutable_graph(); - const auto& fused_graph = fused_node->GetFunctionBody()->Body(); - - for (const auto& node : fused_graph.Nodes()) { - node.ToProto(*(graph_proto->add_node())); - } - - for (const auto& input : fused_node->InputDefs()) { - auto valueInfoProto = graph_proto->add_input(); - *valueInfoProto = input->ToProto(); - } - - for (const auto& output : fused_node->OutputDefs()) { - auto valueInfoProto = graph_proto->add_output(); - *valueInfoProto = output->ToProto(); - } - - for (const auto& initializer : initializers) { - graph_proto->add_initializer()->CopyFrom(initializer); - } - - auto opset = model_proto.add_opset_import(); - opset->set_domain(kOnnxDomain); - opset->set_version(fused_graph.DomainToVersionMap().at(kOnnxDomain)); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); - - return model_proto; + return graph_utils::GetModelProto(*fused_node->GetFunctionBody()); } Status NGRAPHExecutionProvider::Compile(const std::vector& fused_nodes, diff --git a/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc index 83f9d8afed7b0..fbaf0073db5ce 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc @@ -5,6 +5,7 @@ #include "core/framework/compute_capability.h" #include "core/session/onnxruntime_cxx_api.h" #include "core/session/inference_session.h" +#include "core/graph/graph_utils.h" #include "core/graph/model.h" #include "dnnlibrary/ModelBuilder.h" #include "dnnlibrary/OnnxReader.h" @@ -56,35 +57,8 @@ std::vector> NnapiExecutionProvider::GetSupportedNodes(const ON std::vector> NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph, const std::vector& /*kernel_registries*/) const { - // This method is based on that of TRT EP - // Construct modelproto from graph - onnxruntime::Model model(graph.Name(), true, ModelMetaData(), IOnnxRuntimeOpSchemaRegistryList(), graph.DomainToVersionMap()); - onnxruntime::Graph& graph_build = model.MainGraph(); const std::vector& node_index = graph.GetNodesInTopologicalOrder(); - std::set all_node_inputs; - for (const auto& node : graph.Nodes()) { - std::vector inputs, outputs; - for (auto input : node.InputDefs()) { - auto& n_input = graph_build.GetOrCreateNodeArg(input->Name(), input->TypeAsProto()); - inputs.push_back(&n_input); - all_node_inputs.insert(&n_input); - } - for (auto output : node.OutputDefs()) { - auto& n_output = graph_build.GetOrCreateNodeArg(output->Name(), output->TypeAsProto()); - outputs.push_back(&n_output); - } - graph_build.AddNode(node.Name(), node.OpType(), node.Description(), inputs, outputs, &node.GetAttributes(), node.Domain()); - } - const auto graph_outputs = graph.GetOutputs(); - //Add initializer to graph - const auto& init_tensors = graph.GetAllInitializedTensors(); - for (const auto& tensor : init_tensors) { - graph_build.AddInitializedTensor(*(tensor.second)); - } - - ORT_ENFORCE(graph_build.Resolve().IsOK()); - ONNX_NAMESPACE::ModelProto model_proto = model.ToProto(); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(graph); const auto supported_nodes_vector = GetSupportedNodes(model_proto); @@ -174,15 +148,7 @@ NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph, } for (auto it = fused_outputs.begin(), end = fused_outputs.end(); it != end; ++it) { - for (const auto& x : all_node_inputs) { - if (x->Name() == it->first->Name()) { - outputs.insert(std::pair(it->second, it->first)); - break; - } - } - if (std::find(graph_outputs.begin(), graph_outputs.end(), it->first) != graph_outputs.end()) { - outputs.insert(std::pair(it->second, it->first)); - } + outputs.insert(std::pair(it->second, it->first)); } // Assign inputs and outputs to subgraph's meta_def @@ -238,12 +204,7 @@ common::Status NnapiExecutionProvider::Compile(const std::vectorBody(); - onnxruntime::Model model(graph_body.Name(), true, ModelMetaData(), - IOnnxRuntimeOpSchemaRegistryList(), graph_body.DomainToVersionMap()); - ONNX_NAMESPACE::ModelProto model_proto = model.ToProto(); - *(model_proto.mutable_graph()) = graph_body.ToGraphProto(); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(*func_body); dnn::OnnxReader onnx_reader; dnn::ModelBuilder model_builder; diff --git a/onnxruntime/core/providers/openvino/openvino_execution_provider.cc b/onnxruntime/core/providers/openvino/openvino_execution_provider.cc index 5deda941ff2e4..b2e2364e7783e 100644 --- a/onnxruntime/core/providers/openvino/openvino_execution_provider.cc +++ b/onnxruntime/core/providers/openvino/openvino_execution_provider.cc @@ -10,6 +10,7 @@ #include #include "core/common/common.h" +#include "core/graph/graph_utils.h" #include "core/graph/graph_viewer.h" #include "core/framework/compute_capability.h" #include "core/framework/tensorprotoutils.h" @@ -33,36 +34,6 @@ OpenVINOExecutionProvider::OpenVINOExecutionProvider(OpenVINOExecutionProviderIn InsertAllocator(CreateAllocator(device_info)); } -static ONNX_NAMESPACE::ModelProto GetModelProtoFromFusedNode(const onnxruntime::GraphViewer& graph_viewer) { - ONNX_NAMESPACE::ModelProto model_proto; - auto graph_proto = model_proto.mutable_graph(); - - for (const auto& node : graph_viewer.Nodes()) { - node.ToProto(*(graph_proto->add_node())); - } - - for (const auto& input : graph_viewer.GetInputs()) { - auto valueInfoProto = graph_proto->add_input(); - *valueInfoProto = input->ToProto(); - } - - for (const auto& output : graph_viewer.GetOutputs()) { - auto valueInfoProto = graph_proto->add_output(); - *valueInfoProto = output->ToProto(); - } - - for (const auto& initializer : graph_viewer.GetAllInitializedTensors()) { - graph_proto->add_initializer()->CopyFrom(*initializer.second); - } - - auto opset = model_proto.add_opset_import(); - opset->set_domain(kOnnxDomain); - opset->set_version(graph_viewer.DomainToVersionMap().at(kOnnxDomain)); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); - - return model_proto; -} - //Gets the input count of given node int GetInputCount(const Node* node, const InitializedTensorSet& initializer_set){ @@ -208,19 +179,19 @@ bool IsGraphSupported(const onnxruntime::GraphViewer& graph_viewer, std::string auto node_indexes = graph_viewer.GetNodesInTopologicalOrder(); - auto model_proto = GetModelProtoFromFusedNode(graph_viewer); + ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(graph_viewer); + const auto &graph_proto = model_proto.graph(); - auto graph_proto = model_proto.mutable_graph(); int input_dims = 0; int output_dims = 0; int num_inputs = graph_viewer.GetInputs().size(); int num_outputs = graph_viewer.GetOutputs().size(); if (num_inputs != 0) - input_dims = graph_proto->input(0).type().tensor_type().shape().dim_size(); + input_dims = graph_proto.input(0).type().tensor_type().shape().dim_size(); if (num_outputs != 0) - output_dims = graph_proto->output(0).type().tensor_type().shape().dim_size(); + output_dims = graph_proto.output(0).type().tensor_type().shape().dim_size(); //GPU Plugin does not support single dimensional input and 5 dimensional input if (dev_id == "GPU") { @@ -451,7 +422,7 @@ std::vector> OpenVINOExecutionProvider::GetCa int counter = 0; std::unique_ptr sub_graph = std::make_unique(); - auto model_proto = GetModelProtoFromFusedNode(graph_viewer); + ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(graph_viewer); std::set fused_inputs, fused_outputs; diff --git a/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc b/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc index 1ff6839af4dc5..e50cce56fdba2 100644 --- a/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc +++ b/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc @@ -14,6 +14,7 @@ #include "onnx/shape_inference/implementation.h" #include "cuda_runtime_api.h" #include "gsl/pointers" +#include "core/graph/graph_utils.h" #include "core/graph/model.h" #include "cuda_runtime_api.h" @@ -219,32 +220,7 @@ SubGraphCollection_t TensorrtExecutionProvider::GetSupportedList(SubGraphCollect std::vector> TensorrtExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph, const std::vector& /*kernel_registries*/) const { - // Construct modelproto from graph - onnxruntime::Model model(graph.Name(), true, ModelMetaData(), IOnnxRuntimeOpSchemaRegistryList(), graph.DomainToVersionMap()); - onnxruntime::Graph& graph_build = model.MainGraph(); - for (const auto& node : graph.Nodes()) { - std::vector inputs, outputs; - for (auto input : node.InputDefs()) { - auto& n_input = graph_build.GetOrCreateNodeArg(input->Name(), input->TypeAsProto()); - inputs.push_back(&n_input); - } - for (auto output : node.OutputDefs()) { - auto& n_output = graph_build.GetOrCreateNodeArg(output->Name(), output->TypeAsProto()); - outputs.push_back(&n_output); - } - graph_build.AddNode(node.Name(), node.OpType(), node.Description(), inputs, outputs, &node.GetAttributes(), node.Domain()); - } - - //Add initializer to graph - const auto& init_tensors = graph.GetAllInitializedTensors(); - for (const auto& tensor : init_tensors) { - graph_build.AddInitializedTensor(*(tensor.second)); - } - - auto status = graph_build.Resolve(); - ORT_ENFORCE(status.IsOK(), status); - ONNX_NAMESPACE::ModelProto model_proto = model.ToProto(); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(graph); // Serialize modelproto to string string string_buf; @@ -311,11 +287,7 @@ common::Status TensorrtExecutionProvider::Compile(const std::vectorBody(); - onnxruntime::Model model(graph_body.Name(), true, ModelMetaData(), IOnnxRuntimeOpSchemaRegistryList(), graph_body.DomainToVersionMap()); - ONNX_NAMESPACE::ModelProto model_proto = model.ToProto(); - *(model_proto.mutable_graph()) = graph_body.ToGraphProto(); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(*func_body); string string_buf; model_proto.SerializeToString(&string_buf);