Update losses.py

This commit is contained in:
Sihan Wang
2023-12-20 02:50:58 +08:00
parent 97f18c0e26
commit 2e7c173de2
2 changed files with 12 additions and 9 deletions

View File

@@ -67,6 +67,8 @@ class WavLMLoss(torch.nn.Module):
self.wd = wd self.wd = wd
self.resample = torchaudio.transforms.Resample(model_sr, slm_sr) self.resample = torchaudio.transforms.Resample(model_sr, slm_sr)
self.wavlm.eval() self.wavlm.eval()
for param in self.wavlm.parameters():
param.requires_grad = False
def forward(self, wav, y_rec): def forward(self, wav, y_rec):
with torch.no_grad(): with torch.no_grad():

View File

@@ -349,6 +349,13 @@ def run():
scheduler_dur_disc = None scheduler_dur_disc = None
scaler = GradScaler(enabled=hps.train.bf16_run) scaler = GradScaler(enabled=hps.train.bf16_run)
wl = WavLMLoss(
hps.model.slm.model,
net_wd,
hps.data.sampling_rate,
hps.model.slm.sr,
).to(local_rank)
for epoch in range(epoch_str, hps.train.epochs + 1): for epoch in range(epoch_str, hps.train.epochs + 1):
if rank == 0: if rank == 0:
train_and_evaluate( train_and_evaluate(
@@ -356,7 +363,7 @@ def run():
local_rank, local_rank,
epoch, epoch,
hps, hps,
[net_g, net_d, net_dur_disc, net_wd], [net_g, net_d, net_dur_disc, net_wd, wl],
[optim_g, optim_d, optim_dur_disc, optim_wd], [optim_g, optim_d, optim_dur_disc, optim_wd],
[scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd], [scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd],
scaler, scaler,
@@ -370,7 +377,7 @@ def run():
local_rank, local_rank,
epoch, epoch,
hps, hps,
[net_g, net_d, net_dur_disc, net_wd], [net_g, net_d, net_dur_disc, net_wd, wl],
[optim_g, optim_d, optim_dur_disc, optim_wd], [optim_g, optim_d, optim_dur_disc, optim_wd],
[scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd], [scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd],
scaler, scaler,
@@ -397,18 +404,12 @@ def train_and_evaluate(
logger, logger,
writers, writers,
): ):
net_g, net_d, net_dur_disc, net_wd = nets net_g, net_d, net_dur_disc, net_wd, wl = nets
optim_g, optim_d, optim_dur_disc, optim_wd = optims optim_g, optim_d, optim_dur_disc, optim_wd = optims
scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd = schedulers scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd = schedulers
train_loader, eval_loader = loaders train_loader, eval_loader = loaders
if writers is not None: if writers is not None:
writer, writer_eval = writers writer, writer_eval = writers
wl = WavLMLoss(
hps.model.slm.model,
net_wd,
hps.data.sampling_rate,
hps.model.slm.sr,
).to(local_rank)
train_loader.batch_sampler.set_epoch(epoch) train_loader.batch_sampler.set_epoch(epoch)
global global_step global global_step