init (not checked bat script yet)
This commit is contained in:
@@ -3,6 +3,7 @@ import random
|
||||
import torch
|
||||
import torch.utils.data
|
||||
from tqdm import tqdm
|
||||
import numpy as np
|
||||
from tools.log import logger
|
||||
import commons
|
||||
from mel_processing import spectrogram_torch, mel_spectrogram_torch
|
||||
@@ -92,8 +93,19 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
|
||||
spec, wav = self.get_audio(audiopath)
|
||||
sid = torch.LongTensor([int(self.spk_map[sid])])
|
||||
|
||||
return (phones, spec, wav, sid, tone, language, bert, ja_bert, en_bert)
|
||||
style_vec = torch.FloatTensor(np.load(f"{audiopath}.npy"))
|
||||
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)
|
||||
@@ -156,15 +168,15 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
||||
|
||||
if language_str == "ZH":
|
||||
bert = bert_ori
|
||||
ja_bert = torch.randn(1024, len(phone))
|
||||
en_bert = torch.randn(1024, len(phone))
|
||||
ja_bert = torch.zeros(1024, len(phone))
|
||||
en_bert = torch.zeros(1024, len(phone))
|
||||
elif language_str == "JP":
|
||||
bert = torch.randn(1024, len(phone))
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
ja_bert = bert_ori
|
||||
en_bert = torch.randn(1024, len(phone))
|
||||
en_bert = torch.zeros(1024, len(phone))
|
||||
elif language_str == "EN":
|
||||
bert = torch.randn(1024, len(phone))
|
||||
ja_bert = torch.randn(1024, len(phone))
|
||||
bert = torch.zeros(1024, len(phone))
|
||||
ja_bert = torch.zeros(1024, len(phone))
|
||||
en_bert = bert_ori
|
||||
phone = torch.LongTensor(phone)
|
||||
tone = torch.LongTensor(tone)
|
||||
@@ -214,6 +226,7 @@ class TextAudioSpeakerCollate:
|
||||
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)
|
||||
style_vec = torch.FloatTensor(len(batch), 256)
|
||||
|
||||
spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len)
|
||||
wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)
|
||||
@@ -225,6 +238,7 @@ class TextAudioSpeakerCollate:
|
||||
bert_padded.zero_()
|
||||
ja_bert_padded.zero_()
|
||||
en_bert_padded.zero_()
|
||||
style_vec.zero_()
|
||||
|
||||
for i in range(len(ids_sorted_decreasing)):
|
||||
row = batch[ids_sorted_decreasing[i]]
|
||||
@@ -258,6 +272,8 @@ class TextAudioSpeakerCollate:
|
||||
en_bert = row[8]
|
||||
en_bert_padded[i, :, : en_bert.size(1)] = en_bert
|
||||
|
||||
style_vec[i, :] = row[9]
|
||||
|
||||
return (
|
||||
text_padded,
|
||||
text_lengths,
|
||||
@@ -271,6 +287,7 @@ class TextAudioSpeakerCollate:
|
||||
bert_padded,
|
||||
ja_bert_padded,
|
||||
en_bert_padded,
|
||||
style_vec,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user