add bert
This commit is contained in:
@@ -123,20 +123,20 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
assert bert.shape[-1] == len(phone), phone
|
||||
|
||||
if language_str=='ZH':
|
||||
zh_bert = bert
|
||||
bert = bert
|
||||
ja_bert = torch.zeros(768, len(phone))
|
||||
elif language_str=="JA":
|
||||
ja_bert = bert
|
||||
zh_bert = torch.zeros(1024, len(phone))
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
else:
|
||||
zh_bert = torch.zeros(1024, len(phone))
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
ja_bert = torch.zeros(768, len(phone))
|
||||
assert bert.shape[-1] == len(phone), (
|
||||
bert.shape, len(phone), sum(word2ph), p1, p2, t1, t2, pold, pold2, word2ph, text, w2pho)
|
||||
phone = torch.LongTensor(phone)
|
||||
tone = torch.LongTensor(tone)
|
||||
language = torch.LongTensor(language)
|
||||
return bert, phone, tone, language
|
||||
return bert, ja_bert, phone, tone, language
|
||||
|
||||
def get_sid(self, sid):
|
||||
sid = torch.LongTensor([int(sid)])
|
||||
@@ -180,6 +180,7 @@ class TextAudioSpeakerCollate():
|
||||
tone_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)
|
||||
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)
|
||||
wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)
|
||||
@@ -189,6 +190,7 @@ class TextAudioSpeakerCollate():
|
||||
spec_padded.zero_()
|
||||
wav_padded.zero_()
|
||||
bert_padded.zero_()
|
||||
ja_bert_padded.zero_()
|
||||
for i in range(len(ids_sorted_decreasing)):
|
||||
row = batch[ids_sorted_decreasing[i]]
|
||||
|
||||
@@ -215,7 +217,10 @@ class TextAudioSpeakerCollate():
|
||||
bert = row[6]
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user