//! Which runtime, which provider and which model form — decided once per //! device, and the only crate that knows the answer (docs/inference.md). //! //! Consumers ask for a session by [`Role`] and get `ort`'s `Session` back; //! what built it — tract on one core, ONNX Runtime's CPU pool, a TensorRT //! engine, the Hexagon — is this crate's business and shows up in //! [`status`] for the settings row and nowhere else. //! //! The shape follows §3 of the spec: `ort` links nothing (`alternative-backend`), //! and the first call hands it an API table from either a `libonnxruntime` //! found on disk or from `tract`. That choice is once per process, because //! `ort::set_api` is; everything after it — which provider, whether an engine //! has been compiled yet — is per session and may change between two calls. use std::collections::{BTreeSet, HashMap}; use std::path::{Path, PathBuf}; use std::sync::{Arc, Mutex, MutexGuard, OnceLock}; use std::time::{Duration, Instant}; use serde::{Deserialize, Serialize}; mod api; mod engines; mod probe; mod session; pub use api::Runtime; pub use ort::session::Session; /// What a model is for. The role fixes the precision rule (§7): an embedder /// runs in f32 on every rung, a detector may run in fp16 or int8. #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum Role { Detector, Embedder, Segmenter, Scene, /// The dense landmark model behind the eye reading (docs/faces.md §7c). Landmarks, /// The eye-state and sunglasses classifiers, a few hundred kilobytes. EyeClassifier, } /// Which numeric form of a model a session was built from. /// /// `Int8` is a different network from `F32` for a detector — it finds a /// different set of faces — which is why [`form_suffix`] exists and why a /// caller appends it to `model_id`. #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)] pub enum Form { F32, Int8, } /// A rung of the ladder (§2). Ordered: a user override names the highest rung /// the probe may take, and a compiling rung falls back to the one below it /// until its engine exists. #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)] pub enum Rung { /// ONNX Runtime's CPU provider, or tract when no runtime file was found. Cpu, /// NVIDIA, through the CUDA provider. Desktop only. Cuda, /// NVIDIA, through a TensorRT engine compiled on this device. Desktop only. TensorRt, /// Qualcomm's Hexagon NPU through QNN, int8 models only. Android only. Hexagon, } impl Rung { pub fn label(self) -> &'static str { match self { Rung::Cpu => "CPU", Rung::Cuda => "CUDA", Rung::TensorRt => "TensorRT", Rung::Hexagon => "Hexagon NPU", } } /// The rung a request lands on while this one's engine is still being /// compiled (§6 step 2). fn fallback(self) -> Rung { match self { Rung::TensorRt => Rung::Cuda, Rung::Hexagon | Rung::Cuda | Rung::Cpu => Rung::Cpu, } } /// Whether a session on this rung needs an engine built first. fn compiles(self) -> bool { matches!(self, Rung::TensorRt | Rung::Hexagon) } /// The model form this rung wants for a role. fn form(self, role: Role) -> Form { match (self, role) { (Rung::Hexagon, Role::Embedder) => Form::F32, (Rung::Hexagon, _) => Form::Int8, _ => Form::F32, } } } /// How long a session outlives its last use unless [`Config::decay`] says /// otherwise: long enough for the next click, short enough that a session's /// GPU or NPU memory does not sit under the develop view for long. pub const DEFAULT_DECAY: Duration = Duration::from_secs(30); /// What [`init`] is told once, at launch. #[derive(Clone, Debug, Default)] pub struct Config { /// Where to look for `libonnxruntime`, in order. An empty path means "the /// bare library name through the system loader", which is how the APK's /// own copy is found on Android. pub runtime_dirs: Vec, /// Probe cache and compiled engines (§4, §5). Disposable. pub cache_dir: PathBuf, /// The canonical model files on this device, so engines can be compiled /// ahead of the first request for them. pub models: Vec<(Role, PathBuf)>, /// Models compiled into the binary, for the same reason. pub embedded: Vec<(Role, &'static [u8])>, /// The highest rung the user allows; `None` is "the best that works". pub ceiling: Option, /// ONNX Runtime's intra-op pool; 0 picks from the core count. pub threads: usize, /// How long an unused session stays loaded. Zero means the default. pub decay: Duration, } /// One line for the settings row, and the numbers behind the progress row. #[derive(Clone, Debug)] pub struct Status { pub runtime: Runtime, /// The rung selected, or the floor while the probe is still running. pub rung: Rung, /// Why — "probe passed", or the failure that demoted the rung above. pub reason: String, pub probing: bool, /// Engines compiled and engines wanted, for a compiling rung; `(0, 0)` /// otherwise. pub engines: (usize, usize), } impl Status { /// "Hexagon NPU · int8 · ONNX Runtime 1.29" — the settings row's text. pub fn line(&self) -> String { let form = match self.rung { Rung::Hexagon => " · int8", Rung::TensorRt => " · fp16", _ => "", }; format!("{}{} · {}", self.rung.label(), form, self.runtime.label()) } } /// A model the caller can run, whatever is or is not loaded right now. /// /// Holds the bytes, not a session. [`Model::acquire`] finds the loaded copy /// in the registry — shared with every other holder of the same model — /// or loads one, and every acquire refreshes the copy's last-used time. /// The reaper unloads anything idle for [`Config::decay`]; a scan that runs /// the detector on every image never lets it go idle, a click in the /// develop view lets the segmenter go after a quiet spell, and a handle /// used again after that simply loads again. Nobody states a policy. /// /// The registry key includes the rung, so a reload after a compiled engine /// has landed moves up to it by itself (§6 step 4). pub struct Model { role: Role, form: Form, bytes: Arc<[u8]>, } /// A loaded session, held for one `run` and its output decoding. pub struct Acquired { entry: Arc, } struct Loaded { rung: Rung, session: Mutex, last_used: Mutex, } impl Model { /// The loaded session, loading it if the reaper took it. Lock it for /// one run; a scan and a develop click can want the same detector at /// once, and the second waits on the first. pub fn acquire(&self) -> Result { acquire(self.role, self.form, &self.bytes) } pub fn form(&self) -> Form { self.form } } impl Acquired { pub fn lock(&self) -> MutexGuard<'_, Session> { self.entry.session.lock().unwrap_or_else(|e| e.into_inner()) } /// Where this session runs. pub fn rung(&self) -> Rung { self.entry.rung } } impl Drop for Acquired { fn drop(&mut self) { // The clock starts when the use ends, not when it began: a long run // is not idle time. *self.entry.last_used.lock().unwrap() = Instant::now(); } } type Registry = HashMap>; static REGISTRY: OnceLock> = OnceLock::new(); fn registry() -> &'static Mutex { REGISTRY.get_or_init(|| { std::thread::Builder::new() .name("inference-reaper".into()) .spawn(|| loop { std::thread::sleep(Duration::from_secs(5)); release_idle(); }) .expect("spawn inference reaper"); Mutex::new(HashMap::new()) }) } fn acquire(role: Role, form: Form, bytes: &Arc<[u8]>) -> Result { api::ensure_installed(); let (rung, cfg) = { let s = state().lock().unwrap(); let selected = current_rung(&s); ( effective_rung(&s, selected, role, form, bytes), s.config.clone(), ) }; let key = format!("{role:?}:{}", engines::key(rung, bytes)); if let Some(entry) = registry().lock().unwrap().get(&key).cloned() { *entry.last_used.lock().unwrap() = Instant::now(); return Ok(Acquired { entry }); } // Built outside the registry lock: a TensorRT engine load is long enough // that another role's acquire should not wait on it. let session = session::build(rung, role, bytes, &cfg, false)?; log::debug!("inference: {role:?} loaded on {}", rung.label()); let entry = Arc::new(Loaded { rung, session: Mutex::new(session), last_used: Mutex::new(Instant::now()), }); let mut reg = registry().lock().unwrap(); // Two acquires raced; keep the first, drop this one. let entry = reg.entry(key).or_insert_with(|| entry.clone()).clone(); Ok(Acquired { entry }) } /// Unload every session idle for longer than the decay. The reaper does /// this every five seconds. A session in use survives until its run ends: /// the `Acquired` holds it, the registry merely forgets it. pub fn release_idle() { let decay = match state().lock().unwrap().config.decay { Duration::ZERO => DEFAULT_DECAY, d => d, }; let now = Instant::now(); registry() .lock() .unwrap() .retain(|_, e| now.duration_since(*e.last_used.lock().unwrap()) < decay); } /// Unload every session now, decay or not — what a low-memory signal /// asks for. Sessions mid-run finish first. pub fn release_all() { registry().lock().unwrap().clear(); } /// Unload every session of `role` now — "I am done segmenting". pub fn unload(role: Role) { let prefix = format!("{role:?}:"); registry() .lock() .unwrap() .retain(|k, _| !k.starts_with(&prefix)); } /// How many sessions are loaded, for the settings row and the tests. pub fn loaded() -> usize { registry().lock().unwrap().len() } #[derive(Debug, thiserror::Error)] pub enum Error { #[error(transparent)] Inference(#[from] ort::Error), #[error("reading model: {0}")] Io(#[from] std::io::Error), } /// What the probe writes and the next launch reads (§4 step 3). #[derive(Clone, Debug, Default, Serialize, Deserialize)] struct Cache { /// Runtime, driver, hardware and model identity; any change re-probes. fingerprint: String, rung: Option, reason: String, /// Model hashes whose engine exists on disk, per compiling rung. compiled: BTreeSet, /// Rungs that failed under this fingerprint, and why. Not retried until /// the fingerprint changes: a wedged driver must not cost every launch /// thirty seconds. failed: Vec<(Rung, String)>, } struct State { config: Config, cache: Cache, probing: bool, wanted: usize, } static STATE: OnceLock> = OnceLock::new(); fn state() -> &'static Mutex { STATE.get_or_init(|| { Mutex::new(State { config: Config::default(), cache: Cache::default(), probing: false, wanted: 0, }) }) } /// Choose the runtime and start the probe. Idempotent; the first call wins. /// /// Returns at once: the probe and any engine compilation run on their own /// low-priority thread, and every request meanwhile is served by the floor /// (§4). Never blocks the first frame. pub fn init(config: Config) { let runtime = api::install(&config.runtime_dirs); { let mut s = state().lock().unwrap(); if s.probing || s.cache.rung.is_some() { return; } s.config = config; s.probing = true; } log::info!("inference: runtime {}", runtime.label()); std::thread::Builder::new() .name("inference-probe".into()) .spawn(move || { probe::run(runtime); engines::run(); }) .expect("spawn inference probe"); } /// Make sure `ort` has an API table, for code that drives `ort` directly. /// [`open`] does this itself; only the M1 probe example needs it by name. pub fn ensure_runtime() { api::ensure_installed(); } /// The line for the settings row. pub fn status() -> Status { let s = state().lock().unwrap(); let rung = current_rung(&s); Status { runtime: api::runtime(), rung, reason: s.cache.reason.clone(), probing: s.probing, engines: if rung.compiles() { (s.cache.compiled.len(), s.wanted) } else { (0, 0) }, } } fn current_rung(s: &State) -> Rung { if s.probing { Rung::Cpu } else { s.cache.rung.unwrap_or(Rung::Cpu) } } /// The file to load for `role` under the current selection, and its form. /// /// A rung that wants int8 gets the `.int8.onnx` sibling of the canonical file /// if it exists; otherwise the canonical file, on the rung's fallback. A /// caller adds [`form_suffix`] to the `model_id` it records. pub fn resolve_model(role: Role, canonical: &Path) -> (PathBuf, Form) { let rung = current_rung(&state().lock().unwrap()); if rung.form(role) == Form::Int8 { let sibling = int8_sibling(canonical); if sibling.is_file() { return (sibling, Form::Int8); } } (canonical.to_path_buf(), Form::F32) } fn int8_sibling(canonical: &Path) -> PathBuf { let stem = canonical .file_stem() .map(|s| s.to_string_lossy().into_owned()) .unwrap_or_default(); canonical.with_file_name(format!("{stem}.int8.onnx")) } /// What a form appends to a detector's `model_id` (§7). pub fn form_suffix(form: Form) -> &'static str { match form { Form::F32 => "", Form::Int8 => "_i8", } } /// A handle on the model `bytes` in `role`. /// /// Loads it once here, so a graph the runtime rejects fails at /// construction and not on the first image; what happens to that session /// afterwards is the registry's business (see [`Model`]). /// /// Works without [`init`] — a test, or the examples — by installing tract /// and using the CPU rung, which is exactly what every consumer did before /// this crate existed. pub fn open(role: Role, form: Form, bytes: &[u8]) -> Result { let bytes: Arc<[u8]> = Arc::from(bytes); acquire(role, form, &bytes)?; Ok(Model { role, form, bytes }) } /// Where a request lands: the selected rung unless the role's precision rule, /// the form on offer, or a missing engine says one lower (§6 step 4). fn effective_rung(s: &State, selected: Rung, role: Role, form: Form, bytes: &[u8]) -> Rung { let mut rung = selected; if rung.form(role) != form { // The embedder on a Hexagon device, or an f32 detector where the int8 // sibling was missing: neither can go to the NPU. rung = rung.fallback(); } if rung.compiles() && !s.cache.compiled.contains(&engines::key(rung, bytes)) { rung = rung.fallback(); } rung } #[cfg(test)] mod tests { use super::*; /// The registry is one per process, so these run one at a time. static SERIAL: Mutex<()> = Mutex::new(()); fn serial() -> MutexGuard<'static, ()> { SERIAL.lock().unwrap_or_else(|e| e.into_inner()) } /// The smallest shipped graph, if this checkout has the weights; a test /// suite that needs a research-licensed download is one that does not /// run in CI (docs/faces.md §3), so absence is a skip. fn probe_bytes() -> Option> { let path = concat!( env!("CARGO_MANIFEST_DIR"), "/../../models/face/scrfd_500m_640.onnx" ); let bytes = std::fs::read(path).ok()?; (bytes.len() > 100_000).then_some(bytes) } #[test] fn two_handles_on_one_model_share_one_session() { let _serial = serial(); let Some(bytes) = probe_bytes() else { return }; release_all(); let a = open(Role::Detector, Form::F32, &bytes).unwrap(); let b = open(Role::Detector, Form::F32, &bytes).unwrap(); assert_eq!(loaded(), 1); let (x, y) = (a.acquire().unwrap(), b.acquire().unwrap()); assert!(Arc::ptr_eq(&x.entry, &y.entry)); } #[test] fn a_released_model_reloads_on_its_next_use() { let _serial = serial(); let Some(bytes) = probe_bytes() else { return }; release_all(); let model = open(Role::Detector, Form::F32, &bytes).unwrap(); assert_eq!(loaded(), 1); release_all(); assert_eq!(loaded(), 0); let acquired = model.acquire().unwrap(); assert_eq!(loaded(), 1); assert_eq!(acquired.lock().inputs().len(), 1); } #[test] fn an_idle_session_decays_and_a_used_one_does_not() { let _serial = serial(); let Some(bytes) = probe_bytes() else { return }; release_all(); state().lock().unwrap().config.decay = Duration::from_millis(50); let model = open(Role::Detector, Form::F32, &bytes).unwrap(); // Used within the decay: stays. std::thread::sleep(Duration::from_millis(30)); drop(model.acquire().unwrap()); release_idle(); assert_eq!(loaded(), 1); // Idle past it: goes. std::thread::sleep(Duration::from_millis(80)); release_idle(); assert_eq!(loaded(), 0); state().lock().unwrap().config.decay = Duration::ZERO; } #[test] fn unload_by_role_leaves_the_other_roles() { let _serial = serial(); let Some(bytes) = probe_bytes() else { return }; release_all(); let _d = open(Role::Detector, Form::F32, &bytes).unwrap(); let _s = open(Role::Segmenter, Form::F32, &bytes).unwrap(); assert_eq!(loaded(), 2); unload(Role::Segmenter); assert_eq!(loaded(), 1); } #[test] fn the_status_line_reads_as_the_floor_before_init() { let s = status(); assert_eq!(s.rung, Rung::Cpu); assert!(s.line().starts_with("CPU"), "{}", s.line()); } }