Remove unused initializer from GraphProto as well as name_to_initial_tensor_ in CleanUnusedInitializers. (#2320)

* Remove unused initializer from GraphProto as well as name_to_initial_tensor_ in CleanupUnusedInitializers.

This means initializers that have been replaced during graph optimizations are not left in the GraphProto when we save an optimized model.

* Handle edge case where a model has an unused initializer with matching graph input by also removing the graph input.

* Use non-const iterators in std::find_if calls to make centos build happy.
This commit is contained in:
Scott McKay 2019-11-08 16:29:50 +10:00 committed by GitHub
parent da3c0ba14b
commit 5a3ea7469a
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 48 additions and 12 deletions

View file

@ -2051,6 +2051,21 @@ void Graph::AddInitializedTensor(const TensorProto& tensor) {
}
}
template <typename T, typename TIter>
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(...)

View file

@ -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