Fix: CUDAExecutionProvider was not being used to infer Style-Bert-VITS2 ONNX models even though CUDA was available
This commit is contained in:
5
app.py
5
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})")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,9 +141,8 @@ 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)"
|
||||
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())
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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 モデルを探す
|
||||
|
||||
Reference in New Issue
Block a user