fix(ml): register CUDA EP library to avoid plugin EP lookup warning

Since ONNX Runtime 1.26, session creation first looks for a registered
plugin EP device matching the provider name and its `device_id` option,
and only falls back to the built-in provider afterwards. Only the CPU,
WebGPU and DML providers are registered internally, so with CUDA the
lookup always fails and every session logs:

    No registered plugin EP device found for 'CUDAExecutionProvider' with device_id=0

The CUDA provider library exposes an OrtEpFactory, so registering it
makes the device discoverable and the lookup succeeds. Provider options
are still applied to the session, and the built-in provider remains the
fallback if the library is missing or registration fails.

Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
volschin 2026-08-03 20:25:16 +00:00
parent 29e7ea5302
commit 1f9d6b0b0d
2 changed files with 82 additions and 0 deletions

View file

@ -15,11 +15,56 @@ from ..config import log, settings
MigraphxInputSignature = tuple[tuple[str, str, tuple[int, ...]], ...]
# Provider libraries that expose an OrtEpFactory and are looked up as plugin EP devices
# before the built-in provider is used. See _register_ep_library().
_EP_LIBRARIES = {"CUDAExecutionProvider": "libonnxruntime_providers_cuda.so"}
_ep_library_lock = Lock()
_registered_ep_libraries: set[str] = set()
_migraphx_registry_lock = Lock()
_migraphx_model_locks: dict[str, Lock] = {}
_migraphx_compiled_inputs: set[tuple[str, MigraphxInputSignature]] = set()
def _register_ep_library(provider: str) -> None:
"""Register the provider's shared library as a plugin EP library.
Since ORT 1.23 a session first looks for a registered plugin EP device matching the
provider name and its `device_id` option, and only then falls back to the built-in
provider. Without a registered library that lookup logs a warning on every session:
No registered plugin EP device found for 'CUDAExecutionProvider' with device_id=0
Registering the library makes the device discoverable, so the lookup succeeds and the
provider options are applied to the session as usual.
"""
library = _EP_LIBRARIES.get(provider)
register = getattr(ort, "register_execution_provider_library", None)
if library is None or register is None:
return
with _ep_library_lock:
if provider in _registered_ep_libraries:
return
# Mark as attempted regardless of the outcome: a failure here is not fatal, the
# provider is still created from the built-in implementation, and retrying it for
# every session is pointless.
_registered_ep_libraries.add(provider)
try:
library_path = Path(ort.__file__).parent / "capi" / library
if not library_path.is_file():
log.debug(f"EP library {library} not found, skipping plugin EP registration for {provider}")
return
register(provider, library_path.as_posix())
log.debug(f"Registered plugin EP library {library_path} for {provider}")
except Exception as e:
log.debug(f"Could not register plugin EP library {library} for {provider}: {e}")
def _migraphx_get_model_lock(model_key: str) -> Lock:
with _migraphx_registry_lock:
lock = _migraphx_model_locks.get(model_key)
@ -59,6 +104,8 @@ class OrtSession:
self.providers = providers if providers is not None else self._providers_default
self.provider_options = provider_options if provider_options is not None else self._provider_options_default
self.sess_options = sess_options if sess_options is not None else self._sess_options_default
for provider in self.providers:
_register_ep_library(provider)
self.session = ort.InferenceSession(
self.model_path.as_posix(),
providers=self.providers,

View file

@ -30,6 +30,7 @@ from immich_ml.models.ocr.detection import TextDetector
from immich_ml.models.ocr.recognition import TextRecognizer
from immich_ml.models.ocr.schemas import OcrOptions
from immich_ml.schemas import ModelFormat, ModelPrecision, ModelTask, ModelType
from immich_ml.sessions import ort as ort_module
from immich_ml.sessions.ann import AnnSession
from immich_ml.sessions.ort import OrtSession
from immich_ml.sessions.rknn import RknnSession, run_inference
@ -330,6 +331,40 @@ class TestOrtSession:
assert session.provider_options == [{"arena_extend_strategy": "kSameAsRequested", "device_id": "1"}]
def test_registers_cuda_ep_library(self, mocker: MockerFixture) -> None:
mocker.patch.dict(ort_module._registered_ep_libraries, clear=True)
mocker.patch("immich_ml.sessions.ort.Path.is_file", return_value=True)
register = mocker.patch("immich_ml.sessions.ort.ort.register_execution_provider_library")
OrtSession("ViT-B-32__openai", providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
OrtSession("ViT-B-32__openai", providers=["CUDAExecutionProvider", "CPUExecutionProvider"])
register.assert_called_once()
provider, library_path = register.call_args[0]
assert provider == "CUDAExecutionProvider"
assert library_path.endswith("/capi/libonnxruntime_providers_cuda.so")
def test_does_not_register_cuda_ep_library_if_missing(self, mocker: MockerFixture) -> None:
mocker.patch.dict(ort_module._registered_ep_libraries, clear=True)
mocker.patch("immich_ml.sessions.ort.Path.is_file", return_value=False)
register = mocker.patch("immich_ml.sessions.ort.ort.register_execution_provider_library")
OrtSession("ViT-B-32__openai", providers=["CUDAExecutionProvider"])
register.assert_not_called()
def test_ignores_cuda_ep_library_registration_failure(self, mocker: MockerFixture) -> None:
mocker.patch.dict(ort_module._registered_ep_libraries, clear=True)
mocker.patch("immich_ml.sessions.ort.Path.is_file", return_value=True)
mocker.patch(
"immich_ml.sessions.ort.ort.register_execution_provider_library",
side_effect=RuntimeError("registration failed"),
)
session = OrtSession("ViT-B-32__openai", providers=["CUDAExecutionProvider"])
assert session.providers == ["CUDAExecutionProvider"]
def test_sets_provider_options_for_rocm(self, mocker: MockerFixture) -> None:
model_path = "/cache/ViT-B-32__openai/textual/model.onnx"
os.environ["MACHINE_LEARNING_DEVICE_ID"] = "1"