Add: Preparation for ONNX inference support ①
This commit is contained in:
@@ -31,6 +31,7 @@ dependencies = [
|
|||||||
"num2words",
|
"num2words",
|
||||||
"numba",
|
"numba",
|
||||||
"numpy",
|
"numpy",
|
||||||
|
"onnxruntime",
|
||||||
"pydantic>=2.0",
|
"pydantic>=2.0",
|
||||||
"pyopenjtalk-dict",
|
"pyopenjtalk-dict",
|
||||||
"pypinyin",
|
"pypinyin",
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user