Refactor: Use I/O Binding during BERT inference and always release memory after inference

This commit is contained in:
tsukumi
2024-12-22 06:49:43 +09:00
parent 15af441cba
commit 08c439e88e
5 changed files with 223 additions and 78 deletions

View File

@@ -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]