mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
no need putting initializers (for constant node) into graph inputs. (#665)
* constant node should not be put into graph inputs any more. * simplify graph input/output set logic. * refactor comments. * remove adding initializers as graph inputs when creating graph from scratch.
This commit is contained in:
parent
60d71d63b5
commit
41dc3130f5
3 changed files with 103 additions and 182 deletions
|
|
@ -854,13 +854,6 @@ class Graph {
|
|||
|
||||
common::Status PerformTypeAndShapeInferencing();
|
||||
|
||||
enum class Type {
|
||||
// A main graph.
|
||||
Main = 1,
|
||||
// A sub graph (function).
|
||||
Sub = 2,
|
||||
};
|
||||
|
||||
common::Status Resolve(bool no_proto_sync_required);
|
||||
|
||||
// Recursively find all subgraphs including nested subgraphs
|
||||
|
|
@ -925,8 +918,6 @@ class Graph {
|
|||
InitializedTensorSet name_to_initial_tensor_;
|
||||
std::vector<int> removed_initializer_indexes_;
|
||||
|
||||
Type graph_type_ = Type::Main;
|
||||
|
||||
IOnnxRuntimeOpSchemaCollectionPtr schema_registry_;
|
||||
|
||||
std::vector<std::unique_ptr<onnxruntime::Function>> function_container_;
|
||||
|
|
|
|||
|
|
@ -608,7 +608,6 @@ Graph::Graph(GraphProto* graph_proto,
|
|||
Graph* parent_graph,
|
||||
const std::unordered_map<std::string, const ONNX_NAMESPACE::FunctionProto*>& model_functions)
|
||||
: graph_proto_{graph_proto},
|
||||
graph_type_{Type::Main},
|
||||
schema_registry_(schema_registry),
|
||||
graph_resolve_needed_(true),
|
||||
graph_proto_sync_needed_(false),
|
||||
|
|
@ -619,88 +618,72 @@ Graph::Graph(GraphProto* graph_proto,
|
|||
ORT_ENFORCE(graph_proto != nullptr, "graph_proto cannot be null");
|
||||
ArgNameToTypeMap name_to_type_map;
|
||||
|
||||
// these are all empty unless we received a graph_proto as input
|
||||
if (graph_proto != nullptr) {
|
||||
for (auto& node : graph_proto_->node()) {
|
||||
if (node.op_type() != kConstant) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Copy constant nodes _value to name_to_initial_tensor_
|
||||
for (auto& node : graph_proto_->node()) {
|
||||
if (node.op_type() == kConstant) {
|
||||
const gsl::not_null<TensorProto*> tensor{graph_proto_->add_initializer()};
|
||||
*tensor = node.attribute(0).t();
|
||||
*(tensor->mutable_name()) = node.output(0);
|
||||
const gsl::not_null<TensorProto*>
|
||||
tensor{graph_proto_->add_initializer()};
|
||||
*tensor = node.attribute(0).t();
|
||||
*(tensor->mutable_name()) = node.output(0);
|
||||
}
|
||||
|
||||
// we remove the node and add it as an initializer, but still need it to appear in the
|
||||
// graph inputs to make the ONNX checker happy. add a new input due to that.
|
||||
auto graph_inputs = graph_proto_->mutable_input();
|
||||
// Remove constant nodes as they're replaced with initializers above.
|
||||
const gsl::not_null<RepeatedPtrField<NodeProto>*> graph_mutable_nodes{graph_proto_->mutable_node()};
|
||||
graph_mutable_nodes->erase(
|
||||
std::remove_if(graph_mutable_nodes->begin(), graph_mutable_nodes->end(),
|
||||
[](NodeProto& p) {
|
||||
return (p.op_type() == kConstant);
|
||||
}),
|
||||
graph_mutable_nodes->end());
|
||||
|
||||
ValueInfoProto* value_info = graph_inputs->Add();
|
||||
value_info->set_name(node.output(0));
|
||||
value_info->set_doc_string("Input to represent replaced Constant node");
|
||||
// Copy initial tensors to a map.
|
||||
for (auto& tensor : graph_proto_->initializer()) {
|
||||
name_to_initial_tensor_[tensor.name()] = &tensor;
|
||||
|
||||
TypeProto t;
|
||||
t.mutable_tensor_type()->set_elem_type(tensor->data_type());
|
||||
auto shape = t.mutable_tensor_type()->mutable_shape();
|
||||
for (auto dim : tensor->dims())
|
||||
shape->add_dim()->set_dim_value(dim);
|
||||
// v4 does not require initializers to be inputs, so we need to ensure there is a NodeArg created for all
|
||||
// initializers in that case
|
||||
if (ir_version_ > 3) {
|
||||
TypeProto t;
|
||||
t.mutable_tensor_type()->set_elem_type(tensor.data_type());
|
||||
auto shape = t.mutable_tensor_type()->mutable_shape();
|
||||
for (auto dim : tensor.dims())
|
||||
shape->add_dim()->set_dim_value(dim);
|
||||
|
||||
(*value_info->mutable_type()) = t;
|
||||
}
|
||||
GetOrCreateNodeArg(tensor.name(), &t);
|
||||
}
|
||||
}
|
||||
|
||||
// remove constant nodes
|
||||
const gsl::not_null<RepeatedPtrField<NodeProto>*> graph_mutable_nodes{graph_proto_->mutable_node()};
|
||||
graph_mutable_nodes->erase(
|
||||
std::remove_if(graph_mutable_nodes->begin(), graph_mutable_nodes->end(),
|
||||
[](NodeProto& p) {
|
||||
return (p.op_type() == kConstant);
|
||||
}),
|
||||
graph_mutable_nodes->end());
|
||||
|
||||
// Copy initial tensors to a map.
|
||||
for (auto& tensor : graph_proto_->initializer()) {
|
||||
name_to_initial_tensor_[tensor.name()] = &tensor;
|
||||
|
||||
// v4 does not require initializers to be inputs, so we need to ensure there is a NodeArg created for all
|
||||
// initializers in that case
|
||||
if (ir_version > 3) {
|
||||
TypeProto t;
|
||||
t.mutable_tensor_type()->set_elem_type(tensor.data_type());
|
||||
auto shape = t.mutable_tensor_type()->mutable_shape();
|
||||
for (auto dim : tensor.dims())
|
||||
shape->add_dim()->set_dim_value(dim);
|
||||
|
||||
GetOrCreateNodeArg(tensor.name(), &t);
|
||||
}
|
||||
// Collect all node arg name, type, shape information in the graph.
|
||||
// type/shape information will be assigned to each node arg when going
|
||||
// thru all nodes later.
|
||||
for (auto& graph_input : graph_proto_->input()) {
|
||||
if (graph_input.has_name() && graph_input.has_type()) {
|
||||
name_to_type_map[graph_input.name()] = graph_input.type();
|
||||
// always create a NodeArg for graph input in case its from an initializer
|
||||
GetOrCreateNodeArg(graph_input.name(), &graph_input.type());
|
||||
}
|
||||
}
|
||||
|
||||
// Collect all node arg name, type, shape information in the graph.
|
||||
// type/shape information will be assigned to each node arg when going
|
||||
// thru all nodes later.
|
||||
for (auto& graph_input : graph_proto_->input()) {
|
||||
if (graph_input.has_name() && graph_input.has_type()) {
|
||||
name_to_type_map[graph_input.name()] = graph_input.type();
|
||||
// always create a NodeArg for graph input in case its from an initializer
|
||||
GetOrCreateNodeArg(graph_input.name(), &graph_input.type());
|
||||
}
|
||||
for (auto& graph_output : graph_proto_->output()) {
|
||||
if (graph_output.has_name() && graph_output.has_type()) {
|
||||
auto& name = graph_output.name();
|
||||
name_to_type_map[name] = graph_output.type();
|
||||
// always create NodeArg for graph output, in case it's from initializer
|
||||
GetOrCreateNodeArg(name, &graph_output.type());
|
||||
}
|
||||
}
|
||||
|
||||
for (auto& graph_output : graph_proto_->output()) {
|
||||
if (graph_output.has_name() && graph_output.has_type()) {
|
||||
auto& name = graph_output.name();
|
||||
name_to_type_map[name] = graph_output.type();
|
||||
// always create NodeArg for graph output, in case it's from initializer
|
||||
GetOrCreateNodeArg(name, &graph_output.type());
|
||||
}
|
||||
for (auto& node_arg : graph_proto_->value_info()) {
|
||||
if (node_arg.has_name() && node_arg.has_type()) {
|
||||
name_to_type_map[node_arg.name()] = node_arg.type();
|
||||
}
|
||||
}
|
||||
|
||||
for (auto& node_arg : graph_proto_->value_info()) {
|
||||
if (node_arg.has_name() && node_arg.has_type()) {
|
||||
name_to_type_map[node_arg.name()] = node_arg.type();
|
||||
}
|
||||
}
|
||||
|
||||
for (auto node_proto : graph_proto_->node()) {
|
||||
AddNode(node_proto, name_to_type_map);
|
||||
}
|
||||
for (auto node_proto : graph_proto_->node()) {
|
||||
AddNode(node_proto, name_to_type_map);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -952,7 +935,7 @@ Status Graph::BuildConnections(std::vector<std::string>& outer_scope_node_args_c
|
|||
AddEdge(output_node.Index(), node->Index(), entry->second.second, input_slot_index);
|
||||
|
||||
inner_nodes.insert(&output_node);
|
||||
|
||||
|
||||
// If this Graph was built manually, remove the implicit input from the graph outputs if it is present there
|
||||
// and not explicitly listed in the ordered graph outputs (as that implies we should leave it as an output).
|
||||
// If the Graph was loaded from a GraphProto, honor the explicit graph outputs and leave as is.
|
||||
|
|
@ -1834,8 +1817,7 @@ Status Graph::Resolve(bool no_proto_sync_required) {
|
|||
if (parent_graph_) {
|
||||
// Resolve must start at the top level graph in-order to handle outer scope
|
||||
// connections correctly, so recurse up to that level to start
|
||||
auto status = parent_graph_->Resolve(no_proto_sync_required);
|
||||
return status;
|
||||
return parent_graph_->Resolve(no_proto_sync_required);
|
||||
}
|
||||
|
||||
// find all subgraphs including nested ones.
|
||||
|
|
@ -2253,115 +2235,72 @@ Status Graph::SetGraphInputsOutputs() {
|
|||
std::unordered_set<std::string> added_input_names{outer_scope_node_arg_names_};
|
||||
|
||||
if (loaded_from_model_file) {
|
||||
// Collect all graph inputs/outputs specified in original graph proto
|
||||
std::unordered_set<std::string> specified_graph_inputs;
|
||||
std::unordered_set<std::string> specified_graph_outputs;
|
||||
std::unordered_set<std::string> specified_graph_value_info;
|
||||
std::unordered_set<std::string> specified_initializers;
|
||||
std::unordered_map<std::string, const NodeArg*> input_name_to_node_arg;
|
||||
std::unordered_map<std::string, const NodeArg*> output_name_to_node_arg;
|
||||
// Name to NodeArg mapping of all graph initializers.
|
||||
std::unordered_map<std::string, const NodeArg*> graph_initializers;
|
||||
|
||||
for (auto& graph_output : graph_proto_->output()) {
|
||||
specified_graph_outputs.insert(graph_output.name());
|
||||
}
|
||||
// Name to NodeArg mapping of all graph inputs.
|
||||
std::unordered_map<std::string, const NodeArg*> graph_inputs;
|
||||
|
||||
for (auto& graph_value_info : graph_proto_->value_info()) {
|
||||
specified_graph_value_info.insert(graph_value_info.name());
|
||||
}
|
||||
// Name to NodeArg mapping of all graph node outputs.
|
||||
std::unordered_map<std::string, const NodeArg*> nodes_outputs;
|
||||
|
||||
for (auto& initializer : graph_proto_->initializer()) {
|
||||
specified_initializers.insert(initializer.name());
|
||||
auto& initializer_name = initializer.name();
|
||||
auto initializer_arg = GetNodeArg(initializer_name);
|
||||
graph_initializers.insert({initializer_name, initializer_arg});
|
||||
}
|
||||
|
||||
// Set graph inputs.
|
||||
// <graph_inputs_including_initializers_> contains inputs exactly specified in proto.
|
||||
// <graph_inputs_excluding_initializers_> contains inputs without default value (specified as initializer).
|
||||
for (auto& graph_input : graph_proto_->input()) {
|
||||
// add all graph inputs to input_name_to_node_arg
|
||||
auto& name = graph_input.name();
|
||||
const auto* node_arg = GetNodeArg(name);
|
||||
ORT_ENFORCE(node_arg, "Graph ctor should have created NodeArg for initializer.");
|
||||
input_name_to_node_arg.insert({name, node_arg});
|
||||
|
||||
// only add non-initializer to specified_graph_inputs
|
||||
if (specified_initializers.find(name) == specified_initializers.end())
|
||||
specified_graph_inputs.insert(name);
|
||||
}
|
||||
|
||||
// add non-initializer outputs
|
||||
for (const auto& node : Nodes()) {
|
||||
for (const auto* output_def : node.OutputDefs()) {
|
||||
ORT_IGNORE_RETURN_VALUE(specified_graph_outputs.erase(output_def->Name()));
|
||||
output_name_to_node_arg.insert({output_def->Name(), output_def});
|
||||
graph_inputs.insert({name, node_arg});
|
||||
graph_inputs_including_initializers_.push_back(node_arg);
|
||||
if (graph_initializers.end() == graph_initializers.find(name)) {
|
||||
graph_inputs_excluding_initializers_.push_back(node_arg);
|
||||
}
|
||||
}
|
||||
|
||||
// add any outputs using initializer
|
||||
if (specified_graph_outputs.size() > 0) {
|
||||
for (const auto& name : specified_initializers) {
|
||||
ORT_IGNORE_RETURN_VALUE(specified_graph_outputs.erase(name));
|
||||
output_name_to_node_arg.insert({name, GetNodeArg(name)});
|
||||
}
|
||||
}
|
||||
|
||||
if (!specified_graph_outputs.empty()) {
|
||||
std::string missing_list;
|
||||
for (auto& name : specified_graph_outputs)
|
||||
missing_list += name + " ";
|
||||
return Status(ONNXRUNTIME, FAIL, "Some graph outputs do not exist in the graph. (" + missing_list + ")");
|
||||
}
|
||||
|
||||
for (const auto& node : Nodes()) {
|
||||
// Go thru all node's inputs.
|
||||
for (const auto* input_arg : node.InputDefs()) {
|
||||
if (!input_arg->Exists()) {
|
||||
// It's an optional input and does not exist in this case.
|
||||
continue;
|
||||
}
|
||||
|
||||
if (specified_graph_inputs.end() != specified_graph_inputs.find(input_arg->Name())) {
|
||||
if (added_input_names.insert(input_arg->Name()).second) {
|
||||
// The node input is specified as graph input.
|
||||
input_name_to_node_arg.insert({input_arg->Name(), input_arg});
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
auto output_arg_iter = output_name_to_node_arg.find(input_arg->Name());
|
||||
if (output_name_to_node_arg.end() == output_arg_iter &&
|
||||
specified_initializers.end() == specified_initializers.find(input_arg->Name())) {
|
||||
// The node input is not specified as graph input,
|
||||
// and it's not fed by another node neither.
|
||||
if (!IsSubgraph()) {
|
||||
return Status(ONNXRUNTIME, FAIL,
|
||||
"Node input (" + input_arg->Name() + ") should be a graph input or initializer.");
|
||||
}
|
||||
|
||||
// TODO: Do we need to do a comprehensive check that the input is coming from the outer scope or is it
|
||||
// fine to catch this issue later?
|
||||
}
|
||||
|
||||
if (specified_graph_value_info.erase(input_arg->Name()) >= 1) {
|
||||
value_info_.push_back(input_arg);
|
||||
}
|
||||
for (const auto* output_def : node.OutputDefs()) {
|
||||
nodes_outputs.insert({output_def->Name(), output_def});
|
||||
}
|
||||
}
|
||||
|
||||
// preserve input order
|
||||
for (auto& graph_input : graph_proto_->input()) {
|
||||
auto& name = graph_input.name();
|
||||
auto node_arg_iter = input_name_to_node_arg.find(name);
|
||||
ORT_ENFORCE(node_arg_iter != input_name_to_node_arg.cend(),
|
||||
"All inputs and initializers should have entries. Missing ", name);
|
||||
|
||||
graph_inputs_including_initializers_.push_back(node_arg_iter->second);
|
||||
|
||||
if (specified_initializers.find(name) == specified_initializers.end()) {
|
||||
graph_inputs_excluding_initializers_.push_back(node_arg_iter->second);
|
||||
}
|
||||
}
|
||||
|
||||
// preserve output order
|
||||
// Set graph outputs.
|
||||
// Graph outputs specified in the model must be nodes' outputs, initailizer or graph inputs.
|
||||
for (auto& graph_output : graph_proto_->output()) {
|
||||
graph_outputs_.push_back(output_name_to_node_arg.at(graph_output.name()));
|
||||
auto& graph_output_name = graph_output.name();
|
||||
auto iter = nodes_outputs.find(graph_output_name);
|
||||
if (nodes_outputs.end() == iter) {
|
||||
// Graph output is not found as any node's output.
|
||||
auto iter2 = graph_initializers.find(graph_output_name);
|
||||
if (graph_initializers.end() == iter2) {
|
||||
// Graph output is not found as any initializer.
|
||||
auto iter3 = graph_inputs.find(graph_output_name);
|
||||
if (graph_inputs.end() == iter3) {
|
||||
// Graph output is not found as any graph input.
|
||||
return Status(ONNXRUNTIME, FAIL, "Graph output (" + graph_output_name + ") does not exist in the graph.");
|
||||
}
|
||||
graph_outputs_.push_back(iter3->second);
|
||||
continue;
|
||||
}
|
||||
graph_outputs_.push_back(iter2->second);
|
||||
continue;
|
||||
}
|
||||
graph_outputs_.push_back(iter->second);
|
||||
}
|
||||
|
||||
// Set graph value_info_.
|
||||
for (auto& graph_value_info : graph_proto_->value_info()) {
|
||||
auto& name = graph_value_info.name();
|
||||
const auto* node_arg = GetNodeArg(name);
|
||||
value_info_.push_back(node_arg);
|
||||
}
|
||||
|
||||
} else {
|
||||
std::unordered_map<std::string, const NodeArg*> output_name_to_node_arg;
|
||||
std::vector<std::string> ordered_output_names;
|
||||
|
|
@ -2438,14 +2377,6 @@ Status Graph::SetGraphInputsOutputs() {
|
|||
}
|
||||
}
|
||||
|
||||
// Make sure all initializers appear as graph inputs as per ONNX requirements
|
||||
for (auto i : name_to_initial_tensor_) {
|
||||
if (added_input_names.find(i.first) == added_input_names.cend()) {
|
||||
auto* na = GetNodeArg(i.first);
|
||||
graph_inputs_including_initializers_.push_back(na);
|
||||
}
|
||||
}
|
||||
|
||||
// Set graph outputs
|
||||
auto end = graph_output_args.end();
|
||||
for (auto& name : ordered_output_names) {
|
||||
|
|
|
|||
|
|
@ -653,8 +653,7 @@ TEST(ResolvingGraphTest, GraphConstruction_OnlyInitializer) {
|
|||
EXPECT_TRUE(status.IsOK()) << status.ErrorMessage();
|
||||
|
||||
auto& iii = graph.GetInputsIncludingInitializers();
|
||||
EXPECT_TRUE(iii.size() == 1);
|
||||
EXPECT_TRUE(iii.front()->Name() == "node_1_in_2");
|
||||
EXPECT_TRUE(iii.size() == 0);
|
||||
}
|
||||
|
||||
TEST(ResolvingGraphTest, GraphConstruction_TypeInference) {
|
||||
|
|
|
|||
Loading…
Reference in a new issue