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:
@@ -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());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user