diff --git a/app.py b/app.py index 714ebec..d4d08ff 100644 --- a/app.py +++ b/app.py @@ -46,9 +46,13 @@ def tts_fn( use_assist_text, style, style_weight, + given_tone, ): assert model_holder.current_model is not None - + if given_tone == "": + given_tone = None + else: + given_tone = [int(i) for i in given_tone] start_time = datetime.datetime.now() sr, audio = model_holder.current_model.infer( @@ -66,6 +70,7 @@ def tts_fn( use_assist_text=use_assist_text, style=style, style_weight=style_weight, + given_tone=given_tone, ) end_time = datetime.datetime.now() @@ -238,6 +243,7 @@ if __name__ == "__main__": step=0.1, label="分けた場合に挟む無音の長さ(秒)", ) + given_tone = gr.Textbox("トーン、0と1の数値列") language = gr.Dropdown(choices=languages, value="JP", label="Language") with gr.Accordion(label="詳細設定", open=False): sdp_ratio = gr.Slider( @@ -334,6 +340,7 @@ if __name__ == "__main__": use_assist_text, style, style_weight, + given_tone, ], outputs=[text_output, audio_output], ) diff --git a/common/tts_model.py b/common/tts_model.py index 0b5c67d..5fdf551 100644 --- a/common/tts_model.py +++ b/common/tts_model.py @@ -95,6 +95,7 @@ class Model: use_assist_text: bool = False, style: str = DEFAULT_STYLE, style_weight: float = DEFAULT_STYLE_WEIGHT, + given_tone: Optional[list[int]] = None, ) -> tuple[int, np.ndarray]: logger.info(f"Start generating audio data from text:\n{text}") if reference_audio_path == "": @@ -127,6 +128,7 @@ class Model: assist_text=assist_text, assist_text_weight=assist_text_weight, style_vec=style_vector, + given_tone=given_tone, ) else: texts = text.split("\n") diff --git a/configs/config.json b/configs/config.json index 640577a..07b9822 100644 --- a/configs/config.json +++ b/configs/config.json @@ -17,9 +17,10 @@ "c_mel": 45, "c_kl": 1.0, "skip_optimizer": false, - "freeze_ZH_bert": false, - "freeze_JP_bert": false, - "freeze_EN_bert": false + "freeze_ZH_bert": true, + "freeze_JP_bert": true, + "freeze_EN_bert": true, + "freeze_style": true }, "data": { "training_files": "Data/your_model_name/filelists/train.list", diff --git a/infer.py b/infer.py index 31ff6c3..fb0e670 100644 --- a/infer.py +++ b/infer.py @@ -6,6 +6,7 @@ from models import SynthesizerTrn from text import cleaned_text_to_sequence, get_bert from text.cleaner import clean_text from text.symbols import symbols +from common.log import logger # latest_version = "1.0" @@ -29,9 +30,21 @@ def get_net_g(model_path: str, version: str, device: str, hps): return net_g -def get_text(text, language_str, hps, device, assist_text=None, assist_text_weight=0.7): +def get_text( + text, + language_str, + hps, + device, + assist_text=None, + assist_text_weight=0.7, + given_tone=None, +): # 在此处实现当前版本的get_text norm_text, phone, tone, word2ph = clean_text(text, language_str) + logger.info(f"Original tone: {''.join(str(num) for num in tone)}") + if given_tone is not None: + logger.debug(f"Tone given: {given_tone}") + tone = given_tone phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str) if hps.data.add_blank: @@ -88,6 +101,7 @@ def infer( skip_end=False, assist_text=None, assist_text_weight=0.7, + given_tone=None, ): bert, ja_bert, en_bert, phones, tones, lang_ids = get_text( text, @@ -96,6 +110,7 @@ def infer( device, assist_text=assist_text, assist_text_weight=assist_text_weight, + given_tone=given_tone, ) if skip_start: phones = phones[3:] diff --git a/text/japanese.py b/text/japanese.py index 5bc590d..717b169 100644 --- a/text/japanese.py +++ b/text/japanese.py @@ -202,6 +202,7 @@ def g2phone_tone_list(text: str) -> list[tuple[str, int]]: [('k', 0), ('o', 0), ('n', 1), ('n', 1), ('i', 1), ('ch', 1), ('i', 1), ('w', 1), ('a', 1), ('s', 1), ('e', 1), ('k', 0), ('a', 0), ('i', 0), ('i', 0), ('g', 1), ('e', 1), ('n', 0), ('k', 0), ('i', 0)] """ prosodies = pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True) + logger.debug(f"prosodies: {prosodies}") result: list[tuple[str, int]] = [] current_phrase: list[tuple[str, int]] = [] current_tone = 0 diff --git a/train_ms.py b/train_ms.py index b3259e4..c22e4d1 100644 --- a/train_ms.py +++ b/train_ms.py @@ -258,6 +258,10 @@ def run(): logger.info("Freezing JP bert encoder !!!") for param in net_g.enc_p.ja_bert_proj.parameters(): param.requires_grad = False + if getattr(hps.train, "freeze_style", False): + logger.info("Freezing style encoder !!!") + for param in net_g.enc_p.style_proj.parameters(): + param.requires_grad = False net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank) optim_g = torch.optim.AdamW(