mirror of
https://github.com/immich-app/immich
synced 2026-08-15 13:03:57 +00:00
Merge 1f9d6b0b0d into ffc83eae36
This commit is contained in:
commit
2d9db31de9
2 changed files with 82 additions and 0 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -29,6 +29,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
|
||||
|
|
@ -329,6 +330,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"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue