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__":
|
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()
|
||||||
|
|
||||||
# モデルの入出力先ファイルパスを取得
|
# モデルの入出力先ファイルパスを取得
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user