Update finetune_on_pregenerated.py
This commit is contained in:
@@ -314,8 +314,8 @@ def main():
|
|||||||
mean_loss = tr_loss * args.gradient_accumulation_steps / nb_tr_steps
|
mean_loss = tr_loss * args.gradient_accumulation_steps / nb_tr_steps
|
||||||
pbar.set_postfix_str(f"Loss: {mean_loss:.5f}")
|
pbar.set_postfix_str(f"Loss: {mean_loss:.5f}")
|
||||||
if (step + 1) % args.gradient_accumulation_steps == 0:
|
if (step + 1) % args.gradient_accumulation_steps == 0:
|
||||||
scheduler.step() # Update learning rate schedule
|
|
||||||
optimizer.step()
|
optimizer.step()
|
||||||
|
scheduler.step() # Update learning rate schedule
|
||||||
optimizer.zero_grad()
|
optimizer.zero_grad()
|
||||||
global_step += 1
|
global_step += 1
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user