Fix: CUDAExecutionProvider was not being used to infer Style-Bert-VITS2 ONNX models even though CUDA was available

This commit is contained in:
tsukumi
2024-09-18 04:37:23 +09:00
parent e5a05a30cb
commit 7d721fdf2b
8 changed files with 65 additions and 26 deletions

5
app.py
View File

@@ -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 import pyopenjtalk_worker
from style_bert_vits2.nlp.japanese.user_dict import update_dict from style_bert_vits2.nlp.japanese.user_dict import update_dict
from style_bert_vits2.tts_model import TTSModelHolder 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() # download_default_models()
path_config = get_path_config() 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: with gr.Blocks(theme=GRADIO_THEME) as app:
gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})") gr.Markdown(f"# Style-Bert-VITS2 WebUI (version {VERSION})")

View File

@@ -18,11 +18,12 @@ from style_bert_vits2.constants import (
Languages, Languages,
) )
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.nlp import InvalidToneError
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
from style_bert_vits2.tts_model import TTSModelHolder from style_bert_vits2.tts_model import TTSModelHolder
from style_bert_vits2.utils import torch_device_to_onnx_providers
# pyopenjtalk_worker を起動 # pyopenjtalk_worker を起動
@@ -533,6 +534,8 @@ if __name__ == "__main__":
path_config = get_path_config() path_config = get_path_config()
assets_root = path_config.assets_root assets_root = path_config.assets_root
device = "cuda" if torch.cuda.is_available() else "cpu" 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 = create_inference_app(model_holder)
app.launch(inbrowser=True) app.launch(inbrowser=True)

View File

@@ -12,6 +12,7 @@ from config import get_path_config
from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
from style_bert_vits2.utils import torch_device_to_onnx_providers
voice_keys = ["dec"] voice_keys = ["dec"]
@@ -1524,8 +1525,9 @@ def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks:
if __name__ == "__main__": if __name__ == "__main__":
device = "cuda" if torch.cuda.is_available() else "cpu"
model_holder = TTSModelHolder( 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 = create_merge_app(model_holder)
app.launch(inbrowser=True) app.launch(inbrowser=True)

View File

@@ -53,6 +53,7 @@ from style_bert_vits2.nlp.japanese.user_dict import (
update_dict, update_dict,
) )
from style_bert_vits2.tts_model import TTSModelHolder, TTSModelInfo 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 モデル/トークナイザーのみロードする ## server_editor.py は日本語にしか対応していないため、日本語の BERT モデル/トークナイザーのみロードする
bert_models.load_model(Languages.JP, device_map=device) bert_models.load_model(Languages.JP, device_map=device)
bert_models.load_tokenizer(Languages.JP) bert_models.load_tokenizer(Languages.JP)
if device == "cpu": onnx_bert_models.load_model(
onnx_provider = "CPUExecutionProvider" Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
else: )
onnx_provider = ("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"})
onnx_bert_models.load_model(Languages.JP, onnx_providers=[onnx_provider])
onnx_bert_models.load_tokenizer(Languages.JP) 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: if len(model_holder.model_names) == 0:
logger.error(f"Models not found in {model_dir}.") logger.error(f"Models not found in {model_dir}.")
sys.exit(1) sys.exit(1)

View File

@@ -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 import pyopenjtalk_worker as pyopenjtalk
from style_bert_vits2.nlp.japanese.user_dict import update_dict from style_bert_vits2.nlp.japanese.user_dict import update_dict
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
from style_bert_vits2.utils import torch_device_to_onnx_providers
config = get_config() config = get_config()
@@ -103,15 +104,15 @@ if __name__ == "__main__":
bert_models.load_tokenizer(Languages.EN) bert_models.load_tokenizer(Languages.EN)
bert_models.load_model(Languages.ZH, device_map=device) bert_models.load_model(Languages.ZH, device_map=device)
bert_models.load_tokenizer(Languages.ZH) bert_models.load_tokenizer(Languages.ZH)
if device == "cpu": onnx_bert_models.load_model(
onnx_provider = "CPUExecutionProvider" Languages.JP, onnx_providers=torch_device_to_onnx_providers(device)
else: )
onnx_provider = ("CUDAExecutionProvider", {"cudnn_conv_algo_search": "DEFAULT"})
onnx_bert_models.load_model(Languages.JP, onnx_providers=[onnx_provider])
onnx_bert_models.load_tokenizer(Languages.JP) onnx_bert_models.load_tokenizer(Languages.JP)
model_dir = Path(args.dir) 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: if len(model_holder.model_names) == 0:
logger.error(f"Models not found in {model_dir}.") logger.error(f"Models not found in {model_dir}.")
sys.exit(1) sys.exit(1)

View File

@@ -131,6 +131,9 @@ class TTSModel:
device=self.device, device=self.device,
hps=self.hyper_parameters, 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 推論時 # ONNX 推論時
else: else:
@@ -138,9 +141,8 @@ class TTSModel:
path_or_bytes=str(self.model_path), path_or_bytes=str(self.model_path),
providers=self.onnx_providers, providers=self.onnx_providers,
) )
logger.info( logger.info(
f"Model loaded successfully from {self.model_path} ({time.time() - start_time:.2f}s)" 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]: def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
@@ -447,13 +449,18 @@ class TTSModelInfo(BaseModel):
class TTSModelHolder: class TTSModelHolder:
""" """
Style-Bert-Vits2 の音声合成モデルを管理するクラス。 Style-Bert-VITS2 の音声合成モデルを管理するクラス。
model_holder.models_info から指定されたディレクトリ内にある音声合成モデルの一覧を取得できる。 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 のファイル名は自由) 。 音声合成モデルは下記のように配置されていることを前提とする (.safetensors / .onnx のファイル名は自由) 。
``` ```
model_root_dir model_root_dir
@@ -470,11 +477,13 @@ class TTSModelHolder:
Args: Args:
model_root_dir (Path): 音声合成モデルが配置されているディレクトリのパス 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.root_dir: Path = model_root_dir
self.device: str = device 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.model_files_dict: dict[str, list[Path]] = {}
self.current_model: Optional[TTSModel] = None self.current_model: Optional[TTSModel] = None
self.model_names: list[str] = [] self.model_names: list[str] = []
@@ -551,6 +560,7 @@ class TTSModelHolder:
config_path=self.root_dir / model_name / "config.json", config_path=self.root_dir / model_name / "config.json",
style_vec_path=self.root_dir / model_name / "style_vectors.npy", style_vec_path=self.root_dir / model_name / "style_vectors.npy",
device=self.device, device=self.device,
onnx_providers=self.onnx_providers,
) )
return self.current_model return self.current_model
@@ -580,6 +590,7 @@ class TTSModelHolder:
config_path=self.root_dir / model_name / "config.json", config_path=self.root_dir / model_name / "config.json",
style_vec_path=self.root_dir / model_name / "style_vectors.npy", style_vec_path=self.root_dir / model_name / "style_vectors.npy",
device=self.device, device=self.device,
onnx_providers=self.onnx_providers,
) )
speakers = list(self.current_model.spk2id.keys()) speakers = list(self.current_model.spk2id.keys())
styles = list(self.current_model.style2id.keys()) styles = list(self.current_model.style2id.keys())

View File

@@ -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"]

View File

@@ -5,10 +5,13 @@ from style_bert_vits2.constants import BASE_DIR, Languages
from style_bert_vits2.tts_model import TTSModelHolder 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: if len(model_holder.models_info) > 0:
# jvnv-F2-jp モデルを探す # jvnv-F2-jp モデルを探す