Update losses.py
This commit is contained in:
@@ -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():
|
||||
|
||||
19
train_ms.py
19
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
|
||||
|
||||
Reference in New Issue
Block a user