add bert
This commit is contained in:
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user