Fix: Disable optimization on InferenceSession initialization to speed up loading of ONNX models
This commit is contained in:
@@ -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()
|
||||
|
||||
# モデルの入出力先ファイルパスを取得
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user