From d06763ac1ca8c50dc78f74fcd756931c1e4c39f2 Mon Sep 17 00:00:00 2001 From: ashbhandare Date: Fri, 24 Apr 2020 15:28:28 -0700 Subject: [PATCH] Set gradient as output only for easy mode (#3694) --- orttraining/orttraining/python/orttraining_pybind_state.cc | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/orttraining/orttraining/python/orttraining_pybind_state.cc b/orttraining/orttraining/python/orttraining_pybind_state.cc index 8e68e15d5c..4628bad416 100644 --- a/orttraining/orttraining/python/orttraining_pybind_state.cc +++ b/orttraining/orttraining/python/orttraining_pybind_state.cc @@ -94,7 +94,7 @@ TrainingConfigurationResult ConfigureSessionForTraining( config.weight_names_to_not_train = parameters.weights_not_to_train; config.immutable_weights = parameters.immutable_weights; - config.set_gradients_as_graph_outputs = true; + config.set_gradients_as_graph_outputs = false; config.gradient_accumulation_steps = parameters.gradient_accumulation_steps; @@ -115,6 +115,7 @@ TrainingConfigurationResult ConfigureSessionForTraining( config.loss_name = parameters.loss_output_name; if (!parameters.training_optimizer_name.empty()) { + config.set_gradients_as_graph_outputs = true; training::TrainingSession::TrainingConfiguration::OptimizerConfiguration opt{}; opt.name = parameters.training_optimizer_name; opt.learning_rate_input_name = parameters.lr_params_feed_name; @@ -276,4 +277,4 @@ void addObjectMethodsForTraining(py::module& m) { } } // namespace python -} // namespace onnxruntime \ No newline at end of file +} // namespace onnxruntime