Tie weights after preparing the model in run_clm (#18855)
This commit is contained in:
@@ -477,10 +477,6 @@ def main():
|
|||||||
]
|
]
|
||||||
optimizer = torch.optim.AdamW(optimizer_grouped_parameters, lr=args.learning_rate)
|
optimizer = torch.optim.AdamW(optimizer_grouped_parameters, lr=args.learning_rate)
|
||||||
|
|
||||||
# On TPU, the tie weights in our model have been disconnected, so we need to restore the ties.
|
|
||||||
if accelerator.distributed_type == DistributedType.TPU:
|
|
||||||
model.tie_weights()
|
|
||||||
|
|
||||||
# Scheduler and math around the number of training steps.
|
# Scheduler and math around the number of training steps.
|
||||||
overrode_max_train_steps = False
|
overrode_max_train_steps = False
|
||||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
|
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
|
||||||
@@ -500,6 +496,10 @@ def main():
|
|||||||
model, optimizer, train_dataloader, eval_dataloader, lr_scheduler
|
model, optimizer, train_dataloader, eval_dataloader, lr_scheduler
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# On TPU, the tie weights in our model have been disconnected, so we need to restore the ties.
|
||||||
|
if accelerator.distributed_type == DistributedType.TPU:
|
||||||
|
model.tie_weights()
|
||||||
|
|
||||||
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
|
# We need to recalculate our total training steps as the size of the training dataloader may have changed.
|
||||||
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
|
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
|
||||||
if overrode_max_train_steps:
|
if overrode_max_train_steps:
|
||||||
|
|||||||
Reference in New Issue
Block a user