Add: Support for ONNX inference and ONNX conversion of Non-JP-Extra models

This commit is contained in:
tsukumi
2024-09-23 05:59:12 +09:00
parent cc1d6120fe
commit db0acd7b5f
3 changed files with 140 additions and 32 deletions

View File

@@ -21,6 +21,7 @@ from style_bert_vits2.constants import (
Languages,
)
from style_bert_vits2.models.infer import get_text
from style_bert_vits2.models.models import SynthesizerTrn
from style_bert_vits2.models.models_jp_extra import (
SynthesizerTrn as SynthesizerTrnJPExtra,
)
@@ -122,10 +123,10 @@ if __name__ == "__main__":
noise_scale_w = torch.tensor(0.8)
# JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
if tts_model.hyper_parameters.data.use_jp_extra is True:
if isinstance(tts_model.net_g, SynthesizerTrnJPExtra):
# SynthesizerTrnJPExtra の forward メソッドをオーバーライド
def forward(
def forward_jp_extra(
x: torch.Tensor,
x_lengths: torch.Tensor,
sid: torch.Tensor,
@@ -154,7 +155,7 @@ if __name__ == "__main__":
noise_scale_w=noise_scale_w,
)
tts_model.net_g.forward = forward # type: ignore
tts_model.net_g.forward = forward_jp_extra # type: ignore
# モデルを ONNX に変換
print(Rule(characters="=", style=Style(color="blue")))
@@ -208,9 +209,101 @@ if __name__ == "__main__":
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
)
# 非 JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
else:
raise NotImplementedError(
"non-JP-Extra model architecture is not implemented yet"
# SynthesizerTrn の forward メソッドをオーバーライド
def forward_non_jp_extra(
x: torch.Tensor,
x_lengths: torch.Tensor,
sid: torch.Tensor,
tone: torch.Tensor,
language: torch.Tensor,
bert: torch.Tensor,
ja_bert: torch.Tensor,
en_bert: torch.Tensor,
style_vec: torch.Tensor,
length_scale: float = 1.0,
sdp_ratio: float = 0.0,
noise_scale: float = 0.667,
noise_scale_w: float = 0.8,
) -> tuple[
torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...]
]:
return cast(SynthesizerTrn, tts_model.net_g).infer(
x,
x_lengths,
sid,
tone,
language,
bert,
ja_bert,
en_bert,
style_vec,
length_scale=length_scale,
sdp_ratio=sdp_ratio,
noise_scale=noise_scale,
noise_scale_w=noise_scale_w,
)
tts_model.net_g.forward = forward_non_jp_extra # type: ignore
# モデルを ONNX に変換
print(Rule(characters="=", style=Style(color="blue")))
print(
f"[bold cyan]Exporting ONNX model... (Architecture: Non-JP-Extra)[/bold cyan]"
)
print(Rule(characters="=", style=Style(color="blue")))
export_start_time = time.time()
torch.onnx.export(
model=tts_model.net_g,
args=(
x_tst,
x_tst_lengths,
sid_tensor,
tones,
lang_ids,
bert,
ja_bert,
en_bert,
style_vec_tensor,
length_scale,
sdp_ratio,
noise_scale,
noise_scale_w,
),
f=str(onnx_temp_model_path),
verbose=False,
input_names=[
"x_tst",
"x_tst_lengths",
"sid",
"tones",
"language",
"bert",
"ja_bert",
"en_bert",
"style_vec",
"length_scale",
"sdp_ratio",
"noise_scale",
"noise_scale_w",
],
output_names=["output"],
dynamic_axes={
"x_tst": {0: "batch_size", 1: "x_tst_max_length"},
"x_tst_lengths": {0: "batch_size"},
"sid": {0: "batch_size"},
"tones": {0: "batch_size", 1: "x_tst_max_length"},
"language": {0: "batch_size", 1: "x_tst_max_length"},
"bert": {0: "batch_size", 2: "x_tst_max_length"},
"ja_bert": {0: "batch_size", 2: "x_tst_max_length"},
"en_bert": {0: "batch_size", 2: "x_tst_max_length"},
"style_vec": {0: "batch_size"},
},
)
print(
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
)
# ONNX モデルを最適化

View File

@@ -38,7 +38,8 @@ class HyperParametersTrain(BaseModel):
class HyperParametersData(BaseModel):
use_jp_extra: bool = False # このフィールドが存在しない旧モデルとの互換性のために False をデフォルト値とする
# use_jp_extra フィールドが存在しない旧モデルとの互換性のために False をデフォルト値とする
use_jp_extra: bool = False
training_files: str = "Data/Dummy/train.list"
validation_files: str = "Data/Dummy/val.list"
max_wav_value: float = 32768.0

View File

@@ -75,15 +75,15 @@ def get_text_onnx(
if language_str == Languages.ZH:
bert = bert_ori
ja_bert = np.zeros((1024, len(phone)))
en_bert = np.zeros((1024, len(phone)))
ja_bert = np.zeros((1024, len(phone)), dtype=np.float32)
en_bert = np.zeros((1024, len(phone)), dtype=np.float32)
elif language_str == Languages.JP:
bert = np.zeros((1024, len(phone)))
bert = np.zeros((1024, len(phone)), dtype=np.float32)
ja_bert = bert_ori
en_bert = np.zeros((1024, len(phone)))
en_bert = np.zeros((1024, len(phone)), dtype=np.float32)
elif language_str == Languages.EN:
bert = np.zeros((1024, len(phone)))
ja_bert = np.zeros((1024, len(phone)))
bert = np.zeros((1024, len(phone)), dtype=np.float32)
ja_bert = np.zeros((1024, len(phone)), dtype=np.float32)
en_bert = bert_ori
else:
raise ValueError("language_str should be ZH, JP or EN")
@@ -170,6 +170,22 @@ def infer_onnx(
np.array([noise_scale], dtype=np.float32),
np.array([noise_scale_w], dtype=np.float32),
]
else:
input_tensor = [
x_tst,
x_tst_lengths,
sid_tensor,
tones,
lang_ids,
bert,
ja_bert,
en_bert,
style_vec_tensor,
np.array([length_scale], dtype=np.float32),
np.array([sdp_ratio], dtype=np.float32),
np.array([noise_scale], dtype=np.float32),
np.array([noise_scale_w], dtype=np.float32),
]
first_provider = onnx_session.get_providers()[0]
if first_provider == "CUDAExecutionProvider":
@@ -189,8 +205,6 @@ def infer_onnx(
io_binding.bind_output(output_name, device_type)
onnx_session.run_with_iobinding(io_binding)
output = io_binding.get_outputs()
else:
raise NotImplementedError("Not implemented yet")
audio = output[0].numpy()[0, 0]