diff --git a/app.py b/app.py index 987739b..0d4ecbc 100644 --- a/app.py +++ b/app.py @@ -14,6 +14,7 @@ from style_bert_vits2.constants import GRADIO_THEME, VERSION from style_bert_vits2.nlp.japanese import pyopenjtalk_worker from style_bert_vits2.nlp.japanese.user_dict import update_dict from style_bert_vits2.tts_model import TTSModelHolder +from style_bert_vits2.utils import torch_device_to_onnx_providers # このプロセスからはワーカーを起動して辞書を使いたいので、ここで初期化 @@ -40,7 +41,9 @@ if device == "cuda" and not torch.cuda.is_available(): # download_default_models() path_config = get_path_config() -model_holder = TTSModelHolder(Path(path_config.assets_root), device) +model_holder = TTSModelHolder( + Path(path_config.assets_root), device, torch_device_to_onnx_providers(device) +) with gr.Blocks(theme=GRADIO_THEME) as app: gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})") diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index c477a35..f5af097 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -18,11 +18,12 @@ from style_bert_vits2.constants import ( Languages, ) from style_bert_vits2.logging import logger -from style_bert_vits2.models.infer import InvalidToneError +from style_bert_vits2.nlp import InvalidToneError 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 from style_bert_vits2.tts_model import TTSModelHolder +from style_bert_vits2.utils import torch_device_to_onnx_providers # pyopenjtalk_worker を起動 @@ -533,6 +534,8 @@ if __name__ == "__main__": path_config = get_path_config() assets_root = path_config.assets_root device = "cuda" if torch.cuda.is_available() else "cpu" - model_holder = TTSModelHolder(assets_root, device) + model_holder = TTSModelHolder( + assets_root, device, torch_device_to_onnx_providers(device) + ) app = create_inference_app(model_holder) app.launch(inbrowser=True) diff --git a/gradio_tabs/merge.py b/gradio_tabs/merge.py index b4c9204..c5688e0 100644 --- a/gradio_tabs/merge.py +++ b/gradio_tabs/merge.py @@ -12,6 +12,7 @@ from config import get_path_config from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME from style_bert_vits2.logging import logger from style_bert_vits2.tts_model import TTSModel, TTSModelHolder +from style_bert_vits2.utils import torch_device_to_onnx_providers voice_keys = ["dec"] @@ -1524,8 +1525,9 @@ def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks: if __name__ == "__main__": + device = "cuda" if torch.cuda.is_available() else "cpu" model_holder = TTSModelHolder( - assets_root, device="cuda" if torch.cuda.is_available() else "cpu" + assets_root, device, torch_device_to_onnx_providers(device) ) app = create_merge_app(model_holder) app.launch(inbrowser=True) diff --git a/server_editor.py b/server_editor.py index fe7b600..eee8754 100644 --- a/server_editor.py +++ b/server_editor.py @@ -53,6 +53,7 @@ from style_bert_vits2.nlp.japanese.user_dict import ( update_dict, ) from style_bert_vits2.tts_model import TTSModelHolder, TTSModelInfo +from style_bert_vits2.utils import torch_device_to_onnx_providers # ---フロントエンド部分に関する処理--- @@ -193,14 +194,14 @@ skip_static_files = bool(args.skip_static_files) ## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする bert_models.load_model(Languages.JP, device_map=device) bert_models.load_tokenizer(Languages.JP) -if device == "cpu": - onnx_provider = "CPUExecutionProvider" -else: - onnx_provider = ("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"}) -onnx_bert_models.load_model(Languages.JP, onnx_providers=[onnx_provider]) +onnx_bert_models.load_model( + Languages.JP, onnx_providers=torch_device_to_onnx_providers(device) +) onnx_bert_models.load_tokenizer(Languages.JP) -model_holder = TTSModelHolder(model_dir, device) +model_holder = TTSModelHolder( + model_dir, device, torch_device_to_onnx_providers(device) +) if len(model_holder.model_names) == 0: logger.error(f"Models not found in {model_dir}.") sys.exit(1) diff --git a/server_fastapi.py b/server_fastapi.py index e73c3be..eca242e 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -38,6 +38,7 @@ from style_bert_vits2.nlp import bert_models, onnx_bert_models from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk from style_bert_vits2.nlp.japanese.user_dict import update_dict from style_bert_vits2.tts_model import TTSModel, TTSModelHolder +from style_bert_vits2.utils import torch_device_to_onnx_providers config = get_config() @@ -103,15 +104,15 @@ if __name__ == "__main__": bert_models.load_tokenizer(Languages.EN) bert_models.load_model(Languages.ZH, device_map=device) bert_models.load_tokenizer(Languages.ZH) - if device == "cpu": - onnx_provider = "CPUExecutionProvider" - else: - onnx_provider = ("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"}) - onnx_bert_models.load_model(Languages.JP, onnx_providers=[onnx_provider]) + onnx_bert_models.load_model( + Languages.JP, onnx_providers=torch_device_to_onnx_providers(device) + ) onnx_bert_models.load_tokenizer(Languages.JP) model_dir = Path(args.dir) - model_holder = TTSModelHolder(model_dir, device) + model_holder = TTSModelHolder( + model_dir, device, torch_device_to_onnx_providers(device) + ) if len(model_holder.model_names) == 0: logger.error(f"Models not found in {model_dir}.") sys.exit(1) diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index e9ea28b..41ef6cc 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -131,6 +131,9 @@ class TTSModel: device=self.device, hps=self.hyper_parameters, ) + logger.info( + f"Model loaded successfully from {self.model_path} to \"{self.device}\" device ({time.time() - start_time:.2f}s)" + ) # ONNX 推論時 else: @@ -138,10 +141,9 @@ class TTSModel: path_or_bytes=str(self.model_path), providers=self.onnx_providers, ) - - logger.info( - f"Model loaded successfully from {self.model_path} ({time.time() - start_time:.2f}s)" - ) + logger.info( + f"Model loaded successfully from {self.model_path} to {self.onnx_session.get_providers()[0]} ({time.time() - start_time:.2f}s)" + ) def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: """ @@ -447,13 +449,18 @@ class TTSModelInfo(BaseModel): class TTSModelHolder: """ - Style-Bert-Vits2 の音声合成モデルを管理するクラス。 + Style-Bert-VITS2 の音声合成モデルを管理するクラス。 model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。 """ - def __init__(self, model_root_dir: Path, device: str) -> None: + def __init__( + self, + model_root_dir: Path, + device: str, + onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]], + ) -> None: """ - Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。 + Style-Bert-VITS2 の音声合成モデルを管理するクラスを初期化する。 音声合成モデルは下記のように配置されていることを前提とする (.safetensors / .onnx のファイル名は自由) 。 ``` model_root_dir @@ -470,11 +477,13 @@ class TTSModelHolder: Args: model_root_dir (Path): 音声合成モデルが配置されているディレクトリのパス - device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など) + device (str): PyTorch 推論での音声合成時に利用するデバイス (cpu, cuda, mps など) + onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など) """ self.root_dir: Path = model_root_dir self.device: str = device + self.onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = onnx_providers # fmt: skip self.model_files_dict: dict[str, list[Path]] = {} self.current_model: Optional[TTSModel] = None self.model_names: list[str] = [] @@ -551,6 +560,7 @@ class TTSModelHolder: config_path=self.root_dir / model_name / "config.json", style_vec_path=self.root_dir / model_name / "style_vectors.npy", device=self.device, + onnx_providers=self.onnx_providers, ) return self.current_model @@ -580,6 +590,7 @@ class TTSModelHolder: config_path=self.root_dir / model_name / "config.json", style_vec_path=self.root_dir / model_name / "style_vectors.npy", device=self.device, + onnx_providers=self.onnx_providers, ) speakers = list(self.current_model.spk2id.keys()) styles = list(self.current_model.style2id.keys()) diff --git a/style_bert_vits2/utils/__init__.py b/style_bert_vits2/utils/__init__.py index e69de29..d86b476 100644 --- a/style_bert_vits2/utils/__init__.py +++ b/style_bert_vits2/utils/__init__.py @@ -0,0 +1,15 @@ +from typing import Any, Sequence, Union + + +def torch_device_to_onnx_providers( + device: str, +) -> Sequence[Union[str, tuple[str, dict[str, Any]]]]: + if device.startswith("cuda"): + # cudnn_conv_algo_search を DEFAULT にすると推論速度が大幅に向上する + # ref: https://medium.com/neuml/debug-onnx-gpu-performance-c9290fe07459 + return [ + ("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"}), + ("CPUExecutionProvider", {}), + ] + else: + return ["CPUExecutionProvider"] diff --git a/tests/test_main.py b/tests/test_main.py index e2e6530..2aa852f 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -5,10 +5,13 @@ from style_bert_vits2.constants import BASE_DIR, Languages from style_bert_vits2.tts_model import TTSModelHolder -def synthesize(device: str = "cpu"): +def synthesize( + device: str = "cpu", + onnx_providers: list[str] = ["CPUExecutionProvider"], +): # 音声合成モデルが配置されていれば、音声合成を実行 - model_holder = TTSModelHolder(BASE_DIR / "model_assets", device) + model_holder = TTSModelHolder(BASE_DIR / "model_assets", device, onnx_providers) if len(model_holder.models_info) > 0: # jvnv-F2-jp モデルを探す