face-intel / cores /onnx /session.py
Marwan
Restructure + add reverse face search (PimEyes-style)
f5eeb1c
Raw
History Blame Contribute Delete
3.44 kB
"""ONNX Runtime session management — load-once, reuse-many."""
from __future__ import annotations
import threading
from pathlib import Path
from typing import Dict, Optional
import numpy as np
from config.settings import settings as _default_settings
from cores.embedding.cache import EmbeddingCache
# --------------------------------------------------------------------------- #
# Availability check
# --------------------------------------------------------------------------- #
def is_onnx_available() -> bool:
"""True if onnxruntime is importable."""
try:
import onnxruntime # noqa: F401
return True
except ImportError:
return False
# --------------------------------------------------------------------------- #
# Session cache — one session per model path
# --------------------------------------------------------------------------- #
_session_cache: EmbeddingCache = EmbeddingCache()
_init_lock = threading.Lock()
class ONNXModel:
"""Wrapper around an ONNX inference session.
Encapsulates the session + input/output name discovery so providers
don't repeat that boilerplate.
"""
def __init__(self, session) -> None:
self._session = session
self._input_name = session.get_inputs()[0].name
self._output_names = [o.name for o in session.get_outputs()]
@property
def input_name(self) -> str:
return self._input_name
@property
def output_names(self) -> list[str]:
return list(self._output_names)
def run(self, input_array: np.ndarray) -> list[np.ndarray]:
"""Run inference. Returns list of output arrays."""
return self._session.run(
self._output_names,
{self._input_name: input_array},
)
def run_single(self, input_array: np.ndarray) -> np.ndarray:
"""Run inference + return only the first output."""
return self.run(input_array)[0]
def get_session(model_path: Path | str, settings=None) -> ONNXModel:
"""Load (or return cached) ONNXModel for the given model file.
The session is created ONCE per process and reused. Thread-safe.
Args:
model_path: path to .onnx file
settings: Settings (uses default singleton if None)
Returns:
ONNXModel wrapper
"""
s = settings or _default_settings
key = str(model_path)
return _session_cache.get_or_load(key, lambda: _create_session(model_path, s))
def _create_session(model_path: Path | str, settings) -> ONNXModel:
"""Create a new ONNX Runtime session."""
import onnxruntime as ort
path = Path(model_path)
if not path.exists():
raise FileNotFoundError(f"ONNX model not found: {path}")
# CPU-only, limited threads for low-RAM deployments
opts = ort.SessionOptions()
opts.intra_op_num_threads = settings.onnx_intra_op_threads
opts.inter_op_num_threads = 1
opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
session = ort.InferenceSession(
str(path),
sess_options=opts,
providers=["CPUExecutionProvider"], # CPU-only by design
)
return ONNXModel(session)
def run_inference(model_path: Path | str, input_array: np.ndarray,
settings=None) -> list[np.ndarray]:
"""Convenience: load (or get cached) session + run inference."""
model = get_session(model_path, settings)
return model.run(input_array)