mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Expose frozen_weights in PyTorch Frontend (#3317)
This commit is contained in:
parent
66c7579c93
commit
d143b41b81
1 changed files with 9 additions and 3 deletions
|
|
@ -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_:
|
||||
|
|
|
|||
Loading…
Reference in a new issue