diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index f7d4f21..da30f48 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -29,6 +29,7 @@ from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models_jp_extra import ( SynthesizerTrn as SynthesizerTrnJPExtra, ) +from style_bert_vits2.nlp import bert_models from style_bert_vits2.voice import adjust_voice @@ -379,6 +380,13 @@ class TTSModelHolder: def get_model_for_gradio( self, model_name: str, model_path_str: str ) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: + bert_models.load_model(Languages.JP) + bert_models.load_tokenizer(Languages.JP) + bert_models.load_model(Languages.EN) + bert_models.load_tokenizer(Languages.EN) + bert_models.load_model(Languages.ZH) + bert_models.load_tokenizer(Languages.ZH) + model_path = Path(model_path_str) if model_name not in self.model_files_dict: raise ValueError(f"Model `{model_name}` is not found") diff --git a/webui/inference.py b/webui/inference.py index db59829..d711846 100644 --- a/webui/inference.py +++ b/webui/inference.py @@ -19,7 +19,6 @@ from style_bert_vits2.constants import ( ) from style_bert_vits2.logging import logger from style_bert_vits2.models.infer import InvalidToneError -from style_bert_vits2.nlp import bert_models from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone from style_bert_vits2.nlp.japanese.normalizer import normalize_text @@ -30,14 +29,10 @@ from style_bert_vits2.tts_model import TTSModelHolder ## pyopenjtalk_worker は TCP ソケットサーバーのため、ここで起動する pyopenjtalk.initialize_worker() -# 事前に BERT モデル/トークナイザーをロードしておく -## ここでロードしなくても必要になった際に自動ロードされるが、時間がかかるため事前にロードしておいた方が体験が良い -bert_models.load_model(Languages.JP) -bert_models.load_tokenizer(Languages.JP) -bert_models.load_model(Languages.EN) -bert_models.load_tokenizer(Languages.EN) -bert_models.load_model(Languages.ZH) -bert_models.load_tokenizer(Languages.ZH) +# Web UI での学習時の無駄な GPU VRAM 消費を避けるため、あえてここでは BERT モデルの事前ロードを行わない +# データセットの BERT 特徴量は事前に bert_gen.py により抽出されているため、学習時に BERT モデルをロードしておく必要はない +# BERT モデルの事前ロードは「ロード」ボタン押下時に実行される TTSModelHolder.get_model_for_gradio() 内で行われる +# Web UI での学習時、音声合成タブの「ロード」ボタンを押さなければ、BERT モデルが VRAM にロードされていない状態で学習を開始できる languages = [lang.value for lang in Languages]