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