From 99e1af4b587f148d7e217298717a9f5d777fa0f0 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 2 Feb 2024 21:42:34 +0900 Subject: [PATCH] Fix --- configs/configs_jp_extra.json | 1 + data_utils.py | 101 +++++++++++++++++++++------------- train_ms_jp_extra.py | 6 +- 3 files changed, 69 insertions(+), 39 deletions(-) diff --git a/configs/configs_jp_extra.json b/configs/configs_jp_extra.json index addb364..a11f68d 100644 --- a/configs/configs_jp_extra.json +++ b/configs/configs_jp_extra.json @@ -25,6 +25,7 @@ "freeze_style": false }, "data": { + "use_jp_extra": true, "training_files": "filelists/train.list", "validation_files": "filelists/val.list", "max_wav_value": 32768.0, diff --git a/data_utils.py b/data_utils.py index e20aca5..ac038c2 100644 --- a/data_utils.py +++ b/data_utils.py @@ -34,6 +34,7 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset): self.sampling_rate = hparams.sampling_rate self.spk_map = hparams.spk2id self.hparams = hparams + self.use_jp_extra = getattr(hparams, "use_jp_extra", False) self.use_mel_spec_posterior = getattr( hparams, "use_mel_posterior_encoder", False @@ -97,18 +98,21 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset): spec, wav = self.get_audio(audiopath) sid = torch.LongTensor([int(self.spk_map[sid])]) style_vec = torch.FloatTensor(np.load(f"{audiopath}.npy")) - return ( - phones, - spec, - wav, - sid, - tone, - language, - bert, - ja_bert, - en_bert, - style_vec, - ) + if self.use_jp_extra: + return (phones, spec, wav, sid, tone, language, ja_bert, style_vec) + else: + 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) @@ -200,8 +204,9 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset): class TextAudioSpeakerCollate: """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.use_jp_extra = use_jp_extra def __call__(self, batch): """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) tone_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) - ja_bert_padded = torch.FloatTensor(len(batch), 1024, max_text_len) - en_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) + 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) @@ -239,8 +246,9 @@ class TextAudioSpeakerCollate: spec_padded.zero_() wav_padded.zero_() bert_padded.zero_() - ja_bert_padded.zero_() - en_bert_padded.zero_() + if not self.use_jp_extra: + ja_bert_padded.zero_() + en_bert_padded.zero_() style_vec.zero_() for i in range(len(ids_sorted_decreasing)): @@ -269,29 +277,46 @@ class TextAudioSpeakerCollate: bert = row[6] bert_padded[i, :, : bert.size(1)] = bert - ja_bert = row[7] - ja_bert_padded[i, :, : ja_bert.size(1)] = ja_bert + if self.use_jp_extra: + 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_padded[i, :, : en_bert.size(1)] = en_bert + en_bert = row[8] + en_bert_padded[i, :, : en_bert.size(1)] = en_bert + style_vec[i, :] = row[9] - style_vec[i, :] = row[9] - - 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, - ) + 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 ( + 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): diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 8531d64..4ac102a 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -178,7 +178,7 @@ def run(): rank=rank, shuffle=True, ) - collate_fn = TextAudioSpeakerCollate() + collate_fn = TextAudioSpeakerCollate(use_jp_extra=True) train_loader = DataLoader( train_dataset, # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() - 1), @@ -394,6 +394,10 @@ def run(): _ = utils.load_safetensors( 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.") except Exception as e: logger.warning(e)