Fix
This commit is contained in:
101
data_utils.py
101
data_utils.py
@@ -34,6 +34,7 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
self.sampling_rate = hparams.sampling_rate
|
||||
self.spk_map = hparams.spk2id
|
||||
self.hparams = hparams
|
||||
self.use_jp_extra = getattr(hparams, "use_jp_extra", False)
|
||||
|
||||
self.use_mel_spec_posterior = getattr(
|
||||
hparams, "use_mel_posterior_encoder", False
|
||||
@@ -97,18 +98,21 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
spec, wav = self.get_audio(audiopath)
|
||||
sid = torch.LongTensor([int(self.spk_map[sid])])
|
||||
style_vec = torch.FloatTensor(np.load(f"{audiopath}.npy"))
|
||||
return (
|
||||
phones,
|
||||
spec,
|
||||
wav,
|
||||
sid,
|
||||
tone,
|
||||
language,
|
||||
bert,
|
||||
ja_bert,
|
||||
en_bert,
|
||||
style_vec,
|
||||
)
|
||||
if self.use_jp_extra:
|
||||
return (phones, spec, wav, sid, tone, language, ja_bert, style_vec)
|
||||
else:
|
||||
return (
|
||||
phones,
|
||||
spec,
|
||||
wav,
|
||||
sid,
|
||||
tone,
|
||||
language,
|
||||
bert,
|
||||
ja_bert,
|
||||
en_bert,
|
||||
style_vec,
|
||||
)
|
||||
|
||||
def get_audio(self, filename):
|
||||
audio, sampling_rate = load_wav_to_torch(filename)
|
||||
@@ -200,8 +204,9 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
class TextAudioSpeakerCollate:
|
||||
"""Zero-pads model inputs and targets"""
|
||||
|
||||
def __init__(self, return_ids=False):
|
||||
def __init__(self, return_ids=False, use_jp_extra=False):
|
||||
self.return_ids = return_ids
|
||||
self.use_jp_extra = use_jp_extra
|
||||
|
||||
def __call__(self, batch):
|
||||
"""Collate's training batch from normalized text, audio and speaker identities
|
||||
@@ -226,9 +231,11 @@ class TextAudioSpeakerCollate:
|
||||
text_padded = torch.LongTensor(len(batch), max_text_len)
|
||||
tone_padded = torch.LongTensor(len(batch), max_text_len)
|
||||
language_padded = torch.LongTensor(len(batch), max_text_len)
|
||||
# This is ZH bert if not use_jp_extra, JA bert if use_jp_extra
|
||||
bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
|
||||
ja_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
|
||||
en_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
|
||||
if not self.use_jp_extra:
|
||||
ja_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
|
||||
en_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
|
||||
style_vec = torch.FloatTensor(len(batch), 256)
|
||||
|
||||
spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len)
|
||||
@@ -239,8 +246,9 @@ class TextAudioSpeakerCollate:
|
||||
spec_padded.zero_()
|
||||
wav_padded.zero_()
|
||||
bert_padded.zero_()
|
||||
ja_bert_padded.zero_()
|
||||
en_bert_padded.zero_()
|
||||
if not self.use_jp_extra:
|
||||
ja_bert_padded.zero_()
|
||||
en_bert_padded.zero_()
|
||||
style_vec.zero_()
|
||||
|
||||
for i in range(len(ids_sorted_decreasing)):
|
||||
@@ -269,29 +277,46 @@ class TextAudioSpeakerCollate:
|
||||
bert = row[6]
|
||||
bert_padded[i, :, : bert.size(1)] = bert
|
||||
|
||||
ja_bert = row[7]
|
||||
ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert
|
||||
if self.use_jp_extra:
|
||||
style_vec[i, :] = row[7]
|
||||
else:
|
||||
ja_bert = row[7]
|
||||
ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert
|
||||
|
||||
en_bert = row[8]
|
||||
en_bert_padded[i, :, : en_bert.size(1)] = en_bert
|
||||
en_bert = row[8]
|
||||
en_bert_padded[i, :, : en_bert.size(1)] = en_bert
|
||||
style_vec[i, :] = row[9]
|
||||
|
||||
style_vec[i, :] = row[9]
|
||||
|
||||
return (
|
||||
text_padded,
|
||||
text_lengths,
|
||||
spec_padded,
|
||||
spec_lengths,
|
||||
wav_padded,
|
||||
wav_lengths,
|
||||
sid,
|
||||
tone_padded,
|
||||
language_padded,
|
||||
bert_padded,
|
||||
ja_bert_padded,
|
||||
en_bert_padded,
|
||||
style_vec,
|
||||
)
|
||||
if self.use_jp_extra:
|
||||
return (
|
||||
text_padded,
|
||||
text_lengths,
|
||||
spec_padded,
|
||||
spec_lengths,
|
||||
wav_padded,
|
||||
wav_lengths,
|
||||
sid,
|
||||
tone_padded,
|
||||
language_padded,
|
||||
bert_padded,
|
||||
style_vec,
|
||||
)
|
||||
else:
|
||||
return (
|
||||
text_padded,
|
||||
text_lengths,
|
||||
spec_padded,
|
||||
spec_lengths,
|
||||
wav_padded,
|
||||
wav_lengths,
|
||||
sid,
|
||||
tone_padded,
|
||||
language_padded,
|
||||
bert_padded,
|
||||
ja_bert_padded,
|
||||
en_bert_padded,
|
||||
style_vec,
|
||||
)
|
||||
|
||||
|
||||
class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||
|
||||
Reference in New Issue
Block a user