From db0acd7b5f0eda8dad5af2bca1ab1fed9dd6312c Mon Sep 17 00:00:00 2001 From: tsukumi Date: Mon, 23 Sep 2024 05:59:12 +0900 Subject: [PATCH] Add: Support for ONNX inference and ONNX conversion of Non-JP-Extra models --- convert_onnx.py | 103 +++++++++++++++++++- style_bert_vits2/models/hyper_parameters.py | 3 +- style_bert_vits2/models/infer_onnx.py | 66 ++++++++----- 3 files changed, 140 insertions(+), 32 deletions(-) diff --git a/convert_onnx.py b/convert_onnx.py index 7b6c4f7..b323728 100644 --- a/convert_onnx.py +++ b/convert_onnx.py @@ -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 モデルを最適化 diff --git a/style_bert_vits2/models/hyper_parameters.py b/style_bert_vits2/models/hyper_parameters.py index df31573..e27eb42 100644 --- a/style_bert_vits2/models/hyper_parameters.py +++ b/style_bert_vits2/models/hyper_parameters.py @@ -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 diff --git a/style_bert_vits2/models/infer_onnx.py b/style_bert_vits2/models/infer_onnx.py index 41d042e..725340a 100644 --- a/style_bert_vits2/models/infer_onnx.py +++ b/style_bert_vits2/models/infer_onnx.py @@ -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,27 +170,41 @@ def infer_onnx( 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() 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]