This commit is contained in:
litagin02
2024-02-02 21:42:34 +09:00
parent e566ac62f3
commit 99e1af4b58
3 changed files with 69 additions and 39 deletions

View File

@@ -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):