Refactor: Use I/O Binding during BERT inference and always release memory after inference
This commit is contained in:
@@ -11,6 +11,7 @@ from style_bert_vits2.nlp import (
|
||||
cleaned_text_to_sequence,
|
||||
extract_bert_feature_onnx,
|
||||
)
|
||||
from style_bert_vits2.utils import get_onnx_device_options
|
||||
|
||||
|
||||
def __intersperse(lst: list[Any], item: Any) -> list[Any]:
|
||||
@@ -187,34 +188,11 @@ def infer_onnx(
|
||||
np.array(noise_scale_w, dtype=np.float32),
|
||||
]
|
||||
|
||||
# 入力テンソルを転送する GPU デバイスを取得
|
||||
## 本来は device_type="dml" もサポートされているはずだが、手元環境だと常に謎の RuntimeError が発生するため当面無効化している
|
||||
first_provider = onnx_session.get_providers()[0]
|
||||
if first_provider == "CUDAExecutionProvider":
|
||||
device_type = "cuda"
|
||||
# elif first_provider == "DmlExecutionProvider":
|
||||
# device_type = "dml"
|
||||
else:
|
||||
device_type = "cpu"
|
||||
# 入力テンソルの転送に使用するデバイス種別, デバイス ID, 実行オプションを取得
|
||||
device_type, device_id, run_options = get_onnx_device_options(onnx_session, onnx_providers) # fmt: skip
|
||||
|
||||
# 入力テンソルを転送する GPU デバイスの ID を取得
|
||||
## ExecutionProvider に指定したオプションの中から device_id を取得し、入力テンソルの転送先として指定する
|
||||
## InferenceSession で利用するデバイス ID と入力テンソルの転送先デバイス ID は一致している必要がある
|
||||
## 本来は ExecutionProvider に指定したオプションは InferenceSession.get_provider_options() で取得できるはずだが、
|
||||
## 手元環境だと DmlExecutionProvider のみ常に空の辞書が返されるため、当面 onnx_providers から直接オプションを取り出している
|
||||
device_id = 0
|
||||
onnx_providers_dict: dict[str, dict[str, Any]] = {}
|
||||
for provider in onnx_providers:
|
||||
if isinstance(provider, tuple):
|
||||
provider_name, options = provider
|
||||
onnx_providers_dict[provider_name] = options
|
||||
else:
|
||||
onnx_providers_dict[provider] = {}
|
||||
first_provider_options = onnx_providers_dict[first_provider]
|
||||
if "device_id" in first_provider_options:
|
||||
device_id = int(first_provider_options["device_id"])
|
||||
|
||||
# GPU メモリに入力テンソルを割り当て
|
||||
# 推論デバイスに入力テンソルを割り当て
|
||||
## GPU 推論の場合、device_type + device_id に対応する GPU デバイスに入力テンソルが割り当てられる
|
||||
io_binding = onnx_session.io_binding()
|
||||
for name, value in zip(input_names, input_tensor):
|
||||
gpu_tensor = onnxruntime.OrtValue.ortvalue_from_numpy(
|
||||
@@ -224,7 +202,7 @@ def infer_onnx(
|
||||
|
||||
# 推論の実行
|
||||
io_binding.bind_output(output_name, device_type)
|
||||
onnx_session.run_with_iobinding(io_binding)
|
||||
onnx_session.run_with_iobinding(io_binding, run_options=run_options)
|
||||
output = io_binding.get_outputs()
|
||||
|
||||
audio = output[0].numpy()[0, 0]
|
||||
|
||||
Reference in New Issue
Block a user