"""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)