Format code
This commit is contained in:
26
train_ms.py
26
train_ms.py
@@ -173,20 +173,28 @@ def run():
|
|||||||
try:
|
try:
|
||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
_, _, dur_resume_lr, epoch_str = utils.load_checkpoint(
|
_, _, dur_resume_lr, epoch_str = utils.load_checkpoint(
|
||||||
utils.latest_checkpoint_path(hps.model_dir, "DUR_*.pth"), net_dur_disc, optim_dur_disc,
|
utils.latest_checkpoint_path(hps.model_dir, "DUR_*.pth"),
|
||||||
skip_optimizer=True)
|
net_dur_disc,
|
||||||
|
optim_dur_disc,
|
||||||
|
skip_optimizer=True,
|
||||||
|
)
|
||||||
_, optim_g, g_resume_lr, epoch_str = utils.load_checkpoint(
|
_, optim_g, g_resume_lr, epoch_str = utils.load_checkpoint(
|
||||||
utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"), net_g,
|
utils.latest_checkpoint_path(hps.model_dir, "G_*.pth"),
|
||||||
optim_g, skip_optimizer=True)
|
net_g,
|
||||||
|
optim_g,
|
||||||
|
skip_optimizer=True,
|
||||||
|
)
|
||||||
_, optim_d, d_resume_lr, epoch_str = utils.load_checkpoint(
|
_, optim_d, d_resume_lr, epoch_str = utils.load_checkpoint(
|
||||||
utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"), net_d,
|
utils.latest_checkpoint_path(hps.model_dir, "D_*.pth"),
|
||||||
optim_d, skip_optimizer=True)
|
net_d,
|
||||||
|
optim_d,
|
||||||
|
skip_optimizer=True,
|
||||||
|
)
|
||||||
if not optim_g.param_groups[0].get("initial_lr"):
|
if not optim_g.param_groups[0].get("initial_lr"):
|
||||||
optim_g.param_groups[0]["initial_lr"] = g_resume_lr
|
optim_g.param_groups[0]["initial_lr"] = g_resume_lr
|
||||||
if not optim_d.param_groups[0].get("initial_lr"):
|
if not optim_d.param_groups[0].get("initial_lr"):
|
||||||
optim_d.param_groups[0]["initial_lr"] = d_resume_lr
|
optim_d.param_groups[0]["initial_lr"] = d_resume_lr
|
||||||
|
|
||||||
|
|
||||||
epoch_str = max(epoch_str, 1)
|
epoch_str = max(epoch_str, 1)
|
||||||
global_step = (epoch_str - 1) * len(train_loader)
|
global_step = (epoch_str - 1) * len(train_loader)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -203,7 +211,9 @@ def run():
|
|||||||
if net_dur_disc is not None:
|
if net_dur_disc is not None:
|
||||||
if not optim_dur_disc.param_groups[0].get("initial_lr"):
|
if not optim_dur_disc.param_groups[0].get("initial_lr"):
|
||||||
optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr
|
optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr
|
||||||
scheduler_dur_disc = torch.optim.lr_scheduler.ExponentialLR(optim_dur_disc, gamma=hps.train.lr_decay, last_epoch=epoch_str-2)
|
scheduler_dur_disc = torch.optim.lr_scheduler.ExponentialLR(
|
||||||
|
optim_dur_disc, gamma=hps.train.lr_decay, last_epoch=epoch_str - 2
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
scheduler_dur_disc = None
|
scheduler_dur_disc = None
|
||||||
scaler = GradScaler(enabled=hps.train.fp16_run)
|
scaler = GradScaler(enabled=hps.train.fp16_run)
|
||||||
|
|||||||
Reference in New Issue
Block a user