| """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 |
|
|
|
|
| |
| |
| |
| def is_onnx_available() -> bool: |
| """True if onnxruntime is importable.""" |
| try: |
| import onnxruntime |
| return True |
| except ImportError: |
| return False |
|
|
|
|
| |
| |
| |
| _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}") |
|
|
| |
| 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"], |
| ) |
| 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) |
|
|