Expose frozen_weights in PyTorch Frontend (#3317)

This commit is contained in:
Sherlock 2020-03-26 11:26:54 -07:00 committed by GitHub
parent 66c7579c93
commit d143b41b81
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23

View file

@ -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_: