Merge pull request #19 from kale4eat/dev-warmup

Implement warm-up
This commit is contained in:
litagin02
2024-01-08 21:53:48 +09:00
committed by GitHub

View File

@@ -356,15 +356,34 @@ def run():
epoch_str = 1 epoch_str = 1
global_step = 0 global_step = 0
scheduler_g = torch.optim.lr_scheduler.ExponentialLR( steps_per_epoch = len(train_loader)
optim_g, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 warmup_steps = hps.train.warmup_epochs * steps_per_epoch
def lr_lambda(step):
"""for warmup"""
if step < warmup_steps:
return step / warmup_steps
else:
return hps.train.lr_decay ** (step - warmup_steps)
scheduler_last_epoch = -1
if global_step > 0:
scheduler_last_epoch = global_step - 1
scheduler_g = torch.optim.lr_scheduler.LambdaLR(
optim_g, lr_lambda=lr_lambda, last_epoch=scheduler_last_epoch
) )
scheduler_d = torch.optim.lr_scheduler.ExponentialLR( scheduler_d = torch.optim.lr_scheduler.LambdaLR(
optim_d, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 optim_d, lr_lambda=lr_lambda, last_epoch=scheduler_last_epoch
) )
if net_dur_disc is not None: if net_dur_disc is not None:
scheduler_dur_disc = torch.optim.lr_scheduler.ExponentialLR( if not optim_dur_disc.param_groups[0].get("initial_lr"):
optim_dur_disc, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2 if global_step == 0:
optim_dur_disc.param_groups[0]["initial_lr"] = hps.train.learning_rate
else:
optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr
initial_lr = optim_dur_disc.param_groups[0].get("initial_lr")
scheduler_dur_disc = torch.optim.lr_scheduler.LambdaLR(
optim_dur_disc, lr_lambda=lr_lambda, last_epoch=scheduler_last_epoch
) )
else: else:
scheduler_dur_disc = None scheduler_dur_disc = None
@@ -417,10 +436,6 @@ def run():
pbar, pbar,
initial_step, initial_step,
) )
scheduler_g.step()
scheduler_d.step()
if net_dur_disc is not None:
scheduler_dur_disc.step()
if epoch == hps.train.epochs: if epoch == hps.train.epochs:
# Save the final models # Save the final models
@@ -737,6 +752,10 @@ def train_and_evaluate(
for_infer=True, for_infer=True,
) )
scheduler_g.step()
scheduler_d.step()
if net_dur_disc is not None:
scheduler_dur_disc.step()
global_step += 1 global_step += 1
if pbar is not None: if pbar is not None:
pbar.set_description( pbar.set_description(