mirror of
https://github.com/wassname/TTS.git
synced 2026-09-10 11:50:20 +08:00
masked loss
This commit is contained in:
@@ -26,6 +26,7 @@ from utils.model import get_param_size
|
||||
from utils.visual import plot_alignment, plot_spectrogram
|
||||
from datasets.LJSpeech import LJSpeechDataset
|
||||
from models.tacotron import Tacotron
|
||||
from losses import
|
||||
|
||||
|
||||
use_cuda = torch.cuda.is_available()
|
||||
@@ -80,6 +81,7 @@ def train(model, criterion, data_loader, optimizer, epoch):
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
|
||||
current_step = num_iter + args.restore_step + epoch * len(data_loader) + 1
|
||||
|
||||
@@ -93,6 +95,7 @@ def train(model, criterion, data_loader, optimizer, epoch):
|
||||
# convert inputs to variables
|
||||
text_input_var = Variable(text_input)
|
||||
mel_spec_var = Variable(mel_input)
|
||||
mel_length_var = Variable(mel_lengths)
|
||||
linear_spec_var = Variable(linear_input, volatile=True)
|
||||
|
||||
# sort sequence by length for curriculum learning
|
||||
@@ -108,6 +111,7 @@ def train(model, criterion, data_loader, optimizer, epoch):
|
||||
if use_cuda:
|
||||
text_input_var = text_input_var.cuda()
|
||||
mel_spec_var = mel_spec_var.cuda()
|
||||
mel_lengths_var = mel_lengths_var.cuda()
|
||||
linear_spec_var = linear_spec_var.cuda()
|
||||
|
||||
# forward pass
|
||||
@@ -115,10 +119,11 @@ def train(model, criterion, data_loader, optimizer, epoch):
|
||||
model.forward(text_input_var, mel_spec_var)
|
||||
|
||||
# loss computation
|
||||
mel_loss = criterion(mel_output, mel_spec_var)
|
||||
linear_loss = 0.5 * criterion(linear_output, linear_spec_var) \
|
||||
mel_loss = criterion(mel_output, mel_spec_var, mel_lengths)
|
||||
linear_loss = 0.5 * criterion(linear_output, linear_spec_var, mel_lengths) \
|
||||
+ 0.5 * criterion(linear_output[:, :, :n_priority_freq],
|
||||
linear_spec_var[: ,: ,:n_priority_freq])
|
||||
linear_spec_var[: ,: ,:n_priority_freq],
|
||||
mel_lengths)
|
||||
loss = mel_loss + linear_loss
|
||||
|
||||
# backpass and check the grad norm
|
||||
@@ -215,26 +220,30 @@ def evaluate(model, criterion, data_loader, current_step):
|
||||
text_lengths = data[1]
|
||||
linear_input = data[2]
|
||||
mel_input = data[3]
|
||||
mel_lengths = data[4]
|
||||
|
||||
# convert inputs to variables
|
||||
text_input_var = Variable(text_input)
|
||||
mel_spec_var = Variable(mel_input)
|
||||
mel_lengths_var = Variable(mel_lengths)
|
||||
linear_spec_var = Variable(linear_input, volatile=True)
|
||||
|
||||
# dispatch data to GPU
|
||||
if use_cuda:
|
||||
text_input_var = text_input_var.cuda()
|
||||
mel_spec_var = mel_spec_var.cuda()
|
||||
mel_lengths_var = mel_lengths_var.cuda()
|
||||
linear_spec_var = linear_spec_var.cuda()
|
||||
|
||||
# forward pass
|
||||
mel_output, linear_output, alignments = model.forward(text_input_var, mel_spec_var)
|
||||
|
||||
# loss computation
|
||||
mel_loss = criterion(mel_output, mel_spec_var)
|
||||
linear_loss = 0.5 * criterion(linear_output, linear_spec_var) \
|
||||
mel_loss = criterion(mel_output, mel_spec_var, mel_lengths)
|
||||
linear_loss = 0.5 * criterion(linear_output, linear_spec_var, mel_lengths) \
|
||||
+ 0.5 * criterion(linear_output[:, :, :n_priority_freq],
|
||||
linear_spec_var[: ,: ,:n_priority_freq])
|
||||
linear_spec_var[: ,: ,:n_priority_freq],
|
||||
mel_lengths)
|
||||
loss = mel_loss + linear_loss
|
||||
|
||||
step_time = time.time() - start_time
|
||||
|
||||
Reference in New Issue
Block a user