Add: Support for ONNX inference and ONNX conversion of Non-JP-Extra models
This commit is contained in:
103
convert_onnx.py
103
convert_onnx.py
@@ -21,6 +21,7 @@ from style_bert_vits2.constants import (
|
|||||||
Languages,
|
Languages,
|
||||||
)
|
)
|
||||||
from style_bert_vits2.models.infer import get_text
|
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 (
|
from style_bert_vits2.models.models_jp_extra import (
|
||||||
SynthesizerTrn as SynthesizerTrnJPExtra,
|
SynthesizerTrn as SynthesizerTrnJPExtra,
|
||||||
)
|
)
|
||||||
@@ -122,10 +123,10 @@ if __name__ == "__main__":
|
|||||||
noise_scale_w = torch.tensor(0.8)
|
noise_scale_w = torch.tensor(0.8)
|
||||||
|
|
||||||
# JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
|
# JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
|
||||||
if tts_model.hyper_parameters.data.use_jp_extra is True:
|
if isinstance(tts_model.net_g, SynthesizerTrnJPExtra):
|
||||||
|
|
||||||
# SynthesizerTrnJPExtra の forward メソッドをオーバーライド
|
# SynthesizerTrnJPExtra の forward メソッドをオーバーライド
|
||||||
def forward(
|
def forward_jp_extra(
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
x_lengths: torch.Tensor,
|
x_lengths: torch.Tensor,
|
||||||
sid: torch.Tensor,
|
sid: torch.Tensor,
|
||||||
@@ -154,7 +155,7 @@ if __name__ == "__main__":
|
|||||||
noise_scale_w=noise_scale_w,
|
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 に変換
|
# モデルを ONNX に変換
|
||||||
print(Rule(characters="=", style=Style(color="blue")))
|
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]"
|
f"[bold green]ONNX model exported to {onnx_temp_model_path} ({time.time() - export_start_time:.2f}s)[/bold green]"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 非 JP-Extra モデルアーキテクチャ向けの ONNX 変換ロジック
|
||||||
else:
|
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 モデルを最適化
|
# ONNX モデルを最適化
|
||||||
|
|||||||
@@ -38,7 +38,8 @@ class HyperParametersTrain(BaseModel):
|
|||||||
|
|
||||||
|
|
||||||
class HyperParametersData(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"
|
training_files: str = "Data/Dummy/train.list"
|
||||||
validation_files: str = "Data/Dummy/val.list"
|
validation_files: str = "Data/Dummy/val.list"
|
||||||
max_wav_value: float = 32768.0
|
max_wav_value: float = 32768.0
|
||||||
|
|||||||
@@ -75,15 +75,15 @@ def get_text_onnx(
|
|||||||
|
|
||||||
if language_str == Languages.ZH:
|
if language_str == Languages.ZH:
|
||||||
bert = bert_ori
|
bert = bert_ori
|
||||||
ja_bert = np.zeros((1024, len(phone)))
|
ja_bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
en_bert = np.zeros((1024, len(phone)))
|
en_bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
elif language_str == Languages.JP:
|
elif language_str == Languages.JP:
|
||||||
bert = np.zeros((1024, len(phone)))
|
bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
ja_bert = bert_ori
|
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:
|
elif language_str == Languages.EN:
|
||||||
bert = np.zeros((1024, len(phone)))
|
bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
ja_bert = np.zeros((1024, len(phone)))
|
ja_bert = np.zeros((1024, len(phone)), dtype=np.float32)
|
||||||
en_bert = bert_ori
|
en_bert = bert_ori
|
||||||
else:
|
else:
|
||||||
raise ValueError("language_str should be ZH, JP or EN")
|
raise ValueError("language_str should be ZH, JP or EN")
|
||||||
@@ -170,27 +170,41 @@ def infer_onnx(
|
|||||||
np.array([noise_scale], dtype=np.float32),
|
np.array([noise_scale], dtype=np.float32),
|
||||||
np.array([noise_scale_w], dtype=np.float32),
|
np.array([noise_scale_w], dtype=np.float32),
|
||||||
]
|
]
|
||||||
|
|
||||||
first_provider = onnx_session.get_providers()[0]
|
|
||||||
if first_provider == "CUDAExecutionProvider":
|
|
||||||
device_type = "cuda"
|
|
||||||
elif first_provider == "DmlExecutionProvider":
|
|
||||||
device_type = "dml"
|
|
||||||
else:
|
|
||||||
device_type = "cpu"
|
|
||||||
|
|
||||||
# GPU メモリに入力テンソルを割り当て
|
|
||||||
io_binding = onnx_session.io_binding()
|
|
||||||
for name, value in zip(input_names, input_tensor):
|
|
||||||
gpu_tensor = onnxruntime.OrtValue.ortvalue_from_numpy(value, device_type)
|
|
||||||
io_binding.bind_ortvalue_input(name, gpu_tensor)
|
|
||||||
|
|
||||||
# 推論の実行
|
|
||||||
io_binding.bind_output(output_name, device_type)
|
|
||||||
onnx_session.run_with_iobinding(io_binding)
|
|
||||||
output = io_binding.get_outputs()
|
|
||||||
else:
|
else:
|
||||||
raise NotImplementedError("Not implemented yet")
|
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":
|
||||||
|
device_type = "cuda"
|
||||||
|
elif first_provider == "DmlExecutionProvider":
|
||||||
|
device_type = "dml"
|
||||||
|
else:
|
||||||
|
device_type = "cpu"
|
||||||
|
|
||||||
|
# GPU メモリに入力テンソルを割り当て
|
||||||
|
io_binding = onnx_session.io_binding()
|
||||||
|
for name, value in zip(input_names, input_tensor):
|
||||||
|
gpu_tensor = onnxruntime.OrtValue.ortvalue_from_numpy(value, device_type)
|
||||||
|
io_binding.bind_ortvalue_input(name, gpu_tensor)
|
||||||
|
|
||||||
|
# 推論の実行
|
||||||
|
io_binding.bind_output(output_name, device_type)
|
||||||
|
onnx_session.run_with_iobinding(io_binding)
|
||||||
|
output = io_binding.get_outputs()
|
||||||
|
|
||||||
audio = output[0].numpy()[0, 0]
|
audio = output[0].numpy()[0, 0]
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user