Add dr-inference-engine and route every model session through it

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.
This commit is contained in:
2026-09-19 16:02:37 +02:00
parent caf21bea64
commit d15c41e699
16 changed files with 1368 additions and 90 deletions
+542
View File
@@ -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<PathBuf>,
/// 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<Rung>,
/// 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<Loaded>,
}
struct Loaded {
rung: Rung,
session: Mutex<Session>,
last_used: Mutex<Instant>,
}
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<Acquired, Error> {
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<String, Arc<Loaded>>;
static REGISTRY: OnceLock<Mutex<Registry>> = OnceLock::new();
fn registry() -> &'static Mutex<Registry> {
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<Acquired, Error> {
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<Rung>,
reason: String,
/// Model hashes whose engine exists on disk, per compiling rung.
compiled: BTreeSet<String>,
/// 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<Mutex<State>> = OnceLock::new();
fn state() -> &'static Mutex<State> {
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<Model, Error> {
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<Vec<u8>> {
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());
}
}