Add: Preparation for ONNX inference support ①

This commit is contained in:
tsukumi
2024-09-17 06:53:11 +09:00
parent a037b955ac
commit eede796c4e
4 changed files with 141 additions and 37 deletions

View File

@@ -31,6 +31,7 @@ dependencies = [
"num2words", "num2words",
"numba", "numba",
"numpy", "numpy",
"onnxruntime",
"pydantic>=2.0", "pydantic>=2.0",
"pyopenjtalk-dict", "pyopenjtalk-dict",
"pypinyin", "pypinyin",

View File

@@ -4,12 +4,14 @@ cn2an
# faster-whisper==0.10.1 # faster-whisper==0.10.1
g2p_en g2p_en
GPUtil GPUtil
gradio gradio>=4.32
jieba jieba
# librosa==0.9.2 # librosa==0.9.2
loguru loguru
num2words num2words
numpy<2 numpy<2
onnxruntime
# onnxsim
# protobuf==4.25 # protobuf==4.25
psutil psutil
# punctuators # punctuators
@@ -21,5 +23,6 @@ pyworld-prebuilt
# stable_ts # stable_ts
# tensorboard # tensorboard
torch<2.4 torch<2.4
torchaudio<2.4
transformers transformers
umap-learn umap-learn

View File

@@ -10,6 +10,8 @@ librosa==0.9.2
loguru loguru
num2words num2words
numpy<2 numpy<2
onnxruntime
onnxsim
protobuf==4.25 protobuf==4.25
psutil psutil
punctuators punctuators

View File

@@ -1,8 +1,8 @@
from pathlib import Path from pathlib import Path
from typing import Any, Optional, Union from typing import Any, Optional, Sequence, Union
import numpy as np import numpy as np
import torch import onnxruntime
from numpy.typing import NDArray from numpy.typing import NDArray
from pydantic import BaseModel from pydantic import BaseModel
@@ -28,16 +28,9 @@ from style_bert_vits2.models.models_jp_extra import (
from style_bert_vits2.voice import adjust_voice from style_bert_vits2.voice import adjust_voice
# Gradio の import は重いため、ここでは型チェック時のみ import する
# ライブラリとしての利用を考慮し、TTSModelHolder の _for_gradio() 系メソッド以外では Gradio に依存しないようにする
# _for_gradio() 系メソッドの戻り値の型アテーションを文字列としているのは、Gradio なしで実行できるようにするため
# if TYPE_CHECKING:
# import gradio as gr
class TTSModel: class TTSModel:
""" """
Style-Bert-Vits2 の音声合成モデルを操作するクラス。 Style-Bert-VITS2 の音声合成モデルを操作するクラス。
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。 モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える。
""" """
@@ -46,21 +39,33 @@ class TTSModel:
model_path: Path, model_path: Path,
config_path: Union[Path, HyperParameters], config_path: Union[Path, HyperParameters],
style_vec_path: Union[Path, NDArray[Any]], style_vec_path: Union[Path, NDArray[Any]],
device: str, device: str = "cpu",
onnx_providers: list[str] = ["CPUExecutionProvider"],
onnx_provider_options: Optional[Sequence[dict[str, Any]]] = None,
) -> None: ) -> None:
""" """
Style-Bert-Vits2 の音声合成モデルを初期化する。 Style-Bert-VITS2 の音声合成モデルを初期化する。
この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。 この時点ではモデルはロードされていない (明示的にロードしたい場合は model.load() を呼び出す)。
Args: Args:
model_path (Path): モデル (.safetensors) のパス model_path (Path): モデル (.safetensors / .onnx) のパス
config_path (Union[Path, HyperParameters]): ハイパーパラメータ (config.json) のパス (直接 HyperParameters を指定することも可能) config_path (Union[Path, HyperParameters]): ハイパーパラメータ (config.json) のパス (直接 HyperParameters を指定することも可能)
style_vec_path (Union[Path, NDArray[Any]]): スタイルベクトル (style_vectors.npy) のパス (直接 NDArray を指定することも可能) style_vec_path (Union[Path, NDArray[Any]]): スタイルベクトル (style_vectors.npy) のパス (直接 NDArray を指定することも可能)
device (str): 音声合成時に利用するデバイス (cpu, cuda, mps など) device (str): PyTorch 推論での音声合成時に利用するデバイス (cpu, cuda, mps など)
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
""" """
self.model_path: Path = model_path self.model_path: Path = model_path
self.device: str = device self.device: str = device
self.onnx_providers: list[str] = onnx_providers
self.onnx_provider_options: Optional[Sequence[dict[str, Any]]] = onnx_provider_options # fmt: skip
# ONNX 形式のモデルかどうか
if self.model_path.suffix == ".onnx":
self.is_onnx_model = True
else:
self.is_onnx_model = False
# ハイパーパラメータの Pydantic モデルが直接指定された # ハイパーパラメータの Pydantic モデルが直接指定された
if isinstance(config_path, HyperParameters): if isinstance(config_path, HyperParameters):
@@ -101,12 +106,21 @@ class TTSModel:
) )
self.__style_vector_inference: Optional[Any] = None self.__style_vector_inference: Optional[Any] = None
# __net_g は PyTorch 推論時のみ遅延初期化される
self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None self.__net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
# __inference_session* は ONNX 推論時のみ遅延初期化される
self.__onnx_session: Optional[onnxruntime.InferenceSession] = None
self.__onnx_input_names: Optional[list[str]] = None
self.__onnx_output_names: Optional[list[str]] = None
def load(self) -> None: def load(self) -> None:
""" """
音声合成モデルをデバイスにロードする。 音声合成モデルをデバイスにロードする。
""" """
# PyTorch 推論時
if not self.is_onnx_model:
self.__net_g = get_net_g( self.__net_g = get_net_g(
model_path=str(self.model_path), model_path=str(self.model_path),
version=self.hyper_parameters.version, version=self.hyper_parameters.version,
@@ -114,6 +128,20 @@ class TTSModel:
hps=self.hyper_parameters, hps=self.hyper_parameters,
) )
# ONNX 推論時
else:
self.__onnx_session = onnxruntime.InferenceSession(
path_or_bytes=str(self.model_path),
providers=self.onnx_providers,
provider_options=self.onnx_provider_options,
)
self.__onnx_input_names = [
input.name for input in self.__onnx_session.get_inputs()
]
self.__onnx_output_names = [
output.name for output in self.__onnx_session.get_outputs()
]
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]:
""" """
スタイルベクトルを取得する。 スタイルベクトルを取得する。
@@ -155,6 +183,8 @@ class TTSModel:
) )
# スタイルベクトルを取得するための推論モデルを初期化 # スタイルベクトルを取得するための推論モデルを初期化
import torch
self.__style_vector_inference = pyannote.audio.Inference( self.__style_vector_inference = pyannote.audio.Inference(
model=pyannote.audio.Model.from_pretrained( model=pyannote.audio.Model.from_pretrained(
"pyannote/wespeaker-voxceleb-resnet34-LM" "pyannote/wespeaker-voxceleb-resnet34-LM"
@@ -266,9 +296,7 @@ class TTSModel:
if assist_text == "" or not use_assist_text: if assist_text == "" or not use_assist_text:
assist_text = None assist_text = None
if self.__net_g is None: # スタイルベクトルを取得
self.load()
assert self.__net_g is not None
if reference_audio_path is None: if reference_audio_path is None:
style_id = self.style2id[style] style_id = self.style2id[style]
style_vector = self.__get_style_vector(style_id, style_weight) style_vector = self.__get_style_vector(style_id, style_weight)
@@ -276,6 +304,17 @@ class TTSModel:
style_vector = self.__get_style_vector_from_audio( style_vector = self.__get_style_vector_from_audio(
reference_audio_path, style_weight reference_audio_path, style_weight
) )
# PyTorch 推論時
if not self.is_onnx_model:
import torch
# モデルがロードされていない場合はロードする
if self.__net_g is None:
self.load()
assert self.__net_g is not None
# 通常のテキストから音声を生成
if not line_split: if not line_split:
with torch.no_grad(): with torch.no_grad():
audio = infer( audio = infer(
@@ -295,6 +334,8 @@ class TTSModel:
given_phone=given_phone, given_phone=given_phone,
given_tone=given_tone, given_tone=given_tone,
) )
# 改行ごとに分割して音声を生成
else: else:
texts = text.split("\n") texts = text.split("\n")
texts = [t for t in texts if t != ""] texts = [t for t in texts if t != ""]
@@ -321,7 +362,64 @@ class TTSModel:
if i != len(texts) - 1: if i != len(texts) - 1:
audios.append(np.zeros(int(44100 * split_interval))) audios.append(np.zeros(int(44100 * split_interval)))
audio = np.concatenate(audios) audio = np.concatenate(audios)
# ONNX 推論時
else:
# モデルがロードされていない場合はロードする
if self.__onnx_session is None:
self.load()
assert self.__onnx_session is not None
assert self.__onnx_input_names is not None
assert self.__onnx_output_names is not None
# 通常のテキストから音声を生成
if not line_split:
audio = infer_onnx(
text=text,
sdp_ratio=sdp_ratio,
noise_scale=noise,
noise_scale_w=noise_w,
length_scale=length,
sid=speaker_id,
language=language,
hps=self.hyper_parameters,
device=self.device,
assist_text=assist_text,
assist_text_weight=assist_text_weight,
style_vec=style_vector,
given_phone=given_phone,
given_tone=given_tone,
)
# 改行ごとに分割して音声を生成
else:
texts = text.split("\n")
texts = [t for t in texts if t != ""]
audios = []
for i, t in enumerate(texts):
audios.append(
infer_onnx(
text=t,
sdp_ratio=sdp_ratio,
noise_scale=noise,
noise_scale_w=noise_w,
length_scale=length,
sid=speaker_id,
language=language,
hps=self.hyper_parameters,
device=self.device,
assist_text=assist_text,
assist_text_weight=assist_text_weight,
style_vec=style_vector,
)
)
if i != len(texts) - 1:
audios.append(np.zeros(int(44100 * split_interval)))
audio = np.concatenate(audios)
logger.info("Audio data generated successfully") logger.info("Audio data generated successfully")
if not (pitch_scale == 1.0 and intonation_scale == 1.0): if not (pitch_scale == 1.0 and intonation_scale == 1.0):
_, audio = adjust_voice( _, audio = adjust_voice(
fs=self.hyper_parameters.data.sampling_rate, fs=self.hyper_parameters.data.sampling_rate,
@@ -349,7 +447,7 @@ class TTSModelHolder:
def __init__(self, model_root_dir: Path, device: str) -> None: def __init__(self, model_root_dir: Path, device: str) -> None:
""" """
Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。 Style-Bert-Vits2 の音声合成モデルを管理するクラスを初期化する。
音声合成モデルは下記のように配置されていることを前提とする (.safetensors のファイル名は自由) 。 音声合成モデルは下記のように配置されていることを前提とする (.safetensors / .onnx のファイル名は自由) 。
``` ```
model_root_dir model_root_dir
├── model-name-1 ├── model-name-1
@@ -391,7 +489,7 @@ class TTSModelHolder:
model_files = [ model_files = [
f f
for f in model_dir.iterdir() for f in model_dir.iterdir()
if f.suffix in [".pth", ".pt", ".safetensors"] if f.suffix in [".pth", ".pt", ".safetensors", ".onnx"]
] ]
if len(model_files) == 0: if len(model_files) == 0:
logger.warning(f"No model files found in {model_dir}, so skip it") logger.warning(f"No model files found in {model_dir}, so skip it")