mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
remove initializers from original graph
This commit is contained in:
parent
f7f435fc27
commit
c4f827bee1
1 changed files with 26 additions and 8 deletions
|
|
@ -55,12 +55,13 @@ Status ModuleGradientGraphBuilder::Initialize(std::istream& model_istream,
|
|||
ORT_RETURN_IF_ERROR(Model::Load(model_proto, model_, nullptr, *logger_));
|
||||
|
||||
// Handle original model inputs, outputs and trainable initializers.
|
||||
const std::vector<const NodeArg*>& graph_inputs = model_->MainGraph().GetInputsIncludingInitializers();
|
||||
Graph& graph = model_->MainGraph();
|
||||
const std::vector<const NodeArg*>& graph_inputs = graph.GetInputsIncludingInitializers();
|
||||
for (auto& node_arg : graph_inputs) {
|
||||
split_graphs_info_.user_input_names.emplace_back(node_arg->Name());
|
||||
}
|
||||
|
||||
const std::vector<const NodeArg*>& graph_outputs = model_->MainGraph().GetOutputs();
|
||||
const std::vector<const NodeArg*>& graph_outputs = graph.GetOutputs();
|
||||
for (auto& node_arg : graph_outputs) {
|
||||
split_graphs_info_.user_output_names.emplace_back(node_arg->Name());
|
||||
}
|
||||
|
|
@ -68,6 +69,19 @@ Status ModuleGradientGraphBuilder::Initialize(std::istream& model_istream,
|
|||
split_graphs_info_.initializer_names_to_train.assign(config.initializer_names_to_train.begin(),
|
||||
config.initializer_names_to_train.end());
|
||||
|
||||
// Remove the training initializers from the graph and move them to input to save memory.
|
||||
std::vector<const NodeArg*> input_args;
|
||||
for (const auto& input_name : split_graphs_info_.user_input_names) {
|
||||
input_args.emplace_back(graph.GetNodeArg(input_name));
|
||||
}
|
||||
|
||||
for (const auto& initializer_name : split_graphs_info_.initializer_names_to_train) {
|
||||
input_args.emplace_back(graph.GetNodeArg(initializer_name));
|
||||
graph.RemoveInitializedTensor(initializer_name);
|
||||
}
|
||||
|
||||
graph.SetInputs(input_args);
|
||||
|
||||
config_ = config;
|
||||
return Status::OK();
|
||||
}
|
||||
|
|
@ -94,6 +108,12 @@ Status ModuleGradientGraphBuilder::BuildAndSplit(const std::vector<std::vector<i
|
|||
input_index++;
|
||||
}
|
||||
|
||||
// Move over all training initializer inputs. They already have the concrete shapes.
|
||||
const std::vector<const NodeArg*>& graph_inputs = graph.GetInputsIncludingInitializers();
|
||||
for (; input_index < graph_inputs.size(); input_index++) {
|
||||
input_args.emplace_back(graph_inputs[input_index]);
|
||||
}
|
||||
|
||||
graph.SetInputs(input_args);
|
||||
ORT_RETURN_IF_ERROR(graph.Resolve());
|
||||
|
||||
|
|
@ -133,8 +153,7 @@ Status ModuleGradientGraphBuilder::BuildAndSplit(const std::vector<std::vector<i
|
|||
gradient_graph_config.set_gradients_as_graph_outputs = config_.set_gradients_as_graph_outputs;
|
||||
std::unordered_set<std::string> y_node_arg_names(split_graphs_info_.user_output_names.begin(),
|
||||
split_graphs_info_.user_output_names.end());
|
||||
GradientGraphBuilder grad_graph_builder(&graph, y_node_arg_names, x_node_arg_names,
|
||||
"",
|
||||
GradientGraphBuilder grad_graph_builder(&graph, y_node_arg_names, x_node_arg_names, "",
|
||||
gradient_graph_config, *logger_);
|
||||
ORT_RETURN_IF_ERROR(grad_graph_builder.Build());
|
||||
|
||||
|
|
@ -149,8 +168,8 @@ Status ModuleGradientGraphBuilder::BuildAndSplit(const std::vector<std::vector<i
|
|||
}
|
||||
|
||||
input_args.clear();
|
||||
for (auto& input_name : split_graphs_info_.user_input_names) {
|
||||
input_args.emplace_back(graph.GetNodeArg(input_name));
|
||||
for (const NodeArg* input_node_arg: graph.GetInputsIncludingInitializers()) {
|
||||
input_args.emplace_back(input_node_arg);
|
||||
}
|
||||
|
||||
// Add the entry points of gradients (normally loss_gard) to the graph inputs. Using the order of graph outputs.
|
||||
|
|
@ -267,7 +286,6 @@ Status ModuleGradientGraphBuilder::Split() {
|
|||
// Add initializers to forward graph inputs.
|
||||
for (const auto& initializer_name : split_graphs_info_.initializer_names_to_train) {
|
||||
forward_input_args.emplace_back(forward_graph.GetNodeArg(initializer_name));
|
||||
forward_graph.RemoveInitializedTensor(initializer_name);
|
||||
}
|
||||
|
||||
forward_graph.SetInputs(forward_input_args);
|
||||
|
|
@ -305,6 +323,7 @@ Status ModuleGradientGraphBuilder::Split() {
|
|||
}
|
||||
|
||||
RemoveNodes(backward_graph, backward_nodes_to_remove);
|
||||
FilterInitializers(backward_graph, backward_input_names);
|
||||
|
||||
// User inputs to backward graph inputs.
|
||||
split_graphs_info_.backward_user_input_names.clear();
|
||||
|
|
@ -352,7 +371,6 @@ Status ModuleGradientGraphBuilder::Split() {
|
|||
}
|
||||
|
||||
backward_graph.SetOutputs(backward_output_args);
|
||||
FilterInitializers(backward_graph, backward_input_names);
|
||||
backward_graph.Resolve();
|
||||
return Status::OK();
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue