diff --git a/style_bert_vits2/nlp/onnx_bert_models.py b/style_bert_vits2/nlp/onnx_bert_models.py index b0ad870..b14397d 100644 --- a/style_bert_vits2/nlp/onnx_bert_models.py +++ b/style_bert_vits2/nlp/onnx_bert_models.py @@ -93,10 +93,20 @@ def load_model( else: model_path = Path(pretrained_model_name_or_path).resolve() / "model.onnx" - # BERT モデルをロードし、辞書に格納して返す start_time = time.time() + sess_options = onnxruntime.SessionOptions() + # 基本的な最適化のみ有効化 + # ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する + sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC + # エラー以外のログを出力しない + # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している + sess_options.log_severity_level = 3 + onnxruntime.set_default_logger_severity(3) + + # BERT モデルをロードし、辞書に格納して返す __loaded_models[language] = onnxruntime.InferenceSession( model_path, + sess_options=sess_options, providers=onnx_providers, ) logger.info( diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 41ef6cc..2ea10a7 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -137,8 +137,18 @@ class TTSModel: # ONNX 推論時 else: + sess_options = onnxruntime.SessionOptions() + # 基本的な最適化のみ有効化 + # ONNX モデルの作成時にすでに onnxsim により最適化されているため、ここでは基本的な最適化のみ有効化する + sess_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_BASIC + # エラー以外のログを出力しない + # 本来は log_severity_level = 3 だけで効くはずだが、なぜか抑制できないので set_default_logger_severity() も呼び出している + sess_options.log_severity_level = 3 + onnxruntime.set_default_logger_severity(3) + self.onnx_session = onnxruntime.InferenceSession( - path_or_bytes=str(self.model_path), + str(self.model_path), + sess_options=sess_options, providers=self.onnx_providers, ) logger.info(