From 1df7175863a944e7544c43e572e1ab1d4d63aba4 Mon Sep 17 00:00:00 2001 From: daquexian Date: Thu, 4 Jul 2019 14:03:05 +0800 Subject: [PATCH 1/7] Add tographproto in graphviewer --- include/onnxruntime/core/graph/graph_viewer.h | 2 + onnxruntime/core/graph/graph_viewer.cc | 5 ++ .../nnapi/nnapi_execution_provider.cc | 72 ++++++++++--------- 3 files changed, 46 insertions(+), 33 deletions(-) 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_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/nnapi/nnapi_execution_provider.cc b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc index 83f9d8afed7b0..1ffa790b329e4 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc @@ -56,35 +56,40 @@ 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(); + // // 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); + 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; + *model_proto.mutable_graph() = graph.ToGraphProto(); const auto supported_nodes_vector = GetSupportedNodes(model_proto); @@ -174,12 +179,13 @@ 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; - } - } + // for (const auto& x : all_node_inputs) { + // if (x->Name() == it->first->Name()) { + // outputs.insert(std::pair(it->second, it->first)); + // break; + // } + // } + outputs.insert(std::pair(it->second, it->first)); if (std::find(graph_outputs.begin(), graph_outputs.end(), it->first) != graph_outputs.end()) { outputs.insert(std::pair(it->second, it->first)); } From 3aa24bd57c761e1afaa9fe765a901324f7dd10d9 Mon Sep 17 00:00:00 2001 From: daquexian Date: Thu, 4 Jul 2019 18:56:55 +0800 Subject: [PATCH 2/7] Fix the duplicate output --- .../nnapi/nnapi_execution_provider.cc | 39 ------------------- 1 file changed, 39 deletions(-) diff --git a/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc index 1ffa790b329e4..93e30f52d00ce 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc @@ -56,36 +56,6 @@ 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); - const std::vector& node_index = graph.GetNodesInTopologicalOrder(); const auto graph_outputs = graph.GetOutputs(); ONNX_NAMESPACE::ModelProto model_proto; @@ -179,16 +149,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; - // } - // } outputs.insert(std::pair(it->second, it->first)); - if (std::find(graph_outputs.begin(), graph_outputs.end(), it->first) != graph_outputs.end()) { - outputs.insert(std::pair(it->second, it->first)); - } } // Assign inputs and outputs to subgraph's meta_def From 1c27f784374cee59dd790f31e2f30b9e24990838 Mon Sep 17 00:00:00 2001 From: daquexian Date: Fri, 5 Jul 2019 10:38:12 +0800 Subject: [PATCH 3/7] Remove unused variable --- onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc | 1 - 1 file changed, 1 deletion(-) diff --git a/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc index 93e30f52d00ce..f48a499a7ef42 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc @@ -57,7 +57,6 @@ std::vector> NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph, const std::vector& /*kernel_registries*/) const { const std::vector& node_index = graph.GetNodesInTopologicalOrder(); - const auto graph_outputs = graph.GetOutputs(); ONNX_NAMESPACE::ModelProto model_proto; *model_proto.mutable_graph() = graph.ToGraphProto(); From 0e3222b402af1e1828c0bc908dbdc86bab022be0 Mon Sep 17 00:00:00 2001 From: daquexian Date: Fri, 5 Jul 2019 22:20:47 +0800 Subject: [PATCH 4/7] Add ToModelProto for GraphViewer and Function --- onnxruntime/core/graph/graph_utils.cc | 16 ++++++++ onnxruntime/core/graph/graph_utils.h | 5 +++ .../ngraph/ngraph_execution_provider.cc | 33 +-------------- .../nnapi/nnapi_execution_provider.cc | 11 ++--- .../openvino/openvino_execution_provider.cc | 41 +++---------------- .../tensorrt/tensorrt_execution_provider.cc | 34 ++------------- 6 files changed, 35 insertions(+), 105 deletions(-) diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index 2ca9358ac271e..e4f2d1849d012 100644 --- a/onnxruntime/core/graph/graph_utils.cc +++ b/onnxruntime/core/graph/graph_utils.cc @@ -448,6 +448,22 @@ size_t RemoveNodeOutputEdges(Graph& graph, Node& node) { return output_edges.size(); } +ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::GraphViewer& graph) { + ONNX_NAMESPACE::ModelProto model_proto; + *model_proto.mutable_graph() = graph.ToGraphProto(); + model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + 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(); + ONNX_NAMESPACE::ModelProto model_proto; + *model_proto.mutable_graph() = graph_body.ToGraphProto(); + model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + 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 a3be17f2664a9..50ab3d223d61d 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 { @@ -88,6 +89,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/providers/ngraph/ngraph_execution_provider.cc b/onnxruntime/core/providers/ngraph/ngraph_execution_provider.cc index 7749a21532430..f1e0e6cd0f15f 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" @@ -494,37 +495,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 f48a499a7ef42..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" @@ -57,8 +58,7 @@ std::vector> NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph, const std::vector& /*kernel_registries*/) const { const std::vector& node_index = graph.GetNodesInTopologicalOrder(); - ONNX_NAMESPACE::ModelProto model_proto; - *model_proto.mutable_graph() = graph.ToGraphProto(); + ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(graph); const auto supported_nodes_vector = GetSupportedNodes(model_proto); @@ -204,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 f86e5bdaff288..0d703b66aab23 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 af0bbdfe66902..70f0fd469374e 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; @@ -317,11 +293,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); From 6db058125dbf8528a2e31e94a9555974fae54ecd Mon Sep 17 00:00:00 2001 From: daquexian Date: Wed, 10 Jul 2019 13:44:01 +0800 Subject: [PATCH 5/7] Set ir version according to that of the graph --- include/onnxruntime/core/graph/graph.h | 8 ++++---- include/onnxruntime/core/graph/graph_viewer.h | 4 ++++ onnxruntime/core/graph/graph_utils.cc | 4 ++-- 3 files changed, 10 insertions(+), 6 deletions(-) diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index 4abd12499e32c..921e6aba05544 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -715,6 +715,10 @@ class Graph { /** Gets the ISchemaRegistry instances being used with this Graph. */ IOnnxRuntimeOpSchemaCollectionPtr GetSchemaRegistry() const; + Version IrVersion() const noexcept { + return ir_version_; + } + /** Create a single Node that is the result of the a fusion of multiple nodes in this Graph. @param sub_graph A IndexSubGraph instance with details of the nodes to fuse. @@ -794,10 +798,6 @@ class Graph { Node& AddNode(const ONNX_NAMESPACE::NodeProto& node_proto, const ArgNameToTypeMap& name_to_type); - Version IrVersion() const noexcept { - return ir_version_; - } - Graph& GraphResolveNeeded(bool needed) noexcept { graph_resolve_needed_ = needed; return *this; diff --git a/include/onnxruntime/core/graph/graph_viewer.h b/include/onnxruntime/core/graph/graph_viewer.h index 0656db845b31e..b464dd7be5381 100644 --- a/include/onnxruntime/core/graph/graph_viewer.h +++ b/include/onnxruntime/core/graph/graph_viewer.h @@ -102,6 +102,10 @@ class GraphViewer { return graph_->DomainToVersionMap(); } + Version IrVersion() const noexcept { + return graph_->IrVersion(); + } + /** Check if this is a Subgraph */ bool IsSubgraph() const; diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index e4f2d1849d012..c11598d07d7cd 100644 --- a/onnxruntime/core/graph/graph_utils.cc +++ b/onnxruntime/core/graph/graph_utils.cc @@ -451,7 +451,7 @@ size_t RemoveNodeOutputEdges(Graph& graph, Node& node) { ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::GraphViewer& graph) { ONNX_NAMESPACE::ModelProto model_proto; *model_proto.mutable_graph() = graph.ToGraphProto(); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + model_proto.set_ir_version(graph.IrVersion()); return model_proto; } @@ -460,7 +460,7 @@ ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::Function& func_body) const Graph& graph_body = func_body.Body(); ONNX_NAMESPACE::ModelProto model_proto; *model_proto.mutable_graph() = graph_body.ToGraphProto(); - model_proto.set_ir_version(ONNX_NAMESPACE::Version::IR_VERSION); + model_proto.set_ir_version(graph_body.IrVersion()); return model_proto; } From 0159718343591aa6e3198a01f683c760fc2cb95b Mon Sep 17 00:00:00 2001 From: daquexian Date: Wed, 10 Jul 2019 13:44:09 +0800 Subject: [PATCH 6/7] set opset_import --- onnxruntime/core/graph/graph_utils.cc | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index c11598d07d7cd..fdff04dd69fca 100644 --- a/onnxruntime/core/graph/graph_utils.cc +++ b/onnxruntime/core/graph/graph_utils.cc @@ -452,6 +452,11 @@ ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::GraphViewer& graph) ONNX_NAMESPACE::ModelProto model_proto; *model_proto.mutable_graph() = graph.ToGraphProto(); model_proto.set_ir_version(graph.IrVersion()); + for (const auto &opset_import : graph.DomainToVersionMap()) { + auto *opset_import_add = model_proto.add_opset_import(); + opset_import_add->set_domain(opset_import.first); + opset_import_add->set_version(opset_import.second); + } return model_proto; } @@ -461,6 +466,11 @@ ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::Function& func_body) ONNX_NAMESPACE::ModelProto model_proto; *model_proto.mutable_graph() = graph_body.ToGraphProto(); model_proto.set_ir_version(graph_body.IrVersion()); + for (const auto &opset_import : graph_body.DomainToVersionMap()) { + auto *opset_import_add = model_proto.add_opset_import(); + opset_import_add->set_domain(opset_import.first); + opset_import_add->set_version(opset_import.second); + } return model_proto; } From 1314a5513e2d4a78ee7fc34601830c49dfb64661 Mon Sep 17 00:00:00 2001 From: daquexian Date: Wed, 10 Jul 2019 18:07:57 +0800 Subject: [PATCH 7/7] Get ModelProto from onnxruntime::Model --- include/onnxruntime/core/graph/graph.h | 8 ++++---- include/onnxruntime/core/graph/graph_viewer.h | 4 ---- onnxruntime/core/graph/graph_utils.cc | 19 +++++-------------- 3 files changed, 9 insertions(+), 22 deletions(-) diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index 921e6aba05544..4abd12499e32c 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -715,10 +715,6 @@ class Graph { /** Gets the ISchemaRegistry instances being used with this Graph. */ IOnnxRuntimeOpSchemaCollectionPtr GetSchemaRegistry() const; - Version IrVersion() const noexcept { - return ir_version_; - } - /** Create a single Node that is the result of the a fusion of multiple nodes in this Graph. @param sub_graph A IndexSubGraph instance with details of the nodes to fuse. @@ -798,6 +794,10 @@ class Graph { Node& AddNode(const ONNX_NAMESPACE::NodeProto& node_proto, const ArgNameToTypeMap& name_to_type); + Version IrVersion() const noexcept { + return ir_version_; + } + Graph& GraphResolveNeeded(bool needed) noexcept { graph_resolve_needed_ = needed; return *this; diff --git a/include/onnxruntime/core/graph/graph_viewer.h b/include/onnxruntime/core/graph/graph_viewer.h index b464dd7be5381..0656db845b31e 100644 --- a/include/onnxruntime/core/graph/graph_viewer.h +++ b/include/onnxruntime/core/graph/graph_viewer.h @@ -102,10 +102,6 @@ class GraphViewer { return graph_->DomainToVersionMap(); } - Version IrVersion() const noexcept { - return graph_->IrVersion(); - } - /** Check if this is a Subgraph */ bool IsSubgraph() const; diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index fdff04dd69fca..ca105c2063cef 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" @@ -449,28 +450,18 @@ size_t RemoveNodeOutputEdges(Graph& graph, Node& node) { } ONNX_NAMESPACE::ModelProto GetModelProto(const onnxruntime::GraphViewer& graph) { - ONNX_NAMESPACE::ModelProto model_proto; + onnxruntime::Model model(graph.Name(), true, ModelMetaData(), IOnnxRuntimeOpSchemaRegistryList(), graph.DomainToVersionMap()); + ONNX_NAMESPACE::ModelProto model_proto = model.ToProto(); *model_proto.mutable_graph() = graph.ToGraphProto(); - model_proto.set_ir_version(graph.IrVersion()); - for (const auto &opset_import : graph.DomainToVersionMap()) { - auto *opset_import_add = model_proto.add_opset_import(); - opset_import_add->set_domain(opset_import.first); - opset_import_add->set_version(opset_import.second); - } 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(); - ONNX_NAMESPACE::ModelProto model_proto; + 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(graph_body.IrVersion()); - for (const auto &opset_import : graph_body.DomainToVersionMap()) { - auto *opset_import_add = model_proto.add_opset_import(); - opset_import_add->set_domain(opset_import.first); - opset_import_add->set_version(opset_import.second); - } return model_proto; }