Fix: Disable optimization on InferenceSession initialization to speed up loading of ONNX models

This commit is contained in:
tsukumi
2024-09-22 23:08:00 +09:00
parent 16115a366b
commit 0e0b4bea31
4 changed files with 26 additions and 16 deletions

View File

@@ -21,7 +21,11 @@ from style_bert_vits2.nlp import bert_models
if __name__ == "__main__": if __name__ == "__main__":
start_time = time.time() start_time = time.time()
parser = ArgumentParser() 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() args = parser.parse_args()
# モデルの入出力先ファイルパスを取得 # モデルの入出力先ファイルパスを取得

View File

@@ -30,8 +30,14 @@ from style_bert_vits2.tts_model import TTSModel
if __name__ == "__main__": if __name__ == "__main__":
start_time = time.time() start_time = time.time()
parser = ArgumentParser() parser = ArgumentParser()
parser.add_argument("--model", required=True, help="Path to the model file or directory") parser.add_argument(
parser.add_argument("--force-convert", action="store_true", help="Already converted models will be overwritten") "--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() args = parser.parse_args()
# --model に指定されたパスがディレクトリの時、配下にある全ての .safetensors ファイルを対象に変換する # --model に指定されたパスがディレクトリの時、配下にある全ての .safetensors ファイルを対象に変換する
@@ -61,8 +67,12 @@ if __name__ == "__main__":
# すでに ONNX モデルが存在する場合、--force-convert オプションが指定されていない場合はスキップ # すでに ONNX モデルが存在する場合、--force-convert オプションが指定されていない場合はスキップ
if onnx_optimized_model_path.exists() and not args.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(
print("[bold]If you want to overwrite it, use the --force-convert option.[/bold]") 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"))) print(Rule(characters="=", style=Style(color="blue")))
continue continue
@@ -127,7 +137,9 @@ if __name__ == "__main__":
sdp_ratio: float = 0.0, sdp_ratio: float = 0.0,
noise_scale: float = 0.667, noise_scale: float = 0.667,
noise_scale_w: float = 0.8, 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( return cast(SynthesizerTrnJPExtra, tts_model.net_g).infer(
x, x,
x_lengths, x_lengths,

View File

@@ -95,11 +95,8 @@ def load_model(
start_time = time.time() start_time = time.time()
sess_options = onnxruntime.SessionOptions() sess_options = onnxruntime.SessionOptions()
# 基本的な最適化のみ有効化 # ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする
# ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip
sess_options.graph_optimization_level = (
onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC
)
# エラー以外のログを出力しない # エラー以外のログを出力しない
# 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している
sess_options.log_severity_level = 3 sess_options.log_severity_level = 3

View File

@@ -138,11 +138,8 @@ class TTSModel:
# ONNX 推論時 # ONNX 推論時
else: else:
sess_options = onnxruntime.SessionOptions() sess_options = onnxruntime.SessionOptions()
# 基本的な最適化のみ有効化 # ONNX モデルの作成時にすでに onnxsim により最適化されていることから、ロード高速化のため最適化を無効にする
# ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip
sess_options.graph_optimization_level = (
onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC
)
# エラー以外のログを出力しない # エラー以外のログを出力しない
# 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している
sess_options.log_severity_level = 3 sess_options.log_severity_level = 3