diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index 81047b9398..882fdb3a02 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -2051,6 +2051,21 @@ void Graph::AddInitializedTensor(const TensorProto& tensor) { } } +template +static void RemoveRepeatedFieldEntry(T& repeated_field, const TIter& entry_to_remove) { + auto num_entries = repeated_field.size(); + if (num_entries > 1) { + // swap the entry being deleted with the last one, and delete it. + // we do this so we don't have to move all the entries past the one being deleted down one. + auto slot = entry_to_remove - repeated_field.begin(); + auto last_entry = repeated_field.end() - 1; + repeated_field.SwapElements(slot, num_entries - 1); + repeated_field.erase(last_entry); + } else { + repeated_field.erase(entry_to_remove); + } +} + void Graph::RemoveInitializedTensor(const std::string& tensor_name) { bool found = false; auto iter = name_to_initial_tensor_.find(tensor_name); @@ -2065,17 +2080,7 @@ void Graph::RemoveInitializedTensor(const std::string& tensor_name) { [&tensor_name](const TensorProto& entry) { return entry.name() == tensor_name; }); if (proto_entry != mutable_initializers.end()) { - auto num_entries = mutable_initializers.size(); - if (num_entries > 1) { - // swap the entry being deleted with the last one, and delete it. - // we do this so we don't have to move all the entries past the one being deleted down one. - auto slot = proto_entry - mutable_initializers.begin(); - auto last_entry = mutable_initializers.end() - 1; - mutable_initializers.SwapElements(slot, num_entries - 1); - mutable_initializers.erase(last_entry); - } else { - mutable_initializers.erase(proto_entry); - } + RemoveRepeatedFieldEntry(mutable_initializers, proto_entry); } else { // these should always be in sync as the pointer in name_to_initial_tensor_ is to memory owned by graph_proto_ ORT_ENFORCE(!found, "graph_proto_ is not in sync with name_to_initial_tensor_."); @@ -2409,7 +2414,28 @@ void Graph::CleanUnusedInitializers() { } std::for_each(erase_list.cbegin(), erase_list.cend(), - [this](const std::string& name) { name_to_initial_tensor_.erase(name); }); + [this](const std::string& name) { + RemoveInitializedTensor(name); + + // handle edge case where the unused initializer has a matching graph input + auto& proto_inputs = *graph_proto_->mutable_input(); + auto i = std::find_if(proto_inputs.begin(), proto_inputs.end(), + [&name](const ONNX_NAMESPACE::ValueInfoProto& input) { + return input.name() == name; + }); + + if (i != proto_inputs.end()) { + RemoveRepeatedFieldEntry(proto_inputs, i); + } + + auto& inputs_including_initializers = graph_inputs_including_initializers_; + auto j = std::find_if(inputs_including_initializers.begin(), inputs_including_initializers.end(), + [&name](const NodeArg* input) { return input->Name() == name; }); + + if (j != inputs_including_initializers.end()) { + inputs_including_initializers.erase(j); + } + }); } GSL_SUPPRESS(es .84) // warning about ignoring return value from insert(...) diff --git a/onnxruntime/test/ir/graph_test.cc b/onnxruntime/test/ir/graph_test.cc index 6c0fd59339..280e1337a3 100644 --- a/onnxruntime/test/ir/graph_test.cc +++ b/onnxruntime/test/ir/graph_test.cc @@ -20,6 +20,7 @@ #include "core/graph/graph_viewer.h" #include "core/graph/model.h" #include "core/graph/op.h" +#include "test/providers/provider_test_utils.h" #include "gtest/gtest.h" #include "gmock/gmock.h" @@ -1061,6 +1062,15 @@ TEST(GraphUpdateTest, AddRemoveInitializerHandling) { validate_proto(graph_proto_from_const_graph); validate_proto(graph_proto_from_graph); + + // Call Graph::Resolve which should remove the initializers from the Graph instance and proto as they're unused. + ASSERT_STATUS_OK(graph.Resolve()); + ASSERT_EQ(graph.GetAllInitializedTensors().size(), 0); + + ONNX_NAMESPACE::GraphProto graph_proto_from_resolved_graph = graph.ToGraphProto(); + auto num_initializers = graph_proto_from_resolved_graph.initializer_size(); + ASSERT_EQ(num_initializers, 0) << "Expected unused initializers to be removed from proto. " + << num_initializers << " remain."; } } // namespace test } // namespace onnxruntime