Fix: ONNX inference with DirectML fails
This commit is contained in:
@@ -187,16 +187,29 @@ def infer_onnx(
|
|||||||
np.array([noise_scale_w], dtype=np.float32),
|
np.array([noise_scale_w], dtype=np.float32),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
# 入力テンソルを転送する GPU デバイスを取得
|
||||||
|
## 本来は device_type="dml" もサポートされているはずだが、手元環境だと常に謎の RuntimeError が発生するため当面無効化している
|
||||||
first_provider = onnx_session.get_providers()[0]
|
first_provider = onnx_session.get_providers()[0]
|
||||||
if first_provider == "CUDAExecutionProvider":
|
if first_provider == "CUDAExecutionProvider":
|
||||||
device_type = "cuda"
|
device_type = "cuda"
|
||||||
elif first_provider == "DmlExecutionProvider":
|
# elif first_provider == "DmlExecutionProvider":
|
||||||
device_type = "dml"
|
# device_type = "dml"
|
||||||
else:
|
else:
|
||||||
device_type = "cpu"
|
device_type = "cpu"
|
||||||
|
|
||||||
|
# 入力テンソルを転送する GPU デバイスの ID を取得
|
||||||
|
## ExecutionProvider に指定したオプションの中から device_id を取得し、入力テンソルの転送先として指定する
|
||||||
|
## InferenceSession で利用するデバイス ID と入力テンソルの転送先デバイス ID は一致している必要がある
|
||||||
|
## 本来は ExecutionProvider に指定したオプションは InferenceSession.get_provider_options() で取得できるはずだが、
|
||||||
|
## 手元環境だと DmlExecutionProvider のみ常に空の辞書が返されるため、当面 onnx_providers から直接オプションを取り出している
|
||||||
device_id = 0
|
device_id = 0
|
||||||
first_provider_options = onnx_session.get_provider_options()[first_provider]
|
onnx_providers_dict: dict[str, dict[str, Any]] = {}
|
||||||
|
for provider, options in onnx_providers:
|
||||||
|
if isinstance(options, dict):
|
||||||
|
onnx_providers_dict[provider] = options
|
||||||
|
else:
|
||||||
|
onnx_providers_dict[provider] = {}
|
||||||
|
first_provider_options = onnx_providers_dict[first_provider]
|
||||||
if "device_id" in first_provider_options:
|
if "device_id" in first_provider_options:
|
||||||
device_id = int(first_provider_options["device_id"])
|
device_id = int(first_provider_options["device_id"])
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user