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