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

View File

@@ -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)

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.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)