Files
immich/machine-learning/conftest.py
Mert 76c239c345 feat(ml): trt-rtx (#31863)
* trt-rtx

* update dependencies
2026-09-28 21:39:14 +00:00

229 lines
7.3 KiB
Python

import json
from pathlib import Path
from types import SimpleNamespace
from typing import Any, Callable, Iterator
from unittest import mock
import pytest
from fastapi.testclient import TestClient
from PIL import Image
from immich_ml.config import log
from immich_ml.main import app
from immich_ml.models.base import InferenceModel
from immich_ml.schemas import Shape
from immich_ml.sessions.ort import Device
TEST_ASSETS = Path(__file__).parent.parent / "e2e/test-assets"
@pytest.fixture(scope="session")
def asset() -> Callable[[str], bytes]:
return lambda path: (TEST_ASSETS / path).read_bytes()
@pytest.fixture
def pil_image() -> Image.Image:
return Image.new("RGB", (600, 800))
@pytest.fixture(scope="session")
def deployed_app() -> Iterator[TestClient]:
with TestClient(app) as client:
yield client
@pytest.fixture(scope="session")
def responses() -> dict[str, Any]:
responses: dict[str, Any] = json.load(open("responses.json", "r"))
return responses
@pytest.fixture(scope="session")
def clip_model_cfg() -> dict[str, Any]:
return {
"embed_dim": 512,
"vision_cfg": {"image_size": 224, "layers": 12, "width": 768, "patch_size": 32},
"text_cfg": {"context_length": 77, "vocab_size": 49408, "width": 512, "heads": 8, "layers": 12},
}
@pytest.fixture(scope="session")
def clip_preprocess_cfg() -> dict[str, Any]:
return {
"size": [224, 224],
"mode": "RGB",
"mean": [0.48145466, 0.4578275, 0.40821073],
"std": [0.26862954, 0.26130258, 0.27577711],
"interpolation": "bicubic",
"resize_mode": "shortest",
"fill_color": 0,
}
@pytest.fixture(scope="session")
def clip_tokenizer_cfg() -> dict[str, Any]:
return {
"add_prefix_space": False,
"added_tokens_decoder": {
"49406": {
"content": "<|startoftext|>",
"lstrip": False,
"normalized": True,
"rstrip": False,
"single_word": False,
"special": True,
},
"49407": {
"content": "<|endoftext|>",
"lstrip": False,
"normalized": True,
"rstrip": False,
"single_word": False,
"special": True,
},
},
"bos_token": "<|startoftext|>",
"clean_up_tokenization_spaces": True,
"do_lower_case": True,
"eos_token": "<|endoftext|>",
"errors": "replace",
"model_max_length": 77,
"pad_token": "<|endoftext|>",
"tokenizer_class": "CLIPTokenizer",
"unk_token": "<|endoftext|>",
}
@pytest.fixture(scope="function")
def providers(request: pytest.FixtureRequest) -> Iterator[mock.Mock]:
marker = request.node.get_closest_marker("providers")
if marker is None:
raise ValueError("Missing marker 'providers'")
providers = marker.args[0]
with mock.patch("immich_ml.sessions.ort.ort.get_available_providers") as mocked:
mocked.return_value = providers
yield providers
@pytest.fixture(autouse=True)
def gpus() -> Iterator[None]:
with (
mock.patch("immich_ml.sessions.ort._intel_gpu", return_value=Device("12.71.4-128eu", "26.22.38646.4")),
mock.patch("immich_ml.sessions.ort._amd_gpu", return_value=Device("gfx1100", "7.2.0")),
mock.patch("immich_ml.sessions.ort._nvidia_gpu", return_value=Device("sm89", "617.14 10601")),
):
yield
@pytest.fixture(scope="function")
def ort_pybind() -> Iterator[mock.Mock]:
with mock.patch("immich_ml.sessions.ort.ort.capi._pybind_state") as mocked:
yield mocked
@pytest.fixture(scope="function")
def ov_device_ids(request: pytest.FixtureRequest, ort_pybind: mock.Mock) -> Iterator[mock.Mock]:
marker = request.node.get_closest_marker("ov_device_ids")
if marker is None:
raise ValueError("Missing marker 'ov_device_ids'")
ort_pybind.get_available_openvino_device_ids.return_value = marker.args[0]
return ort_pybind
@pytest.fixture(scope="function")
def ort_session() -> Iterator[mock.Mock]:
# a graph is opened as it is: preparing one takes a child process and real files
with (
mock.patch("immich_ml.sessions.ort.ort.InferenceSession") as mocked,
mock.patch("immich_ml.sessions.ort.prepared", side_effect=lambda spec: spec.model_path),
):
yield mocked
@pytest.fixture(scope="function")
def stub_session() -> Callable[..., mock.Mock]:
def _make(
shape: tuple[Any, ...],
outputs: Any = None,
name: str = "input.1",
normalizes_input: bool = False,
shapes: tuple[Shape, ...] = (Shape(batch=1),),
) -> mock.Mock:
session = mock.Mock()
session.get_inputs.return_value = [SimpleNamespace(name=name, shape=shape)]
session.normalizes_input = normalizes_input
session.shapes = shapes
session.batches = tuple(sorted({shape.batch for shape in shapes}, reverse=True))
session.for_shape.return_value = session # a stub is the session and the one graph in it
if outputs is not None:
session.run.return_value = outputs
return session
return _make
@pytest.fixture(scope="function")
def ann_session() -> Iterator[mock.Mock]:
with mock.patch("immich_ml.sessions.ann.Ann") as mocked:
mocked.return_value.input_shapes.__getitem__.return_value = [(1, 3, 112, 112)] # armnn compiles one batch
yield mocked
@pytest.fixture(scope="function")
def rknn_session() -> Iterator[mock.Mock]:
with mock.patch("immich_ml.sessions.rknn.RknnPoolExecutor") as mocked:
mocked.return_value.custom_string = "" # an unstamped binary, i.e. one that wants normalized input
compiled = SimpleNamespace(name="input", shape=(1, 224, 224, 3)) # the compiler reports NHWC either way
mocked.return_value.inputs = [compiled, compiled]
yield mocked
@pytest.fixture(scope="function")
def rmtree() -> Iterator[mock.Mock]:
with mock.patch("immich_ml.models.base.rmtree", autospec=True) as mocked:
mocked.avoids_symlink_attacks = True
yield mocked
@pytest.fixture(scope="function")
def path() -> Iterator[mock.Mock]:
path = mock.MagicMock()
path.exists.return_value = True
path.is_dir.return_value = True
path.is_file.return_value = True
path.with_suffix.return_value = path
path.return_value = path
with mock.patch("immich_ml.models.base.Path", return_value=path) as mocked:
yield mocked
@pytest.fixture(scope="function")
def info() -> Iterator[mock.Mock]:
with mock.patch.object(log, "info") as mocked:
yield mocked
@pytest.fixture(scope="function")
def warning() -> Iterator[mock.Mock]:
with mock.patch.object(log, "warning") as mocked:
yield mocked
@pytest.fixture(scope="function")
def exception() -> Iterator[mock.Mock]:
with mock.patch.object(log, "exception") as mocked:
yield mocked
@pytest.fixture(scope="function")
def snapshot_download() -> Iterator[mock.Mock]:
# a download that leaves the model cached, as one that found the artifact does
with (
mock.patch("immich_ml.models.base.snapshot_download") as mocked,
mock.patch.object(InferenceModel, "cached", new_callable=mock.PropertyMock, side_effect=lambda: mocked.called),
):
yield mocked