This commit is contained in:
Stardust·减
2023-09-05 09:54:11 +08:00
committed by GitHub
parent 02423d892f
commit bb4ebb80f7

View File

@@ -123,20 +123,20 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
assert bert.shape[-1] == len(phone), phone assert bert.shape[-1] == len(phone), phone
if language_str=='ZH': if language_str=='ZH':
zh_bert = bert bert = bert
ja_bert = torch.zeros(768, len(phone)) ja_bert = torch.zeros(768, len(phone))
elif language_str=="JA": elif language_str=="JA":
ja_bert = bert ja_bert = bert
zh_bert = torch.zeros(1024, len(phone)) bert = torch.zeros(1024, len(phone))
else: else:
zh_bert = torch.zeros(1024, len(phone)) bert = torch.zeros(1024, len(phone))
ja_bert = torch.zeros(768, len(phone)) ja_bert = torch.zeros(768, len(phone))
assert bert.shape[-1] == len(phone), ( assert bert.shape[-1] == len(phone), (
bert.shape, len(phone), sum(word2ph), p1, p2, t1, t2, pold, pold2, word2ph, text, w2pho) bert.shape, len(phone), sum(word2ph), p1, p2, t1, t2, pold, pold2, word2ph, text, w2pho)
phone = torch.LongTensor(phone) phone = torch.LongTensor(phone)
tone = torch.LongTensor(tone) tone = torch.LongTensor(tone)
language = torch.LongTensor(language) language = torch.LongTensor(language)
return bert, phone, tone, language return bert, ja_bert, phone, tone, language
def get_sid(self, sid): def get_sid(self, sid):
sid = torch.LongTensor([int(sid)]) sid = torch.LongTensor([int(sid)])
@@ -180,6 +180,7 @@ class TextAudioSpeakerCollate():
tone_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) language_padded = torch.LongTensor(len(batch), max_text_len)
bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
ja_bert_padded = torch.FloatTensor(len(batch), 768, max_text_len)
spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len) spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len)
wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len) wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)
@@ -189,6 +190,7 @@ class TextAudioSpeakerCollate():
spec_padded.zero_() spec_padded.zero_()
wav_padded.zero_() wav_padded.zero_()
bert_padded.zero_() bert_padded.zero_()
ja_bert_padded.zero_()
for i in range(len(ids_sorted_decreasing)): for i in range(len(ids_sorted_decreasing)):
row = batch[ids_sorted_decreasing[i]] row = batch[ids_sorted_decreasing[i]]
@@ -215,7 +217,10 @@ class TextAudioSpeakerCollate():
bert = row[6] bert = row[6]
bert_padded[i, :, :bert.size(1)] = bert bert_padded[i, :, :bert.size(1)] = bert
return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, sid, tone_padded, language_padded, bert_padded ja_bert = row[7]
ja_bert_padded[i, :, :ja_bert.size(1)] = ja_bert
return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, sid, tone_padded, language_padded, bert_padded, ja_bert_padded
class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler): class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):