Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions include/onnxruntime/core/graph/graph_viewer.h
Original file line number Diff line number Diff line change
Expand Up @@ -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);

Expand Down
17 changes: 17 additions & 0 deletions onnxruntime/core/graph/graph_utils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down Expand Up @@ -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;
Comment thread
daquexian marked this conversation as resolved.
}

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
5 changes: 5 additions & 0 deletions onnxruntime/core/graph/graph_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
5 changes: 5 additions & 0 deletions onnxruntime/core/graph/graph_viewer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
#endif

#include "core/graph/graph_viewer.h"
#include "core/graph/model.h"

namespace onnxruntime {

Expand Down Expand Up @@ -109,4 +110,8 @@ bool GraphViewer::IsSubgraph() const {
return graph_->IsSubgraph();
}

ONNX_NAMESPACE::GraphProto GraphViewer::ToGraphProto() const {
return graph_->ToGraphProto();
}

} // namespace onnxruntime
33 changes: 2 additions & 31 deletions onnxruntime/core/providers/ngraph/ngraph_execution_provider.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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<onnxruntime::Node*>& fused_nodes,
Expand Down
47 changes: 4 additions & 43 deletions onnxruntime/core/providers/nnapi/nnapi_execution_provider.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -56,35 +57,8 @@ std::vector<std::vector<int>> NnapiExecutionProvider::GetSupportedNodes(const ON
std::vector<std::unique_ptr<ComputeCapability>>
NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph,
const std::vector<const KernelRegistry*>& /*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<NodeIndex>& node_index = graph.GetNodesInTopologicalOrder();
std::set<NodeArg*> all_node_inputs;
for (const auto& node : graph.Nodes()) {
std::vector<onnxruntime::NodeArg*> 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);

Expand Down Expand Up @@ -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<int, const NodeArg*>(it->second, it->first));
break;
}
}
if (std::find(graph_outputs.begin(), graph_outputs.end(), it->first) != graph_outputs.end()) {
outputs.insert(std::pair<int, const NodeArg*>(it->second, it->first));
}
outputs.insert(std::pair<int, const NodeArg*>(it->second, it->first));
}

// Assign inputs and outputs to subgraph's meta_def
Expand Down Expand Up @@ -238,12 +204,7 @@ common::Status NnapiExecutionProvider::Compile(const std::vector<onnxruntime::No
if (!func_body) {
return common::Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, "Function body is empty");
}
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();
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;
Expand Down
41 changes: 6 additions & 35 deletions onnxruntime/core/providers/openvino/openvino_execution_provider.cc
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
#include <vector>

#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"
Expand All @@ -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){

Expand Down Expand Up @@ -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") {
Expand Down Expand Up @@ -451,7 +422,7 @@ std::vector<std::unique_ptr<ComputeCapability>> OpenVINOExecutionProvider::GetCa
int counter = 0;
std::unique_ptr<IndexedSubGraph> sub_graph = std::make_unique<IndexedSubGraph>();

auto model_proto = GetModelProtoFromFusedNode(graph_viewer);
ONNX_NAMESPACE::ModelProto model_proto = graph_utils::GetModelProto(graph_viewer);

std::set<const onnxruntime::NodeArg*> fused_inputs, fused_outputs;

Expand Down
34 changes: 3 additions & 31 deletions onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down Expand Up @@ -219,32 +220,7 @@ SubGraphCollection_t TensorrtExecutionProvider::GetSupportedList(SubGraphCollect
std::vector<std::unique_ptr<ComputeCapability>>
TensorrtExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph,
const std::vector<const KernelRegistry*>& /*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<onnxruntime::NodeArg*> 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;
Expand Down Expand Up @@ -311,11 +287,7 @@ common::Status TensorrtExecutionProvider::Compile(const std::vector<onnxruntime:
if (!func_body) {
return common::Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, "Function body is empty");
}
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();
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);

Expand Down