Refactor: separate module for utilities related to loading/saving checkpoints and safetensors

This commit is contained in:
tsukumi
2024-03-09 15:58:57 +00:00
parent c594f7ea7a
commit 98ab8e7978
7 changed files with 510 additions and 457 deletions

View File

@@ -247,10 +247,7 @@ def run():
drop_last=False,
collate_fn=collate_fn,
)
if (
"use_noise_scaled_mas" in hps.model.keys()
and hps.model.use_noise_scaled_mas is True
):
if hps.model.use_noise_scaled_mas is True:
logger.info("Using noise scaled MAS for VITS2")
mas_noise_scale_initial = 0.01
noise_scale_delta = 2e-6
@@ -258,10 +255,7 @@ def run():
logger.info("Using normal MAS for VITS1")
mas_noise_scale_initial = 0.0
noise_scale_delta = 0.0
if (
"use_duration_discriminator" in hps.model.keys()
and hps.model.use_duration_discriminator is True
):
if hps.model.use_duration_discriminator is True:
logger.info("Using duration discriminator for VITS2")
net_dur_disc = DurationDiscriminator(
hps.model.hidden_channels,
@@ -272,19 +266,13 @@ def run():
).cuda(local_rank)
else:
net_dur_disc = None
if (
"use_wavlm_discriminator" in hps.model.keys()
and hps.model.use_wavlm_discriminator is True
):
if hps.model.use_wavlm_discriminator is True:
net_wd = WavLMDiscriminator(
hps.model.slm.hidden, hps.model.slm.nlayers, hps.model.slm.initial_channel
).cuda(local_rank)
else:
net_wd = None
if (
"use_spk_conditioned_encoder" in hps.model.keys()
and hps.model.use_spk_conditioned_encoder is True
):
if hps.model.use_spk_conditioned_encoder is True:
if hps.data.n_speakers == 0:
raise ValueError(
"n_speakers must be > 0 when using spk conditioned encoder to train multi-speaker model"
@@ -394,15 +382,11 @@ def run():
if utils.is_resuming(model_dir):
if net_dur_disc is not None:
try:
_, _, dur_resume_lr, epoch_str = utils.load_checkpoint(
utils.latest_checkpoint_path(model_dir, "DUR_*.pth"),
_, _, dur_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
utils.checkpoints.get_latest_checkpoint_path(model_dir, "DUR_*.pth"),
net_dur_disc,
optim_dur_disc,
skip_optimizer=(
hps.train.skip_optimizer
if "skip_optimizer" in hps.train
else True
),
skip_optimizer=hps.train.skip_optimizer,
)
if not optim_dur_disc.param_groups[0].get("initial_lr"):
optim_dur_disc.param_groups[0]["initial_lr"] = dur_resume_lr
@@ -412,15 +396,11 @@ def run():
print("Initialize dur_disc")
if net_wd is not None:
try:
_, optim_wd, wd_resume_lr, epoch_str = utils.load_checkpoint(
utils.latest_checkpoint_path(model_dir, "WD_*.pth"),
_, optim_wd, wd_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
utils.checkpoints.get_latest_checkpoint_path(model_dir, "WD_*.pth"),
net_wd,
optim_wd,
skip_optimizer=(
hps.train.skip_optimizer
if "skip_optimizer" in hps.train
else True
),
skip_optimizer=hps.train.skip_optimizer,
)
if not optim_wd.param_groups[0].get("initial_lr"):
optim_wd.param_groups[0]["initial_lr"] = wd_resume_lr
@@ -430,21 +410,17 @@ def run():
logger.info("Initialize wavlm")
try:
_, optim_g, g_resume_lr, epoch_str = utils.load_checkpoint(
utils.latest_checkpoint_path(model_dir, "G_*.pth"),
_, optim_g, g_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth"),
net_g,
optim_g,
skip_optimizer=(
hps.train.skip_optimizer if "skip_optimizer" in hps.train else True
),
skip_optimizer=hps.train.skip_optimizer,
)
_, optim_d, d_resume_lr, epoch_str = utils.load_checkpoint(
utils.latest_checkpoint_path(model_dir, "D_*.pth"),
_, optim_d, d_resume_lr, epoch_str = utils.checkpoints.load_checkpoint(
utils.checkpoints.get_latest_checkpoint_path(model_dir, "D_*.pth"),
net_d,
optim_d,
skip_optimizer=(
hps.train.skip_optimizer if "skip_optimizer" in hps.train else True
),
skip_optimizer=hps.train.skip_optimizer,
)
if not optim_g.param_groups[0].get("initial_lr"):
optim_g.param_groups[0]["initial_lr"] = g_resume_lr
@@ -454,7 +430,7 @@ def run():
epoch_str = max(epoch_str, 1)
# global_step = (epoch_str - 1) * len(train_loader)
global_step = int(
utils.get_steps(utils.latest_checkpoint_path(model_dir, "G_*.pth"))
utils.get_steps(utils.checkpoints.get_latest_checkpoint_path(model_dir, "G_*.pth"))
)
logger.info(
f"******************Found the model. Current epoch is {epoch_str}, gloabl step is {global_step}*********************"
@@ -468,18 +444,18 @@ def run():
global_step = 0
else:
try:
_ = utils.load_safetensors(
_ = utils.safetensors.load_safetensors(
os.path.join(model_dir, "G_0.safetensors"), net_g
)
_ = utils.load_safetensors(
_ = utils.safetensors.load_safetensors(
os.path.join(model_dir, "D_0.safetensors"), net_d
)
if net_dur_disc is not None:
_ = utils.load_safetensors(
_ = utils.safetensors.load_safetensors(
os.path.join(model_dir, "DUR_0.safetensors"), net_dur_disc
)
if net_wd is not None:
_ = utils.load_safetensors(
_ = utils.safetensors.load_safetensors(
os.path.join(model_dir, "WD_0.safetensors"), net_wd
)
logger.info("Loaded the pretrained models.")
@@ -586,14 +562,16 @@ def run():
scheduler_wd.step()
if epoch == hps.train.epochs:
# Save the final models
utils.save_checkpoint(
assert optim_g is not None
utils.checkpoints.save_checkpoint(
net_g,
optim_g,
hps.train.learning_rate,
epoch,
os.path.join(model_dir, "G_{}.pth".format(global_step)),
)
utils.save_checkpoint(
assert optim_d is not None
utils.checkpoints.save_checkpoint(
net_d,
optim_d,
hps.train.learning_rate,
@@ -601,7 +579,8 @@ def run():
os.path.join(model_dir, "D_{}.pth".format(global_step)),
)
if net_dur_disc is not None:
utils.save_checkpoint(
assert optim_dur_disc is not None
utils.checkpoints.save_checkpoint(
net_dur_disc,
optim_dur_disc,
hps.train.learning_rate,
@@ -609,14 +588,15 @@ def run():
os.path.join(model_dir, "DUR_{}.pth".format(global_step)),
)
if net_wd is not None:
utils.save_checkpoint(
assert optim_wd is not None
utils.checkpoints.save_checkpoint(
net_wd,
optim_wd,
hps.train.learning_rate,
epoch,
os.path.join(model_dir, "WD_{}.pth".format(global_step)),
)
utils.save_safetensors(
utils.safetensors.save_safetensors(
net_g,
epoch,
os.path.join(
@@ -949,14 +929,14 @@ def train_and_evaluate(
):
if not hps.speedup:
evaluate(hps, net_g, eval_loader, writer_eval)
utils.save_checkpoint(
utils.checkpoints.save_checkpoint(
net_g,
optim_g,
hps.train.learning_rate,
epoch,
os.path.join(hps.model_dir, "G_{}.pth".format(global_step)),
)
utils.save_checkpoint(
utils.checkpoints.save_checkpoint(
net_d,
optim_d,
hps.train.learning_rate,
@@ -964,7 +944,7 @@ def train_and_evaluate(
os.path.join(hps.model_dir, "D_{}.pth".format(global_step)),
)
if net_dur_disc is not None:
utils.save_checkpoint(
utils.checkpoints.save_checkpoint(
net_dur_disc,
optim_dur_disc,
hps.train.learning_rate,
@@ -972,7 +952,7 @@ def train_and_evaluate(
os.path.join(hps.model_dir, "DUR_{}.pth".format(global_step)),
)
if net_wd is not None:
utils.save_checkpoint(
utils.checkpoints.save_checkpoint(
net_wd,
optim_wd,
hps.train.learning_rate,
@@ -981,13 +961,13 @@ def train_and_evaluate(
)
keep_ckpts = config.train_ms_config.keep_ckpts
if keep_ckpts > 0:
utils.clean_checkpoints(
path_to_models=hps.model_dir,
utils.checkpoints.clean_checkpoints(
model_dir_path=hps.model_dir,
n_ckpts_to_keep=keep_ckpts,
sort_by_time=True,
)
# Save safetensors (for inference) to `model_assets/{model_name}`
utils.save_safetensors(
utils.safetensors.save_safetensors(
net_g,
epoch,
os.path.join(