Update train_ms.py

This commit is contained in:
Stardust·减
2023-08-23 16:42:11 +08:00
committed by GitHub
parent d09c7cbfeb
commit 991dfaa7fb

View File

@@ -98,7 +98,7 @@ def run(rank, n_gpus, hps):
if "use_duration_discriminator" in hps.model.keys() and hps.model.use_duration_discriminator == True: if "use_duration_discriminator" in hps.model.keys() and hps.model.use_duration_discriminator == True:
print("Using duration discriminator for VITS2") print("Using duration discriminator for VITS2")
use_duration_discriminator = True use_duration_discriminator = True
net_dur_disc = DurationDiscriminator( net_dur_disc = DurationDiscriminator(
hps.model.hidden_channels, hps.model.hidden_channels,
hps.model.hidden_channels, hps.model.hidden_channels,
3, 3,
@@ -265,7 +265,7 @@ def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loade
# Generator # Generator
y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat) y_d_hat_r, y_d_hat_g, fmap_r, fmap_g = net_d(y, y_hat)
if net_dur_disc is not None: if net_dur_disc is not None:
y_dur_hat_r, y_dur_hat_g = net_dur_disc(logw, logw_) y_dur_hat_r, y_dur_hat_g = net_dur_disc(x, x_mask, logw, logw_)
with autocast(enabled=False): with autocast(enabled=False):
loss_dur = torch.sum(l_length.float()) loss_dur = torch.sum(l_length.float())
loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel loss_mel = F.l1_loss(y_mel, y_hat_mel) * hps.train.c_mel
@@ -274,6 +274,9 @@ def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loade
loss_fm = feature_loss(fmap_r, fmap_g) loss_fm = feature_loss(fmap_r, fmap_g)
loss_gen, losses_gen = generator_loss(y_d_hat_g) loss_gen, losses_gen = generator_loss(y_d_hat_g)
loss_gen_all = loss_gen + loss_fm + loss_mel + loss_dur + loss_kl loss_gen_all = loss_gen + loss_fm + loss_mel + loss_dur + loss_kl
if net_dur_disc is not None:
loss_dur_gen, losses_dur_gen = generator_loss(y_dur_hat_g)
loss_gen_all += loss_dur_gen
optim_g.zero_grad() optim_g.zero_grad()
scaler.scale(loss_gen_all).backward() scaler.scale(loss_gen_all).backward()
scaler.unscale_(optim_g) scaler.unscale_(optim_g)
@@ -302,6 +305,7 @@ def train_and_evaluate(rank, epoch, hps, nets, optims, schedulers, scaler, loade
if net_dur_disc is not None: if net_dur_disc is not None:
scalar_dict.update({f"loss/dur_disc_r/{losses_dur_disc_r}"}) scalar_dict.update({f"loss/dur_disc_r/{losses_dur_disc_r}"})
scalar_dict.update({f"loss/dur_disc_g/{losses_dur_disc_g}"}) scalar_dict.update({f"loss/dur_disc_g/{losses_dur_disc_g}"})
scalar_dict.update({f"loss/dur_gen/{loss_dur_gen}"})
image_dict = { image_dict = {
"slice/mel_org": utils.plot_spectrogram_to_numpy(y_mel[0].data.cpu().numpy()), "slice/mel_org": utils.plot_spectrogram_to_numpy(y_mel[0].data.cpu().numpy()),
"slice/mel_gen": utils.plot_spectrogram_to_numpy(y_hat_mel[0].data.cpu().numpy()), "slice/mel_gen": utils.plot_spectrogram_to_numpy(y_hat_mel[0].data.cpu().numpy()),