This commit is contained in:
litagin02
2024-02-02 21:42:34 +09:00
parent e566ac62f3
commit 99e1af4b58
3 changed files with 69 additions and 39 deletions

View File

@@ -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,

View File

@@ -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,6 +98,9 @@ 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"))
if self.use_jp_extra:
return (phones, spec, wav, sid, tone, language, ja_bert, style_vec)
else:
return ( return (
phones, phones,
spec, spec,
@@ -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,7 +231,9 @@ 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)
if not self.use_jp_extra:
ja_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) 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)
@@ -239,6 +246,7 @@ class TextAudioSpeakerCollate:
spec_padded.zero_() spec_padded.zero_()
wav_padded.zero_() wav_padded.zero_()
bert_padded.zero_() bert_padded.zero_()
if not self.use_jp_extra:
ja_bert_padded.zero_() ja_bert_padded.zero_()
en_bert_padded.zero_() en_bert_padded.zero_()
style_vec.zero_() style_vec.zero_()
@@ -269,14 +277,31 @@ class TextAudioSpeakerCollate:
bert = row[6] bert = row[6]
bert_padded[i, :, : bert.size(1)] = bert bert_padded[i, :, : bert.size(1)] = bert
if self.use_jp_extra:
style_vec[i, :] = row[7]
else:
ja_bert = row[7] ja_bert = row[7]
ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert 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 (
text_padded,
text_lengths,
spec_padded,
spec_lengths,
wav_padded,
wav_lengths,
sid,
tone_padded,
language_padded,
bert_padded,
style_vec,
)
else:
return ( return (
text_padded, text_padded,
text_lengths, text_lengths,

View File

@@ -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)