From 9e5222619fd6d00ca9837c2325b728df0ead0b7c Mon Sep 17 00:00:00 2001 From: tsukumi Date: Tue, 12 Mar 2024 18:22:34 +0000 Subject: [PATCH] Refactor: No preloading of BERT models to avoid unnecessary GPU VRAM consumption during training in the Web UI Since the BERT features of the dataset are pre-extracted by bert_gen.py, there is no need to load the BERT model at training time. --- style_bert_vits2/tts_model.py | 8 ++++++++ webui/inference.py | 13 ++++--------- 2 files changed, 12 insertions(+), 9 deletions(-) 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]