From d143b41b8175ac227bbe24778ae51d8594bee14b Mon Sep 17 00:00:00 2001 From: Sherlock Date: Thu, 26 Mar 2020 11:26:54 -0700 Subject: [PATCH] Expose frozen_weights in PyTorch Frontend (#3317) --- orttraining/orttraining/python/ort_trainer.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/orttraining/orttraining/python/ort_trainer.py b/orttraining/orttraining/python/ort_trainer.py index e81f6ed740..240c9ad412 100644 --- a/orttraining/orttraining/python/ort_trainer.py +++ b/orttraining/orttraining/python/ort_trainer.py @@ -400,7 +400,8 @@ 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, enable_grad_norm_clip=True): + partition_optimizer=False, enable_grad_norm_clip=True, + frozen_weights=[]): output_name = model.graph.output[0].name ort_parameters = ort.TrainingParameters() ort_parameters.loss_output_name = output_name @@ -430,6 +431,8 @@ def create_ort_training_session_with_optimizer(model, device, training_optimizer optimizer_int_attributes_map = {} weights_to_train = set() for initializer in model.graph.initializer: + if initializer.name in frozen_weights: + continue weights_to_train.add(initializer.name) if map_optimizer_attributes is not None: attributes = map_optimizer_attributes(initializer.name) @@ -534,7 +537,8 @@ 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, enable_grad_norm_clip=True): + global_step=0, get_lr_this_step=None, loss_scaler=None, partition_optimizer=False, + enable_grad_norm_clip=True, frozen_weights=[]): super(ORTTrainer, self).__init__() """ Initializes ORTTrainer. @@ -609,6 +613,7 @@ class ORTTrainer(): self.allreduce_post_accumulation_ = allreduce_post_accumulation self.partition_optimizer_ = partition_optimizer self.enable_grad_norm_clip_ = enable_grad_norm_clip + self.frozen_weights_ = frozen_weights self.loss_scale_input_name = '' self._init_session() @@ -632,7 +637,8 @@ class ORTTrainer(): self.gradient_accumulation_steps, bind_parameters=False, use_mixed_precision=self.use_mixed_precision, allreduce_post_accumulation=self.allreduce_post_accumulation_, loss_scale_input_name=self.loss_scale_input_name, scaled_loss_output_name=self.scaled_loss_output_name, - partition_optimizer=self.partition_optimizer_, enable_grad_norm_clip=self.enable_grad_norm_clip_) + partition_optimizer=self.partition_optimizer_, enable_grad_norm_clip=self.enable_grad_norm_clip_, + frozen_weights=self.frozen_weights_) # ORT backend has modified model output dtype from float32 to float16. for o_desc in self.model_desc_.outputs_: