From c4f827bee1a112538458cbdb965673520ddb01c4 Mon Sep 17 00:00:00 2001 From: Vincent Wang Date: Tue, 8 Dec 2020 03:38:09 +0000 Subject: [PATCH] remove initializers from original graph --- .../module_gradient_graph_builder.cc | 34 ++++++++++++++----- 1 file changed, 26 insertions(+), 8 deletions(-) diff --git a/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc b/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc index e9373e3554..fb7b21b3fa 100644 --- a/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc +++ b/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc @@ -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& graph_inputs = model_->MainGraph().GetInputsIncludingInitializers(); + Graph& graph = model_->MainGraph(); + const std::vector& graph_inputs = graph.GetInputsIncludingInitializers(); for (auto& node_arg : graph_inputs) { split_graphs_info_.user_input_names.emplace_back(node_arg->Name()); } - const std::vector& graph_outputs = model_->MainGraph().GetOutputs(); + const std::vector& 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 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& 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 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