diff --git a/convert_bert_onnx.py b/convert_bert_onnx.py index a37a971..8140a83 100644 --- a/convert_bert_onnx.py +++ b/convert_bert_onnx.py @@ -21,7 +21,11 @@ from style_bert_vits2.nlp import bert_models if __name__ == "__main__": start_time = time.time() parser = ArgumentParser() - parser.add_argument("--language", default=Languages.JP, help="Language of the BERT model to be converted") + parser.add_argument( + "--language", + default=Languages.JP, + help="Language of the BERT model to be converted", + ) args = parser.parse_args() # モデルの入出力先ファイルパスを取得 diff --git a/convert_onnx.py b/convert_onnx.py index 6f34ceb..7b6c4f7 100644 --- a/convert_onnx.py +++ b/convert_onnx.py @@ -30,8 +30,14 @@ from style_bert_vits2.tts_model import TTSModel if __name__ == "__main__": start_time = time.time() parser = ArgumentParser() - parser.add_argument("--model", required=True, help="Path to the model file or directory") - parser.add_argument("--force-convert", action="store_true", help="Already converted models will be overwritten") + parser.add_argument( + "--model", required=True, help="Path to the model file or directory" + ) + parser.add_argument( + "--force-convert", + action="store_true", + help="Already converted models will be overwritten", + ) args = parser.parse_args() # --model に指定されたパスがディレクトリの時、配下にある全ての .safetensors ファイルを対象に変換する @@ -61,8 +67,12 @@ if __name__ == "__main__": # すでに ONNX モデルが存在する場合、--force-convert オプションが指定されていない場合はスキップ if onnx_optimized_model_path.exists() and not args.force_convert: - print(f"[bold yellow]ONNX model already exists: {onnx_optimized_model_path}[/bold yellow]") - print("[bold]If you want to overwrite it, use the --force-convert option.[/bold]") + print( + f"[bold yellow]ONNX model already exists: {onnx_optimized_model_path}[/bold yellow]" + ) + print( + "[bold]If you want to overwrite it, use the --force-convert option.[/bold]" + ) print(Rule(characters="=", style=Style(color="blue"))) continue @@ -127,7 +137,9 @@ if __name__ == "__main__": 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, ...]]: + ) -> tuple[ + torch.Tensor, torch.Tensor, torch.Tensor, tuple[torch.Tensor, ...] + ]: return cast(SynthesizerTrnJPExtra, tts_model.net_g).infer( x, x_lengths, diff --git a/style_bert_vits2/nlp/onnx_bert_models.py b/style_bert_vits2/nlp/onnx_bert_models.py index 442f2ad..04c8d65 100644 --- a/style_bert_vits2/nlp/onnx_bert_models.py +++ b/style_bert_vits2/nlp/onnx_bert_models.py @@ -95,11 +95,8 @@ def load_model( start_time = time.time() sess_options = onnxruntime.SessionOptions() - # 基本的な最適化のみ有効化 - # ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する - sess_options.graph_optimization_level = ( - onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC - ) + # ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする + sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip # エラー以外のログを出力しない # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している sess_options.log_severity_level = 3 diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index bd56de7..8c9570c 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -138,11 +138,8 @@ class TTSModel: # ONNX 推論時 else: sess_options = onnxruntime.SessionOptions() - # 基本的な最適化のみ有効化 - # ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する - sess_options.graph_optimization_level = ( - onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC - ) + # ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする + sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip # エラー以外のログを出力しない # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している sess_options.log_severity_level = 3