From 2e7c173de27456dc7288602c74ef96639bca8d38 Mon Sep 17 00:00:00 2001 From: Sihan Wang Date: Wed, 20 Dec 2023 02:50:58 +0800 Subject: [PATCH] Update losses.py --- losses.py | 2 ++ train_ms.py | 19 ++++++++++--------- 2 files changed, 12 insertions(+), 9 deletions(-) diff --git a/losses.py b/losses.py index 6a4dafa..62982fc 100644 --- a/losses.py +++ b/losses.py @@ -67,6 +67,8 @@ class WavLMLoss(torch.nn.Module): self.wd = wd self.resample = torchaudio.transforms.Resample(model_sr, slm_sr) self.wavlm.eval() + for param in self.wavlm.parameters(): + param.requires_grad = False def forward(self, wav, y_rec): with torch.no_grad(): diff --git a/train_ms.py b/train_ms.py index fe795c0..f76534d 100644 --- a/train_ms.py +++ b/train_ms.py @@ -349,6 +349,13 @@ def run(): scheduler_dur_disc = None 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): if rank == 0: train_and_evaluate( @@ -356,7 +363,7 @@ def run(): local_rank, epoch, 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], [scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd], scaler, @@ -370,7 +377,7 @@ def run(): local_rank, epoch, 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], [scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd], scaler, @@ -397,18 +404,12 @@ def train_and_evaluate( logger, 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 scheduler_g, scheduler_d, scheduler_dur_disc, scheduler_wd = schedulers train_loader, eval_loader = loaders if writers is not None: 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) global global_step