Fix
This commit is contained in:
@@ -25,6 +25,7 @@
|
|||||||
"freeze_style": false
|
"freeze_style": false
|
||||||
},
|
},
|
||||||
"data": {
|
"data": {
|
||||||
|
"use_jp_extra": true,
|
||||||
"training_files": "filelists/train.list",
|
"training_files": "filelists/train.list",
|
||||||
"validation_files": "filelists/val.list",
|
"validation_files": "filelists/val.list",
|
||||||
"max_wav_value": 32768.0,
|
"max_wav_value": 32768.0,
|
||||||
|
|||||||
101
data_utils.py
101
data_utils.py
@@ -34,6 +34,7 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
|||||||
self.sampling_rate = hparams.sampling_rate
|
self.sampling_rate = hparams.sampling_rate
|
||||||
self.spk_map = hparams.spk2id
|
self.spk_map = hparams.spk2id
|
||||||
self.hparams = hparams
|
self.hparams = hparams
|
||||||
|
self.use_jp_extra = getattr(hparams, "use_jp_extra", False)
|
||||||
|
|
||||||
self.use_mel_spec_posterior = getattr(
|
self.use_mel_spec_posterior = getattr(
|
||||||
hparams, "use_mel_posterior_encoder", False
|
hparams, "use_mel_posterior_encoder", False
|
||||||
@@ -97,18 +98,21 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
|||||||
spec, wav = self.get_audio(audiopath)
|
spec, wav = self.get_audio(audiopath)
|
||||||
sid = torch.LongTensor([int(self.spk_map[sid])])
|
sid = torch.LongTensor([int(self.spk_map[sid])])
|
||||||
style_vec = torch.FloatTensor(np.load(f"{audiopath}.npy"))
|
style_vec = torch.FloatTensor(np.load(f"{audiopath}.npy"))
|
||||||
return (
|
if self.use_jp_extra:
|
||||||
phones,
|
return (phones, spec, wav, sid, tone, language, ja_bert, style_vec)
|
||||||
spec,
|
else:
|
||||||
wav,
|
return (
|
||||||
sid,
|
phones,
|
||||||
tone,
|
spec,
|
||||||
language,
|
wav,
|
||||||
bert,
|
sid,
|
||||||
ja_bert,
|
tone,
|
||||||
en_bert,
|
language,
|
||||||
style_vec,
|
bert,
|
||||||
)
|
ja_bert,
|
||||||
|
en_bert,
|
||||||
|
style_vec,
|
||||||
|
)
|
||||||
|
|
||||||
def get_audio(self, filename):
|
def get_audio(self, filename):
|
||||||
audio, sampling_rate = load_wav_to_torch(filename)
|
audio, sampling_rate = load_wav_to_torch(filename)
|
||||||
@@ -200,8 +204,9 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset):
|
|||||||
class TextAudioSpeakerCollate:
|
class TextAudioSpeakerCollate:
|
||||||
"""Zero-pads model inputs and targets"""
|
"""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.return_ids = return_ids
|
||||||
|
self.use_jp_extra = use_jp_extra
|
||||||
|
|
||||||
def __call__(self, batch):
|
def __call__(self, batch):
|
||||||
"""Collate's training batch from normalized text, audio and speaker identities
|
"""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)
|
text_padded = torch.LongTensor(len(batch), max_text_len)
|
||||||
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)
|
||||||
|
# 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)
|
bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
|
||||||
ja_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len)
|
if not self.use_jp_extra:
|
||||||
en_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)
|
style_vec = torch.FloatTensor(len(batch), 256)
|
||||||
|
|
||||||
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)
|
||||||
@@ -239,8 +246,9 @@ class TextAudioSpeakerCollate:
|
|||||||
spec_padded.zero_()
|
spec_padded.zero_()
|
||||||
wav_padded.zero_()
|
wav_padded.zero_()
|
||||||
bert_padded.zero_()
|
bert_padded.zero_()
|
||||||
ja_bert_padded.zero_()
|
if not self.use_jp_extra:
|
||||||
en_bert_padded.zero_()
|
ja_bert_padded.zero_()
|
||||||
|
en_bert_padded.zero_()
|
||||||
style_vec.zero_()
|
style_vec.zero_()
|
||||||
|
|
||||||
for i in range(len(ids_sorted_decreasing)):
|
for i in range(len(ids_sorted_decreasing)):
|
||||||
@@ -269,29 +277,46 @@ class TextAudioSpeakerCollate:
|
|||||||
bert = row[6]
|
bert = row[6]
|
||||||
bert_padded[i, :, : bert.size(1)] = bert
|
bert_padded[i, :, : bert.size(1)] = bert
|
||||||
|
|
||||||
ja_bert = row[7]
|
if self.use_jp_extra:
|
||||||
ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert
|
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 = row[8]
|
||||||
en_bert_padded[i, :, : en_bert.size(1)] = en_bert
|
en_bert_padded[i, :, : en_bert.size(1)] = en_bert
|
||||||
|
style_vec[i, :] = row[9]
|
||||||
|
|
||||||
style_vec[i, :] = row[9]
|
if self.use_jp_extra:
|
||||||
|
return (
|
||||||
return (
|
text_padded,
|
||||||
text_padded,
|
text_lengths,
|
||||||
text_lengths,
|
spec_padded,
|
||||||
spec_padded,
|
spec_lengths,
|
||||||
spec_lengths,
|
wav_padded,
|
||||||
wav_padded,
|
wav_lengths,
|
||||||
wav_lengths,
|
sid,
|
||||||
sid,
|
tone_padded,
|
||||||
tone_padded,
|
language_padded,
|
||||||
language_padded,
|
bert_padded,
|
||||||
bert_padded,
|
style_vec,
|
||||||
ja_bert_padded,
|
)
|
||||||
en_bert_padded,
|
else:
|
||||||
style_vec,
|
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):
|
class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):
|
||||||
|
|||||||
@@ -178,7 +178,7 @@ def run():
|
|||||||
rank=rank,
|
rank=rank,
|
||||||
shuffle=True,
|
shuffle=True,
|
||||||
)
|
)
|
||||||
collate_fn = TextAudioSpeakerCollate()
|
collate_fn = TextAudioSpeakerCollate(use_jp_extra=True)
|
||||||
train_loader = DataLoader(
|
train_loader = DataLoader(
|
||||||
train_dataset,
|
train_dataset,
|
||||||
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() - 1),
|
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() - 1),
|
||||||
@@ -394,6 +394,10 @@ def run():
|
|||||||
_ = utils.load_safetensors(
|
_ = utils.load_safetensors(
|
||||||
os.path.join(model_dir, "DUR_0.safetensors"), net_dur_disc
|
os.path.join(model_dir, "DUR_0.safetensors"), net_dur_disc
|
||||||
)
|
)
|
||||||
|
if net_wd is not None:
|
||||||
|
_ = utils.load_safetensors(
|
||||||
|
os.path.join(model_dir, "WD_0.safetensors"), net_wd
|
||||||
|
)
|
||||||
logger.info("Loaded the pretrained models.")
|
logger.info("Loaded the pretrained models.")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(e)
|
logger.warning(e)
|
||||||
|
|||||||
Reference in New Issue
Block a user