From d15c41e69966e4679edfded1d8b249a1f0716e31 Mon Sep 17 00:00:00 2001 From: Duncan Tourolle Date: Sat, 19 Sep 2026 14:15:56 +0200 Subject: [PATCH] Add dr-inference-engine and route every model session through it MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit One crate names the runtime, the providers and the devices; dr-face and dr-segment ask it for a session by role. It hands ort an API table once per process — from a libonnxruntime it dlopens when the app names a directory holding one, otherwise from tract — so the Rust build stays free of C on every target and a package can install the runtime as a file (docs/inference.md §3). Sessions live in a registry behind a Model handle that holds the bytes, not the session: every use refreshes a timestamp and a reaper unloads whatever sat idle past the decay. A scan that runs the detector on each image never lets it go idle; a click in the develop view lets the segmenter go after thirty seconds; a handle used after that reloads, and reloads on a higher rung if a compiled engine has landed meanwhile. The probe walks the platform's ladder by building strict sessions and timing them against the CPU provider, caches the choice against a fingerprint of the runtime, driver, hardware and models, and compiles engines for the selected rung in the background, smallest model first. Nothing in this commit turns the native path on: the apps still run on tract until they call init with a runtime directory. --- Cargo.lock | 18 +- Cargo.toml | 4 + core/dr-face/Cargo.toml | 11 +- core/dr-face/src/detect.rs | 50 ++- core/dr-face/src/embed.rs | 28 +- core/dr-face/src/lib.rs | 26 +- core/dr-inference-engine/Cargo.toml | 44 ++ core/dr-inference-engine/src/api.rs | 142 +++++++ core/dr-inference-engine/src/engines.rs | 104 +++++ core/dr-inference-engine/src/lib.rs | 542 ++++++++++++++++++++++++ core/dr-inference-engine/src/probe.rs | 279 ++++++++++++ core/dr-inference-engine/src/session.rs | 129 ++++++ core/dr-segment/Cargo.toml | 9 +- core/dr-segment/src/lib.rs | 10 + core/dr-segment/src/scene.rs | 26 +- core/dr-segment/src/semantic.rs | 36 +- 16 files changed, 1368 insertions(+), 90 deletions(-) create mode 100644 core/dr-inference-engine/Cargo.toml create mode 100644 core/dr-inference-engine/src/api.rs create mode 100644 core/dr-inference-engine/src/engines.rs create mode 100644 core/dr-inference-engine/src/lib.rs create mode 100644 core/dr-inference-engine/src/probe.rs create mode 100644 core/dr-inference-engine/src/session.rs diff --git a/Cargo.lock b/Cargo.lock index 490835d..67edfcc 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1475,11 +1475,11 @@ dependencies = [ name = "dr-face" version = "0.12.2" dependencies = [ + "dr-inference-engine", "env_logger", "log", "ndarray", "ort", - "ort-tract", "thiserror 2.0.20", "zune-jpeg 0.4.21", ] @@ -1511,6 +1511,20 @@ dependencies = [ "wgpu", ] +[[package]] +name = "dr-inference-engine" +version = "0.12.2" +dependencies = [ + "libloading", + "log", + "ort", + "ort-sys", + "ort-tract", + "serde", + "serde_json", + "thiserror 2.0.20", +] + [[package]] name = "dr-ingest" version = "0.12.2" @@ -1584,11 +1598,11 @@ dependencies = [ name = "dr-segment" version = "0.12.2" dependencies = [ + "dr-inference-engine", "env_logger", "log", "ndarray", "ort", - "ort-tract", "thiserror 2.0.20", "zune-jpeg 0.4.21", ] diff --git a/Cargo.toml b/Cargo.toml index 7b3ae01..52f66a7 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -8,6 +8,7 @@ members = [ "core/dr-export", "core/dr-face", "core/dr-film", + "core/dr-inference-engine", "core/dr-ingest", "core/dr-gpu", "core/dr-lens", @@ -46,6 +47,9 @@ dr-export = { path = "core/dr-export" } # `features = ["inference"]`. dr-face = { path = "core/dr-face", default-features = false } dr-film = { path = "core/dr-film" } +# `tract` on by default so a test binary can open a session with nothing +# installed; the apps add `native` to look for a runtime file (docs/inference.md §3). +dr-inference-engine = { path = "core/dr-inference-engine" } dr-ingest = { path = "core/dr-ingest" } dr-gpu = { path = "core/dr-gpu" } dr-lens = { path = "core/dr-lens" } diff --git a/core/dr-face/Cargo.toml b/core/dr-face/Cargo.toml index 27e5ae0..4e32c8b 100644 --- a/core/dr-face/Cargo.toml +++ b/core/dr-face/Cargo.toml @@ -9,10 +9,11 @@ license.workspace = true thiserror.workspace = true log.workspace = true -# Inference. `ort` is the API; **tract is the engine** — see the workspace -# manifest, and docs/faces.md §3, for why the C++ ONNX Runtime is not linked. +# Inference. `ort` is the API; **what runs it is `dr-inference-engine`'s +# business** — tract, or an ONNX Runtime the app found on disk, on whichever +# provider the device has (docs/inference.md). This crate never names either. ort = { workspace = true, optional = true } -ort-tract = { workspace = true, optional = true } +dr-inference-engine = { workspace = true, optional = true } ndarray = { workspace = true, optional = true } [dev-dependencies] @@ -20,7 +21,7 @@ zune-jpeg.workspace = true env_logger.workspace = true # The M1 probe drives `ort` directly so it can print the raw load error. ort = { workspace = true } -ort-tract = { workspace = true } +dr-inference-engine = { workspace = true } [[example]] name = "probe" @@ -48,4 +49,4 @@ default = [] # must be testable against synthetic embeddings on a machine with no weights on # it — a test suite that needs a research-licensed download is a test suite # that does not run in CI. -inference = ["dep:ort", "dep:ort-tract", "dep:ndarray"] +inference = ["dep:ort", "dep:dr-inference-engine", "dep:ndarray"] diff --git a/core/dr-face/src/detect.rs b/core/dr-face/src/detect.rs index e6727a9..cea345c 100644 --- a/core/dr-face/src/detect.rs +++ b/core/dr-face/src/detect.rs @@ -14,7 +14,8 @@ use ndarray::Array4; -use crate::{install_backend, FaceError}; +use crate::FaceError; +use dr_inference_engine::{Form, Model, Role}; /// The graph's input edge, in pixels. See the module note: not configurable. pub const INPUT_EDGE: usize = 640; @@ -135,7 +136,10 @@ impl Detection { /// A loaded SCRFD graph. pub struct Detector { - session: ort::session::Session, + session: Model, + /// f32 or int8 — the int8 form finds a different set of faces and is a + /// different detector in `model_id` (docs/inference.md §7). + form: Form, /// Feature-map count: 3 for strides {8,16,32}, 4 for {8,16,32,64}. /// /// Discovered from the output count rather than assumed, because both @@ -145,18 +149,29 @@ pub struct Detector { } impl Detector { - pub fn from_path(path: impl AsRef) -> Result { - let bytes = std::fs::read(path).map_err(FaceError::ModelRead)?; - Self::from_bytes(&bytes) + /// Which form this detector was loaded from. + pub fn form(&self) -> Form { + self.form } - pub fn from_bytes(bytes: &[u8]) -> Result { - install_backend(); + /// Load the canonical f32 file at `path`, or the form the device's + /// backend wants instead — the `.int8.onnx` beside it on a Hexagon — + /// which [`Detector::form`] then reports. + pub fn from_path(path: impl AsRef) -> Result { + let (path, form) = dr_inference_engine::resolve_model(Role::Detector, path.as_ref()); + let bytes = std::fs::read(path).map_err(FaceError::ModelRead)?; + Self::from_bytes_in(&bytes, form) + } - let session = ort::session::Session::builder() - .map_err(FaceError::Inference)? - .commit_from_memory(bytes) - .map_err(FaceError::Inference)?; + /// An f32 graph from memory. + pub fn from_bytes(bytes: &[u8]) -> Result { + Self::from_bytes_in(bytes, Form::F32) + } + + fn from_bytes_in(bytes: &[u8], form: Form) -> Result { + let model = dr_inference_engine::open(Role::Detector, form, bytes)?; + let acquired = model.acquire()?; + let session = acquired.lock(); let n_out = session.outputs().len(); if n_out % 3 != 0 || !(9..=12).contains(&n_out) { @@ -191,7 +206,13 @@ impl Detector { } } - Ok(Self { session, fmc }) + drop(session); + drop(acquired); + Ok(Self { + session: model, + form, + fmc, + }) } /// Stride levels this graph emits. @@ -223,8 +244,9 @@ impl Detector { let lb = Letterbox::fit(width as f32, height as f32); let input = lb.sample(rgb, width, height); - let outputs = self - .session + let acquired = self.session.acquire()?; + let mut session = acquired.lock(); + let outputs = session .run(ort::inputs![ ort::value::Tensor::from_array(input).map_err(FaceError::Inference)? ]) diff --git a/core/dr-face/src/embed.rs b/core/dr-face/src/embed.rs index 7e568b7..28c2441 100644 --- a/core/dr-face/src/embed.rs +++ b/core/dr-face/src/embed.rs @@ -14,7 +14,8 @@ use ndarray::Array4; use crate::align::{Aligned112, ALIGNED_EDGE}; use crate::embedding::{normalise, Embedding, ModelId, EMBEDDING_DIM}; -use crate::{install_backend, FaceError}; +use crate::FaceError; +use dr_inference_engine::{Form, Model, Role}; /// What one pass of the embedder produces: the direction, and the length. /// @@ -53,7 +54,7 @@ impl Embedded { /// A loaded ArcFace graph. pub struct Embedder { - session: ort::session::Session, + session: Model, model: ModelId, } @@ -64,12 +65,11 @@ impl Embedder { } pub fn from_bytes(bytes: &[u8], model: ModelId) -> Result { - install_backend(); - - let session = ort::session::Session::builder() - .map_err(FaceError::Inference)? - .commit_from_memory(bytes) - .map_err(FaceError::Inference)?; + // Always the f32 form: an embedding must compare across devices + // (docs/inference.md §7), and the engine pins this role to it. + let loaded = dr_inference_engine::open(Role::Embedder, Form::F32, bytes)?; + let acquired = loaded.acquire()?; + let session = acquired.lock(); // One output, `[1, 512]`. Checked because an ArcFace variant with a // different embedding width would otherwise be read as a truncated @@ -90,7 +90,12 @@ impl Embedder { }); } - Ok(Self { session, model }) + drop(session); + drop(acquired); + Ok(Self { + session: loaded, + model, + }) } pub fn model(&self) -> &ModelId { @@ -111,8 +116,9 @@ impl Embedder { } } - let outputs = self - .session + let acquired = self.session.acquire()?; + let mut session = acquired.lock(); + let outputs = session .run(ort::inputs![ ort::value::Tensor::from_array(input).map_err(FaceError::Inference)? ]) diff --git a/core/dr-face/src/lib.rs b/core/dr-face/src/lib.rs index 89633e8..af0f290 100644 --- a/core/dr-face/src/lib.rs +++ b/core/dr-face/src/lib.rs @@ -124,25 +124,21 @@ pub enum FaceError { ImageShape { expected: usize, got: usize }, } -/// Install tract as `ort`'s backend. -/// -/// Idempotent, and it must happen before any other `ort` call: with -/// `alternative-backend` there is no linked runtime to fall back on, so an -/// un-set API is a panic rather than a slow path. Same helper as -/// `dr-segment::semantic`, for the same reason. #[cfg(feature = "inference")] -pub(crate) fn install_backend() { - use std::sync::Once; - static ONCE: Once = Once::new(); - ONCE.call_once(|| { - let _ = ort::set_api(ort_tract::api()); - }); +impl From for FaceError { + fn from(e: dr_inference_engine::Error) -> Self { + match e { + dr_inference_engine::Error::Inference(e) => FaceError::Inference(e), + dr_inference_engine::Error::Io(e) => FaceError::ModelRead(e), + } + } } -/// [`install_backend`] for the M1 probe example, which drives `ort` directly -/// rather than through [`detect::Detector`] so it can report the raw error. +/// Make sure `ort` has a backend, for the M1 probe example, which drives +/// `ort` directly rather than through [`detect::Detector`] so it can report +/// the raw error. Every other path goes through `dr-inference-engine`. #[cfg(feature = "inference")] #[doc(hidden)] pub fn install_backend_for_probe() { - install_backend(); + dr_inference_engine::ensure_runtime(); } diff --git a/core/dr-inference-engine/Cargo.toml b/core/dr-inference-engine/Cargo.toml new file mode 100644 index 0000000..46da443 --- /dev/null +++ b/core/dr-inference-engine/Cargo.toml @@ -0,0 +1,44 @@ +[package] +name = "dr-inference-engine" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true + +# The one crate that names a runtime, a provider, a vendor library or a +# device (docs/inference.md §8). `dr-face` and `dr-segment` ask it for a +# session by role and never see which of these answered. + +[dependencies] +thiserror.workspace = true +log.workspace = true +serde.workspace = true +serde_json.workspace = true + +# `ort` is the API; what supplies it is decided once per process (§3): +# `libonnxruntime` found on disk, or `tract`. Both are behind +# `alternative-backend`, so nothing here links C on any target. +ort = { workspace = true } +ort-tract = { workspace = true, optional = true } +# dlopen, and the C types of the table it fetches. Both pure Rust; +# `libloading` is already in the tree through wgpu. +libloading = { version = "0.8", optional = true } +ort-sys = { version = "2.0.0-rc.13", default-features = false, features = ["disable-linking"], optional = true } + +# The NVIDIA rungs exist on the desktop only. These features add `ort`'s +# option builders and nothing else — no linking under `alternative-backend` — +# but an Android binary has no business carrying even the option names, and +# the packaging must never be tempted to (§2, §3.1). +[target.'cfg(not(target_os = "android"))'.dependencies] +ort = { workspace = true, features = ["cuda", "tensorrt"] } + +[target.'cfg(target_os = "android")'.dependencies] +ort = { workspace = true, features = ["qnn"] } + +[features] +# The floor: `tract` supplies the API table when no runtime file is found, or +# always, in a build without `native`. Tests want this and nothing else. +default = ["tract"] +tract = ["dep:ort-tract"] +# Look for `libonnxruntime` on disk and hand its table to `ort`. +native = ["dep:libloading", "dep:ort-sys"] diff --git a/core/dr-inference-engine/src/api.rs b/core/dr-inference-engine/src/api.rs new file mode 100644 index 0000000..d18c527 --- /dev/null +++ b/core/dr-inference-engine/src/api.rs @@ -0,0 +1,142 @@ +//! The API table `ort` runs on, chosen once (docs/inference.md §3). +//! +//! `ort` with `alternative-backend` links no runtime and asks, on first use, +//! for an `OrtApi` — a struct of function pointers. Two things can fill it: +//! a `libonnxruntime` this module `dlopen`s, or `ort-tract`. The Rust build +//! is identical either way; the difference is whether a file was found. + +use std::path::PathBuf; +use std::sync::OnceLock; + +/// What supplied the table. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum Runtime { + /// Pure Rust, one core, every operator these graphs use. The floor. + Tract, + /// The C++ ONNX Runtime, loaded from `path`. + OnnxRuntime { path: PathBuf, version: String }, +} + +impl Runtime { + pub fn label(&self) -> String { + match self { + Runtime::Tract => "tract".into(), + Runtime::OnnxRuntime { version, .. } => format!("ONNX Runtime {version}"), + } + } + + pub fn is_native(&self) -> bool { + matches!(self, Runtime::OnnxRuntime { .. }) + } +} + +static RUNTIME: OnceLock = OnceLock::new(); + +/// The runtime in use; tract until something installs another. +pub fn runtime() -> Runtime { + RUNTIME.get().cloned().unwrap_or(Runtime::Tract) +} + +/// Install a table if none is installed yet — tract, since no directories +/// were named. What a test or an example gets. +pub fn ensure_installed() { + if RUNTIME.get().is_none() { + install(&[]); + } +} + +/// Look for `libonnxruntime` in `dirs`, in order, and hand `ort` the first +/// table that loads; otherwise tract. Once per process. +pub fn install(dirs: &[PathBuf]) -> Runtime { + RUNTIME + .get_or_init(|| { + #[cfg(feature = "native")] + for dir in dirs { + match load_native(dir) { + Ok(rt) => return rt, + Err(e) => log::info!("inference: no runtime in {}: {e}", dir.display()), + } + } + #[cfg(not(feature = "native"))] + let _ = dirs; + install_tract() + }) + .clone() +} + +#[cfg(feature = "tract")] +fn install_tract() -> Runtime { + let _ = ort::set_api(ort_tract::api()); + Runtime::Tract +} + +#[cfg(not(feature = "tract"))] +fn install_tract() -> Runtime { + // A build with neither tract nor a runtime file has nothing to run + // models on; every `open` will report the un-set API rather than panic + // somewhere deeper. + log::error!("inference: no ONNX Runtime found and tract is not compiled in"); + Runtime::Tract +} + +#[cfg(feature = "native")] +fn load_native(dir: &std::path::Path) -> Result { + let name = if cfg!(target_os = "windows") { + "onnxruntime.dll" + } else if cfg!(any(target_os = "macos", target_os = "ios")) { + "libonnxruntime.dylib" + } else { + "libonnxruntime.so" + }; + // An empty dir means the bare name: the system loader's search, which on + // Android includes the APK's own native libraries. + let path = if dir.as_os_str().is_empty() { + PathBuf::from(name) + } else { + let p = dir.join(name); + if !p.is_file() { + return Err("not present".into()); + } + p + }; + + // SAFETY: the library's initialisers are ONNX Runtime's own; the symbol + // is the documented entry point with the documented signature; the table + // is copied out and the library handle is leaked, so every pointer in + // the copy stays valid for the life of the process. + unsafe { + let lib = libloading::Library::new(&path).map_err(|e| e.to_string())?; + let get_base: libloading::Symbol< + unsafe extern "system" fn() -> *const ort_sys::OrtApiBase, + > = lib.get(b"OrtGetApiBase\0").map_err(|e| e.to_string())?; + let base = get_base(); + if base.is_null() { + return Err("OrtGetApiBase returned null".into()); + } + let version = std::ffi::CStr::from_ptr(((*base).GetVersionString)()) + .to_string_lossy() + .into_owned(); + let api = ((*base).GetApi)(ort_sys::ORT_API_VERSION); + if api.is_null() { + return Err(format!( + "ONNX Runtime {version} is older than API version {}", + ort_sys::ORT_API_VERSION + )); + } + if !ort::set_api((*api).clone()) { + return Err("an API table was already installed".into()); + } + std::mem::forget(lib); + + // Qualcomm's DSP loader finds the Hexagon skel through this variable, + // and only through it; the runtime's own directory is where the APK + // put it. Harmless anywhere else. + #[cfg(target_os = "android")] + if !dir.as_os_str().is_empty() { + std::env::set_var("ADSP_LIBRARY_PATH", dir); + } + + log::info!("inference: ONNX Runtime {version} from {}", path.display()); + Ok(Runtime::OnnxRuntime { path, version }) + } +} diff --git a/core/dr-inference-engine/src/engines.rs b/core/dr-inference-engine/src/engines.rs new file mode 100644 index 0000000..d565758 --- /dev/null +++ b/core/dr-inference-engine/src/engines.rs @@ -0,0 +1,104 @@ +//! Compiled engines: what a rung builds once per device, and the thread that +//! builds them before anyone asks (docs/inference.md §5, §6). +//! +//! TensorRT keeps its own engine cache keyed by graph hash; QNN writes a +//! context model. Both are opaque to this crate, which tracks only *that* a +//! model compiled — by the hash of its bytes — so [`crate::open`] can tell a +//! request whether to expect the rung or its fallback. + +use std::path::PathBuf; + +use crate::{state, Config, Form, Rung}; + +/// 64-bit FNV-1a. A cache key, not a checksum: two model files that collide +/// here would have to also be the same size and the same role, and the cost +/// of that is a rebuilt engine. +pub fn hash(bytes: &[u8]) -> u64 { + let mut h = 0xcbf2_9ce4_8422_2325u64; + for &b in bytes { + h ^= b as u64; + h = h.wrapping_mul(0x0000_0100_0000_01b3); + } + h +} + +/// The cache entry for `bytes` compiled on `rung`. +pub fn key(rung: Rung, bytes: &[u8]) -> String { + format!("{}:{:016x}", rung.label(), hash(bytes)) +} + +/// Where QNN's compiled context for `bytes` lives. +pub fn context_path(cfg: &Config, bytes: &[u8]) -> PathBuf { + cfg.cache_dir + .join("qnn") + .join(format!("{:016x}_ctx.onnx", hash(bytes))) +} + +/// After the probe: compile every configured model the selected rung can +/// take, smallest first, recording each as it lands. +pub fn run() { + let (rung, cfg) = { + let s = state().lock().unwrap(); + (crate::current_rung(&s), s.config.clone()) + }; + if !rung.compiles() { + return; + } + + // Smallest first, so the detector — the one that runs per image — is + // ready soonest (§6 step 3). + let mut jobs: Vec<(crate::Role, PathBuf, u64)> = cfg + .models + .iter() + .filter(|(role, _)| rung.form(*role) != Form::F32 || rung != Rung::Hexagon) + .filter_map(|(role, path)| { + let (path, form) = crate::resolve_model(*role, path); + (form == rung.form(*role)).then(|| { + let size = std::fs::metadata(&path).map(|m| m.len()).unwrap_or(0); + (*role, path, size) + }) + }) + .collect(); + jobs.sort_by_key(|j| j.2); + state().lock().unwrap().wanted = jobs.len(); + + for (role, path, _) in jobs { + let Ok(bytes) = std::fs::read(&path) else { + continue; + }; + let key = key(rung, &bytes); + if state().lock().unwrap().cache.compiled.contains(&key) { + continue; + } + log::info!( + "inference: compiling {} for {}", + path.display(), + rung.label() + ); + let started = std::time::Instant::now(); + match crate::session::build(rung, role, &bytes, &cfg, false) { + Ok(session) => { + drop(session); + let mut s = state().lock().unwrap(); + s.cache.compiled.insert(key); + crate::probe::write_cache(&s.config, &s.cache); + log::info!( + "inference: {} ready on {} in {:.1} s", + path.display(), + rung.label(), + started.elapsed().as_secs_f64() + ); + } + Err(e) => { + // This model stays on the fallback; the others still get + // their engine. A corrected model file changes the hash and + // is retried. + log::warn!( + "inference: {} will not compile for {}: {e}", + path.display(), + rung.label() + ); + } + } + } +} diff --git a/core/dr-inference-engine/src/lib.rs b/core/dr-inference-engine/src/lib.rs new file mode 100644 index 0000000..2eac684 --- /dev/null +++ b/core/dr-inference-engine/src/lib.rs @@ -0,0 +1,542 @@ +//! 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, +} + +/// 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)>, + /// 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()); + } +} diff --git a/core/dr-inference-engine/src/probe.rs b/core/dr-inference-engine/src/probe.rs new file mode 100644 index 0000000..ffff81b --- /dev/null +++ b/core/dr-inference-engine/src/probe.rs @@ -0,0 +1,279 @@ +//! Walk the ladder, once, by building real sessions (docs/inference.md §4). +//! +//! A rung is taken when a strict session builds on it, runs, and is faster +//! than the floor. Both halves matter: a provider can register and then fail +//! at partition time, and a provider can take a graph and run it slower than +//! the CPU would have. The outcome is cached against a fingerprint of the +//! runtime, the driver, the hardware and the models, and trusted until any +//! of those changes. + +use std::path::{Path, PathBuf}; +use std::time::Instant; + +use crate::{api::Runtime, state, Cache, Config, Form, Role, Rung}; + +/// The rungs to try on this platform, best first, under the user's ceiling. +fn ladder(ceiling: Option) -> Vec { + #[cfg(target_os = "android")] + let all = [Rung::Hexagon]; + #[cfg(not(target_os = "android"))] + let all = [Rung::TensorRt, Rung::Cuda]; + all.into_iter() + .filter(|r| ceiling.is_none_or(|c| *r <= c)) + .collect() +} + +/// The probe body. Sets the cache and clears `probing` when done; never +/// panics out, because a failed probe is a result (the floor) and not an +/// error. +pub fn run(runtime: Runtime) { + let cfg = state().lock().unwrap().config.clone(); + let fingerprint = fingerprint(&runtime, &cfg); + + if let Some(cached) = read_cache(&cfg) { + if cached.fingerprint == fingerprint && cached.rung.is_some() { + log::info!( + "inference: cached selection {} ({})", + cached.rung.unwrap().label(), + cached.reason + ); + finish(cached); + return; + } + } + + let mut cache = Cache { + fingerprint, + ..Cache::default() + }; + + if !runtime.is_native() { + cache.rung = Some(Rung::Cpu); + cache.reason = "no ONNX Runtime found; tract on one core".into(); + write_cache(&cfg, &cache); + finish(cache); + return; + } + + let Some((role, canonical)) = probe_model(&cfg) else { + cache.rung = Some(Rung::Cpu); + cache.reason = "no model to probe with".into(); + write_cache(&cfg, &cache); + finish(cache); + return; + }; + + let floor = match time_rung(Rung::Cpu, role, &canonical, &cfg) { + Ok((ms, _)) => ms, + Err(e) => { + // The CPU provider failing is the runtime failing; there is + // nothing below it to try, and the reason is worth reading. + cache.rung = Some(Rung::Cpu); + cache.reason = format!("CPU provider failed: {e}"); + write_cache(&cfg, &cache); + finish(cache); + return; + } + }; + log::info!("inference: floor {floor:.1} ms on the CPU provider"); + + for rung in ladder(cfg.ceiling) { + match time_rung(rung, role, &canonical, &cfg) { + Ok((ms, key)) if ms < floor => { + cache.rung = Some(rung); + cache.reason = format!("{ms:.1} ms against {floor:.1} ms on the CPU"); + if let Some(key) = key { + cache.compiled.insert(key); + } + break; + } + Ok((ms, _)) => { + let why = format!("{ms:.1} ms, slower than the CPU's {floor:.1} ms"); + log::info!("inference: {} rejected: {why}", rung.label()); + cache.failed.push((rung, why)); + } + Err(e) => { + log::info!("inference: {} failed: {e}", rung.label()); + cache.failed.push((rung, e)); + } + } + } + if cache.rung.is_none() { + cache.rung = Some(Rung::Cpu); + cache.reason = match cache.failed.first() { + Some((r, why)) => format!("{} {}", r.label(), first_line(why)), + None => "the only rung on this platform".into(), + }; + } + write_cache(&cfg, &cache); + finish(cache); +} + +fn finish(cache: Cache) { + let mut s = state().lock().unwrap(); + s.cache = cache; + s.probing = false; +} + +/// The smallest configured model: the detector on every device shipped +/// today, and a ~2 MB graph is the cheapest real test of a provider. +fn probe_model(cfg: &Config) -> Option<(Role, PathBuf)> { + cfg.models + .iter() + .filter_map(|(role, path)| { + let size = std::fs::metadata(path).ok()?.len(); + Some((size, *role, path.clone())) + }) + .min_by_key(|(size, _, _)| *size) + .map(|(_, role, path)| (role, path)) +} + +/// Build strictly, run once for the engine, then time three runs; the +/// median in milliseconds and, for a compiling rung, the cache key of the +/// engine this just built. +fn time_rung( + rung: Rung, + role: Role, + canonical: &Path, + cfg: &Config, +) -> Result<(f64, Option), String> { + let want = rung.form(role); + let path = match want { + Form::Int8 => { + let p = crate::int8_sibling(canonical); + if !p.is_file() { + return Err(format!("no int8 form of {}", canonical.display())); + } + p + } + Form::F32 => canonical.to_path_buf(), + }; + let bytes = std::fs::read(&path).map_err(|e| e.to_string())?; + let started = Instant::now(); + let mut session = crate::session::build(rung, role, &bytes, cfg, true) + .map_err(|e| first_line(&e.to_string()))?; + log::info!( + "inference: {} session built in {:.1} s", + rung.label(), + started.elapsed().as_secs_f64() + ); + + let shape: Vec = session.inputs()[0] + .dtype() + .tensor_shape() + .ok_or("model input is not a tensor")? + .iter() + .map(|&d| if d > 0 { d as usize } else { 1 }) + .collect(); + let zeros = vec![0f32; shape.iter().product()]; + let run = |session: &mut ort::session::Session| -> Result { + let input = ort::value::Tensor::from_array((shape.clone(), zeros.clone())) + .map_err(|e| e.to_string())?; + let t = Instant::now(); + let out = session + .run(ort::inputs![input]) + .map_err(|e| e.to_string())?; + let _ = out[0] + .try_extract_tensor::() + .map_err(|e| e.to_string())?; + Ok(t.elapsed().as_secs_f64() * 1e3) + }; + run(&mut session)?; + let mut times = [run(&mut session)?, run(&mut session)?, run(&mut session)?]; + times.sort_by(|a, b| a.partial_cmp(b).unwrap()); + let key = rung.compiles().then(|| crate::engines::key(rung, &bytes)); + Ok((times[1], key)) +} + +fn first_line(s: &str) -> String { + s.lines().next().unwrap_or("").chars().take(160).collect() +} + +/// Everything a change of which should re-probe: the runtime and where it +/// came from, this crate, the platform, the driver or SoC, and the models. +fn fingerprint(runtime: &Runtime, cfg: &Config) -> String { + let mut parts = vec![ + format!("engine {}", env!("CARGO_PKG_VERSION")), + format!("{} {}", std::env::consts::OS, std::env::consts::ARCH), + match runtime { + Runtime::Tract => "tract".to_string(), + Runtime::OnnxRuntime { path, version } => format!("ort {version} {}", path.display()), + }, + device_identity(), + ]; + for (role, path) in &cfg.models { + let hash = std::fs::read(path) + .map(|b| crate::engines::hash(&b)) + .unwrap_or(0); + parts.push(format!("{role:?} {hash:016x}")); + let int8 = crate::int8_sibling(path); + if let Ok(b) = std::fs::read(&int8) { + parts.push(format!("{role:?} int8 {:016x}", crate::engines::hash(&b))); + } + } + parts.join("\n") +} + +#[cfg(target_os = "linux")] +fn device_identity() -> String { + // The NVIDIA driver's version line; absent means no NVIDIA driver. + std::fs::read_to_string("/proc/driver/nvidia/version") + .ok() + .and_then(|s| s.lines().next().map(str::to_string)) + .unwrap_or_else(|| "no nvidia driver".into()) +} + +#[cfg(target_os = "android")] +fn device_identity() -> String { + // The SoC and the vendor's build: a Hexagon appears or disappears with + // either. + format!( + "{} {}", + system_property("ro.soc.model"), + system_property("ro.build.version.incremental") + ) +} + +#[cfg(target_os = "android")] +fn system_property(name: &str) -> String { + extern "C" { + fn __system_property_get( + name: *const std::ffi::c_char, + value: *mut std::ffi::c_char, + ) -> i32; + } + let name = std::ffi::CString::new(name).unwrap(); + let mut buf = [0u8; 92]; // PROP_VALUE_MAX + // SAFETY: bionic's documented call; the buffer is PROP_VALUE_MAX bytes. + let n = unsafe { __system_property_get(name.as_ptr(), buf.as_mut_ptr().cast()) }; + String::from_utf8_lossy(&buf[..n.max(0) as usize]).into_owned() +} + +#[cfg(not(any(target_os = "linux", target_os = "android")))] +fn device_identity() -> String { + String::new() +} + +fn cache_path(cfg: &Config) -> PathBuf { + cfg.cache_dir.join("backend.json") +} + +fn read_cache(cfg: &Config) -> Option { + let text = std::fs::read_to_string(cache_path(cfg)).ok()?; + serde_json::from_str(&text).ok() +} + +/// Written whole and renamed into place, so a reader never sees half. +pub fn write_cache(cfg: &Config, cache: &Cache) { + if cfg.cache_dir.as_os_str().is_empty() { + return; + } + let path = cache_path(cfg); + let tmp = path.with_extension("json.tmp"); + let _ = std::fs::create_dir_all(&cfg.cache_dir); + if let Ok(text) = serde_json::to_string_pretty(cache) { + if std::fs::write(&tmp, text).is_ok() { + let _ = std::fs::rename(&tmp, &path); + } + } +} diff --git a/core/dr-inference-engine/src/session.rs b/core/dr-inference-engine/src/session.rs new file mode 100644 index 0000000..f899f6b --- /dev/null +++ b/core/dr-inference-engine/src/session.rs @@ -0,0 +1,129 @@ +//! One session builder per rung (docs/inference.md §2, §7, §9). + +use ort::session::builder::GraphOptimizationLevel; +use ort::session::Session; + +use crate::{Config, Role, Rung}; + +/// Build a session for `bytes` on `rung`. +/// +/// `strict` is the probe's flag: with it, a provider that would hand any +/// node to the CPU fails the build instead, so "the session built" means +/// "the provider took the graph" and not "the provider registered" (§4). +pub fn build( + rung: Rung, + role: Role, + bytes: &[u8], + cfg: &Config, + strict: bool, +) -> ort::Result { + let mut b = Session::builder()? + .with_optimization_level(GraphOptimizationLevel::Level3)? + .with_intra_threads(threads(cfg))?; + if strict { + b = b.with_config_entry("session.disable_cpu_ep_fallback", "1")?; + } + // A Hexagon session loads the compiled context when there is one and + // compiles it from the model when there is not; the engine thread is + // what makes the second case rare (§6). + let context = (rung == Rung::Hexagon).then(|| crate::engines::context_path(cfg, bytes)); + let ready = context.as_ref().is_some_and(|p| p.is_file()); + b = providers( + b, + rung, + role, + cfg, + if ready { None } else { context.as_deref() }, + )?; + match (ready, context) { + (true, Some(path)) => b.commit_from_file(path), + _ => b.commit_from_memory(bytes), + } +} + +/// The intra-op pool: what the config says, else the cores less two for +/// the compositor and the decoder (§9). tract ignores it. +fn threads(cfg: &Config) -> usize { + if cfg.threads > 0 { + return cfg.threads; + } + std::thread::available_parallelism() + .map(|n| n.get().saturating_sub(2).max(1)) + .unwrap_or(1) +} + +#[cfg(not(target_os = "android"))] +fn providers( + b: ort::session::builder::SessionBuilder, + rung: Rung, + role: Role, + cfg: &Config, + _generate_context: Option<&std::path::Path>, +) -> ort::Result { + use ort::ep; + match rung { + Rung::Cpu => Ok(b), + Rung::Cuda => { + Ok(b.with_execution_providers([ep::CUDA::default().build().error_on_failure()])?) + } + Rung::TensorRt => { + let cache = cfg.cache_dir.join("tensorrt"); + let _ = std::fs::create_dir_all(&cache); + let cache = cache.to_string_lossy().into_owned(); + // fp16 for everything but the embedder, whose comparability + // across devices is worth more than its 0.2 ms (§7). The + // workspace cap keeps the develop view's tiles on the card + // (NFR-RES-2). CUDA behind it takes any node TensorRT declines. + Ok(b.with_execution_providers([ + ep::TensorRT::default() + .with_fp16(role != Role::Embedder) + .with_engine_cache(true) + .with_engine_cache_path(&cache) + .with_timing_cache(true) + .with_timing_cache_path(&cache) + .with_max_workspace_size(512 << 20) + .build() + .error_on_failure(), + ep::CUDA::default().build(), + ])?) + } + Rung::Hexagon => unreachable!("the Hexagon rung is not on a desktop ladder"), + } +} + +#[cfg(target_os = "android")] +fn providers( + b: ort::session::builder::SessionBuilder, + rung: Rung, + _role: Role, + _cfg: &Config, + generate_context: Option<&std::path::Path>, +) -> ort::Result { + use ort::ep; + match rung { + Rung::Cpu => Ok(b), + Rung::Hexagon => { + // The HTP compiles the graph once per device (0.8–1.7 s here). + // With `ep.context_enable` ONNX Runtime writes the compiled + // context beside the probe cache; the next session loads that + // file as its model and skips the compile (§5). + let mut b = b; + if let Some(ctx) = generate_context { + let _ = std::fs::create_dir_all(ctx.parent().unwrap()); + b = b + .with_config_entry("ep.context_enable", "1")? + .with_config_entry("ep.context_file_path", ctx.to_string_lossy())? + .with_config_entry("ep.context_embed_mode", "0")?; + } + // Quantise/dequantise at the graph's edges stay on the NPU too, + // so a strict build is a whole-graph build. + Ok(b.with_execution_providers([ep::QNN::default() + .with_backend_path("libQnnHtp.so") + .with_performance_mode(ep::qnn::PerformanceMode::Burst) + .with_offload_graph_io_quantization(false) + .build() + .error_on_failure()])?) + } + Rung::Cuda | Rung::TensorRt => unreachable!("no NVIDIA rung on Android"), + } +} diff --git a/core/dr-segment/Cargo.toml b/core/dr-segment/Cargo.toml index 0733375..6f0f021 100644 --- a/core/dr-segment/Cargo.toml +++ b/core/dr-segment/Cargo.toml @@ -11,10 +11,11 @@ build = "build.rs" thiserror.workspace = true log.workspace = true -# Inference. `ort` is the API; **tract is the engine** — see the workspace -# manifest for why the C++ ONNX Runtime is not linked here. +# Inference. `ort` is the API; **what runs it is `dr-inference-engine`'s +# business** — tract, or an ONNX Runtime the app found on disk, on whichever +# provider the device has (docs/inference.md). This crate never names either. ort = { workspace = true, optional = true } -ort-tract = { workspace = true, optional = true } +dr-inference-engine = { workspace = true, optional = true } ndarray = { workspace = true, optional = true } [dev-dependencies] @@ -35,7 +36,7 @@ default = ["semantic", "embedded-model"] # Separable because the watershed half is genuinely independent of it: with # this off, `dr-segment` is a pure-CPU graph algorithm crate with no model to # carry, which is what the headless hierarchy tests want. -semantic = ["dep:ort", "dep:ort-tract", "dep:ndarray"] +semantic = ["dep:ort", "dep:dr-inference-engine", "dep:ndarray"] # Compile the weights into the binary. # diff --git a/core/dr-segment/src/lib.rs b/core/dr-segment/src/lib.rs index 1a10271..7dbbb25 100644 --- a/core/dr-segment/src/lib.rs +++ b/core/dr-segment/src/lib.rs @@ -98,3 +98,13 @@ pub enum SegmentError { #[error("category descriptor: {0}")] CategoryDescriptor(String), } + +#[cfg(feature = "semantic")] +impl From for SegmentError { + fn from(e: dr_inference_engine::Error) -> Self { + match e { + dr_inference_engine::Error::Inference(e) => SegmentError::Inference(e), + dr_inference_engine::Error::Io(e) => SegmentError::ModelRead(e), + } + } +} diff --git a/core/dr-segment/src/scene.rs b/core/dr-segment/src/scene.rs index e469d87..f92fe73 100644 --- a/core/dr-segment/src/scene.rs +++ b/core/dr-segment/src/scene.rs @@ -59,7 +59,7 @@ use ndarray::ArrayView3; #[cfg(test)] use crate::semantic::INPUT_EDGE; -use crate::semantic::{install_backend, Letterbox, Window}; +use crate::semantic::{Letterbox, Window}; use crate::SegmentError; /// Classes in the ADE20K vocabulary the scene model was trained on. @@ -84,7 +84,7 @@ pub struct Category { /// The scene model, and the categories it has been told to report. pub struct SceneModel { - session: ort::session::Session, + session: dr_inference_engine::Model, categories: Vec, } @@ -132,12 +132,12 @@ impl SceneModel { } pub fn from_bytes(bytes: &[u8], categories: Vec) -> Result { - install_backend(); - - let session = ort::session::Session::builder() - .map_err(SegmentError::Inference)? - .commit_from_memory(bytes) - .map_err(SegmentError::Inference)?; + // f32, as for `SemanticModel`; see there. + let session = dr_inference_engine::open( + dr_inference_engine::Role::Scene, + dr_inference_engine::Form::F32, + bytes, + )?; Ok(Self { session, @@ -170,13 +170,9 @@ impl SceneModel { }); } - // Split the borrow: `run` needs the session mutably while - // `marginalise` needs the categories, and going through `self` for - // both at once is what the borrow checker objects to. - let Self { - session, - categories, - } = self; + let categories = &self.categories; + let acquired = self.session.acquire()?; + let mut session = acquired.lock(); let window = Window { x: 0.0, diff --git a/core/dr-segment/src/semantic.rs b/core/dr-segment/src/semantic.rs index 1beb291..e64fa15 100644 --- a/core/dr-segment/src/semantic.rs +++ b/core/dr-segment/src/semantic.rs @@ -194,7 +194,7 @@ impl Instance { /// Holds an `ort` session, so it is neither `Clone` nor cheap to build — /// construct once and keep it. Loading is ~50 ms. pub struct SemanticModel { - session: ort::session::Session, + session: dr_inference_engine::Model, classes: Vec>, } @@ -229,15 +229,14 @@ impl SemanticModel { } pub fn from_bytes(bytes: &[u8], classes: Vec>) -> Result { - // Idempotent, and it must happen before any other `ort` call: with - // `alternative-backend` there is no linked runtime to fall back on, so - // an un-set API is a panic rather than a slow path. - install_backend(); - - let session = ort::session::Session::builder() - .map_err(SegmentError::Inference)? - .commit_from_memory(bytes) - .map_err(SegmentError::Inference)?; + // The f32 graph on whatever the device's backend is. An int8 form + // for the Hexagon waits on docs/inference.md §10 M7 — the mask + // boundary has to be measured before it moves. + let session = dr_inference_engine::open( + dr_inference_engine::Role::Segmenter, + dr_inference_engine::Form::F32, + bytes, + )?; Ok(Self { session, classes }) } @@ -332,8 +331,9 @@ impl SemanticModel { let letterbox = Letterbox::fit(window.w, window.h); let input = letterbox.sample(rgb, width, height, window); - let outputs = self - .session + let acquired = self.session.acquire()?; + let mut session = acquired.lock(); + let outputs = session .run(ort::inputs![ ort::value::Tensor::from_array(input).map_err(SegmentError::Inference)? ]) @@ -674,18 +674,6 @@ fn steps(extent: f32, edge: f32, stride: f32) -> usize { } } -/// Point `ort` at tract, exactly once per process. -pub(crate) fn install_backend() { - use std::sync::Once; - static ONCE: Once = Once::new(); - ONCE.call_once(|| { - // Returns false if an API was already installed, which is not an error - // — it means something else got here first, and there is only one - // backend compiled in for it to have chosen. - let _ = ort::set_api(ort_tract::api()); - }); -} - /// Read the class list written beside the model by `tools/export-seg-model.sh`. /// /// A deliberately small hand-rolled reader for a flat array of strings, rather