Improve: DirectML inference performance

This commit is contained in:
tsukumi
2024-09-23 20:20:46 +09:00
parent 08af692835
commit b67e50fd8d

View File

@@ -139,7 +139,17 @@ class TTSModel:
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 ## DmlExecutionProvider が先頭に指定されているときのみ、DirectML 推論の高速化のためすべての最適化を有効にする
assert len(self.onnx_providers) > 0
first_provider_name = (
self.onnx_providers[0]
if type(self.onnx_providers[0]) is str
else self.onnx_providers[0][0]
)
if first_provider_name == "DmlExecutionProvider":
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_ALL # fmt: skip
else:
sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL # fmt: skip
# エラー以外のログを出力しない # エラー以外のログを出力しない
# 本来は 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