From 7f610caca07d26094413172f53dcbe24b978a740 Mon Sep 17 00:00:00 2001 From: Tixxx Date: Mon, 23 Mar 2020 12:21:48 -0700 Subject: [PATCH] Make gradient clipping configurable. (#3243) * Make gradient clipping configurable. add control flag to c++ and python frontend --- orttraining/orttraining/core/graph/optimizer_config.h | 1 + .../orttraining/core/graph/optimizer_graph_builder.cc | 2 +- orttraining/orttraining/core/session/training_session.cc | 1 + orttraining/orttraining/core/session/training_session.h | 2 ++ orttraining/orttraining/models/bert/main.cc | 5 ++++- orttraining/orttraining/models/gpt2/main.cc | 5 ++++- orttraining/orttraining/models/runner/training_runner.cc | 1 + orttraining/orttraining/models/runner/training_runner.h | 2 ++ orttraining/orttraining/python/ort_trainer.py | 8 +++++--- .../orttraining/python/orttraining_pybind_state.cc | 8 +++++++- 10 files changed, 28 insertions(+), 7 deletions(-) diff --git a/orttraining/orttraining/core/graph/optimizer_config.h b/orttraining/orttraining/core/graph/optimizer_config.h index 09e64ba9be..1af10b256a 100644 --- a/orttraining/orttraining/core/graph/optimizer_config.h +++ b/orttraining/orttraining/core/graph/optimizer_config.h @@ -45,6 +45,7 @@ struct OptimizerGraphConfig { int64_t horovod_reduce_op{1}; std::string loss_scale_input_name{}; // empty string means no loss scaling factor is applied AdasumReductionType adasum_reduction_type{AdasumReductionType::None}; + bool enable_grad_norm_clip{true}; }; } // namespace training diff --git a/orttraining/orttraining/core/graph/optimizer_graph_builder.cc b/orttraining/orttraining/core/graph/optimizer_graph_builder.cc index 8ce9ea4d30..3993860b10 100644 --- a/orttraining/orttraining/core/graph/optimizer_graph_builder.cc +++ b/orttraining/orttraining/core/graph/optimizer_graph_builder.cc @@ -190,7 +190,7 @@ Status OptimizerGraphBuilder::BuildOptimizerNode( global_gradient_norm_argdef, global_gradient_norm_finite_argdef, opt_configs, graph_defs, new_initializers, - output_weight_argdefs, output_gradient_argdefs)); + output_weight_argdefs, output_gradient_argdefs, opt_graph_config_.enable_grad_norm_clip)); return Status::OK(); } diff --git a/orttraining/orttraining/core/session/training_session.cc b/orttraining/orttraining/core/session/training_session.cc index a39331cc58..4cee57fbf0 100644 --- a/orttraining/orttraining/core/session/training_session.cc +++ b/orttraining/orttraining/core/session/training_session.cc @@ -76,6 +76,7 @@ Status SetupOptimizerParams( opt_graph_config.allreduce_in_fp16 = optimizer_config.do_all_reduce_in_fp16; opt_graph_config.use_nccl = optimizer_config.use_nccl; opt_graph_config.adasum_reduction_type = optimizer_config.adasum_reduction_type; + opt_graph_config.enable_grad_norm_clip = optimizer_config.enable_grad_norm_clip; #if USE_HOROVOD opt_graph_config.horovod_reduce_op = opt_graph_config.adasum_reduction_type == AdasumReductionType::None diff --git a/orttraining/orttraining/core/session/training_session.h b/orttraining/orttraining/core/session/training_session.h index dcaa2e0fae..3d3710bc1c 100644 --- a/orttraining/orttraining/core/session/training_session.h +++ b/orttraining/orttraining/core/session/training_session.h @@ -129,6 +129,8 @@ class TrainingSession : public InferenceSession { bool partition_optimizer{}; // Selects the reduction algorithm for Adasum. AdasumReductionType adasum_reduction_type{AdasumReductionType::None}; + // Whether to enable gradient clipping. + bool enable_grad_norm_clip{true}; }; // The optimizer configuration. // If not provided, no optimizer is added. diff --git a/orttraining/orttraining/models/bert/main.cc b/orttraining/orttraining/models/bert/main.cc index d78dd2e556..ba1f91d4a3 100644 --- a/orttraining/orttraining/models/bert/main.cc +++ b/orttraining/orttraining/models/bert/main.cc @@ -140,7 +140,9 @@ Status ParseArguments(int argc, char* argv[], BertParameters& params, OrtParamet ("ratio_max", "Lamb max ratio parameter", cxxopts::value()->default_value("5.0")) ("cuda_mem_limit_in_gb", "Max cuda memory ort can use, in GB", cxxopts::value()->default_value("-1.0")) ("data_parallel_size", "Data parallel group size.", cxxopts::value()->default_value("1")) - ("horizontal_parallel_size", "Horizontal model parallel group size.", cxxopts::value()->default_value("1")); + ("horizontal_parallel_size", "Horizontal model parallel group size.", cxxopts::value()->default_value("1")) + ("enable_grad_norm_clip", "Specify whether to enable gradient clipping for optimizers.", + cxxopts::value()->default_value("true")); options .add_options("ORT configuration") ("ort_log_severity", "ORT minimum logging severity (see onnxruntime::logging::Severity values)", @@ -305,6 +307,7 @@ Status ParseArguments(int argc, char* argv[], BertParameters& params, OrtParamet } params.partition_optimizer = flags["partition_optimizer"].as(); + params.enable_grad_norm_clip = flags["enable_grad_norm_clip"].as(); float alpha = flags["alpha"].as(); float beta = flags["beta"].as(); float lambda = flags["lambda"].as(); diff --git a/orttraining/orttraining/models/gpt2/main.cc b/orttraining/orttraining/models/gpt2/main.cc index 0bf09479b5..f9f7e1038d 100644 --- a/orttraining/orttraining/models/gpt2/main.cc +++ b/orttraining/orttraining/models/gpt2/main.cc @@ -81,7 +81,9 @@ Status ParseArguments(int argc, char* argv[], GPT2Parameters& params, OrtParamet ("lambda", "Adam/Lamb lambda parameter", cxxopts::value()->default_value("0.01")) ("epsilon", "Adam/Lamb epsilon parameter", cxxopts::value()->default_value("1e-8")) ("data_parallel_size", "Data parallel group size.", cxxopts::value()->default_value("1")) - ("horizontal_parallel_size", "Horizontal model parallel group size.", cxxopts::value()->default_value("1")); + ("horizontal_parallel_size", "Horizontal model parallel group size.", cxxopts::value()->default_value("1")) + ("enable_grad_norm_clip", "Specify whether to enable gradient clipping for optimizers.", + cxxopts::value()->default_value("true")); options .add_options("ORT configuration") ("ort_log_severity", "ORT minimum logging severity (see onnxruntime::logging::Severity values)", @@ -197,6 +199,7 @@ Status ParseArguments(int argc, char* argv[], GPT2Parameters& params, OrtParamet } params.partition_optimizer = flags["partition_optimizer"].as(); + params.enable_grad_norm_clip = flags["enable_grad_norm_clip"].as(); float alpha = flags["alpha"].as(); float beta = flags["beta"].as(); float lambda = flags["lambda"].as(); diff --git a/orttraining/orttraining/models/runner/training_runner.cc b/orttraining/orttraining/models/runner/training_runner.cc index b6d2b32afa..ba79201d81 100644 --- a/orttraining/orttraining/models/runner/training_runner.cc +++ b/orttraining/orttraining/models/runner/training_runner.cc @@ -109,6 +109,7 @@ Status TrainingRunner::Initialize() { opt.use_nccl = params_.use_nccl; opt.partition_optimizer = params_.partition_optimizer; opt.adasum_reduction_type = params_.GetAdasumReductionType(); + opt.enable_grad_norm_clip = params_.enable_grad_norm_clip; config.optimizer_config = opt; } diff --git a/orttraining/orttraining/models/runner/training_runner.h b/orttraining/orttraining/models/runner/training_runner.h index 33c2d629f5..12ef6472d0 100644 --- a/orttraining/orttraining/models/runner/training_runner.h +++ b/orttraining/orttraining/models/runner/training_runner.h @@ -149,6 +149,8 @@ class TrainingRunner { int data_parallel_size = 1; int horizontal_parallel_size = 1; + // Enable gradient clipping. + bool enable_grad_norm_clip=true; }; TrainingRunner(Parameters params); diff --git a/orttraining/orttraining/python/ort_trainer.py b/orttraining/orttraining/python/ort_trainer.py index 6e9411eebd..32f2011fde 100644 --- a/orttraining/orttraining/python/ort_trainer.py +++ b/orttraining/orttraining/python/ort_trainer.py @@ -393,7 +393,7 @@ def create_ort_training_session_with_optimizer(model, device, training_optimizer gradient_accumulation_steps=1, bind_parameters=False, use_mixed_precision=False, allreduce_post_accumulation=False, loss_scale_input_name='', scaled_loss_output_name='', - partition_optimizer=False): + partition_optimizer=False, enable_grad_norm_clip=True): output_name = model.graph.output[0].name ort_parameters = ort.TrainingParameters() ort_parameters.loss_output_name = output_name @@ -410,6 +410,7 @@ def create_ort_training_session_with_optimizer(model, device, training_optimizer ort_parameters.scaled_loss_output_name = scaled_loss_output_name ort_parameters.allreduce_post_accumulation = allreduce_post_accumulation ort_parameters.partition_optimizer = partition_optimizer + ort_parameters.enable_grad_norm_clip = enable_grad_norm_clip output_types = {} for output in model.graph.output: @@ -515,7 +516,7 @@ class ORTTrainer(): def __init__(self, model, loss_fn, model_desc, training_optimizer_name, map_optimizer_attributes, learning_rate_description, device, gradient_accumulation_steps=1, postprocess_model=None, world_rank=0, world_size=1, use_mixed_precision=False, allreduce_post_accumulation=False, - global_step=0, get_lr_this_step=None, loss_scaler=None, partition_optimizer=False): + global_step=0, get_lr_this_step=None, loss_scaler=None, partition_optimizer=False, enable_grad_norm_clip=True): super(ORTTrainer, self).__init__() """ Initializes ORTTrainer. @@ -596,6 +597,7 @@ class ORTTrainer(): self.input_desc_with_lr_and_loss_scale = [*self.input_desc_with_lr, IODescription(self.loss_scale_input_name, [], torch.float32)] else: self.loss_scale_input_name, self.scaled_loss_output_name = '', '' + self.enable_grad_norm_clip_ = enable_grad_norm_clip self.verify_fully_optimized_model(self.onnx_model_) self.session, self.train_io_binding, self.eval_io_binding, self.output_name, _, self.output_types = \ @@ -606,7 +608,7 @@ class ORTTrainer(): self.gradient_accumulation_steps, bind_parameters=False, use_mixed_precision=use_mixed_precision, allreduce_post_accumulation=allreduce_post_accumulation, loss_scale_input_name=self.loss_scale_input_name, scaled_loss_output_name=self.scaled_loss_output_name, - partition_optimizer=partition_optimizer) + partition_optimizer=partition_optimizer, enable_grad_norm_clip=self.enable_grad_norm_clip_) # ORT backend has modified model output dtype from float32 to float16. for o_desc in self.model_desc_.outputs_: diff --git a/orttraining/orttraining/python/orttraining_pybind_state.cc b/orttraining/orttraining/python/orttraining_pybind_state.cc index eafbe9f616..c115f44600 100644 --- a/orttraining/orttraining/python/orttraining_pybind_state.cc +++ b/orttraining/orttraining/python/orttraining_pybind_state.cc @@ -52,6 +52,7 @@ struct TrainingParameters { int horizontal_parallel_size = 1; bool partition_optimizer = false; int seed = -1; + bool enable_grad_norm_clip = true; }; // TODO: this method does not handle parallel optimization. @@ -133,6 +134,9 @@ void ConfigureSessionForTraining( // an allreduce_post_accumulation option and remove the use_nccl option. opt.use_nccl = parameters.allreduce_post_accumulation; opt.partition_optimizer = parameters.partition_optimizer; + // TODO: The norm clipping value is 1.0f which is the default used in most frameworks. + // Need to have another option to support more values in the future. + opt.enable_grad_norm_clip = parameters.enable_grad_norm_clip; config.optimizer_config = opt; } @@ -166,7 +170,9 @@ void addObjectMethodsForTraining(py::module& m) { .def_readwrite("world_rank", &TrainingParameters::world_rank) .def_readwrite("world_size", &TrainingParameters::world_size) .def_readwrite("gradient_accumulation_steps", &TrainingParameters::gradient_accumulation_steps) - .def_readwrite("partition_optimizer", &TrainingParameters::partition_optimizer); + .def_readwrite("partition_optimizer", &TrainingParameters::partition_optimizer) + .def_readwrite("enable_grad_norm_clip", &TrainingParameters::enable_grad_norm_clip); + py::class_ training_session(m, "TrainingSession"); training_session.def(py::init())