The desktop names where a package may have put libonnxruntime — an override variable, beside the executable, the package's own library directory, the Flatpak prefix, the system library directory — and Android points at the APK's native library directory, which is also what Qualcomm's DSP loader must be told for the Hexagon skel. Android starts the engine at the end of the model unpack rather than at launch, because the probe fingerprints the model files and a first launch has none until then. The About panel gains an Inference row beside Graphics, re-read every two seconds while the probe runs and engines land, and faces.model_id carries the detector's form: an int8 detector finds a different set of faces and is a different population (docs/inference.md §7). A low-memory signal drops every idle session with the GPU caches. The APK assembly bundles ONNX Runtime and the Qualcomm HTP libraries from Maven, fetched by tools/fetch-android-runtime.sh with their published checksums; RUNTIME_DIR=none builds the tract-only APK, which is a slower app and not a broken one. The desktop packages carry no runtime yet. Two probe fixes from the first desktop run: the floor must not be built with CPU fallback disabled, and a versioned libonnxruntime.so is a runtime too. On the reference desktop the probe now loads ONNX Runtime 1.30, measures 30 ms on the CPU provider, and selects TensorRT at 1.5 ms.
549 lines
18 KiB
Rust
549 lines
18 KiB
Rust
//! Which runtime, which provider and which model form — decided once per
|
|
//! device, and the only crate that knows the answer (docs/inference.md).
|
|
//!
|
|
//! Consumers ask for a session by [`Role`] and get `ort`'s `Session` back;
|
|
//! what built it — tract on one core, ONNX Runtime's CPU pool, a TensorRT
|
|
//! engine, the Hexagon — is this crate's business and shows up in
|
|
//! [`status`] for the settings row and nowhere else.
|
|
//!
|
|
//! The shape follows §3 of the spec: `ort` links nothing (`alternative-backend`),
|
|
//! and the first call hands it an API table from either a `libonnxruntime`
|
|
//! found on disk or from `tract`. That choice is once per process, because
|
|
//! `ort::set_api` is; everything after it — which provider, whether an engine
|
|
//! has been compiled yet — is per session and may change between two calls.
|
|
|
|
use std::collections::{BTreeSet, HashMap};
|
|
use std::path::{Path, PathBuf};
|
|
use std::sync::{Arc, Mutex, MutexGuard, OnceLock};
|
|
use std::time::{Duration, Instant};
|
|
|
|
use serde::{Deserialize, Serialize};
|
|
|
|
mod api;
|
|
mod engines;
|
|
mod probe;
|
|
mod session;
|
|
|
|
pub use api::Runtime;
|
|
pub use ort::session::Session;
|
|
|
|
/// What a model is for. The role fixes the precision rule (§7): an embedder
|
|
/// runs in f32 on every rung, a detector may run in fp16 or int8.
|
|
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
|
pub enum Role {
|
|
Detector,
|
|
Embedder,
|
|
Segmenter,
|
|
Scene,
|
|
/// The dense landmark model behind the eye reading (docs/faces.md §7c).
|
|
Landmarks,
|
|
/// The eye-state and sunglasses classifiers, a few hundred kilobytes.
|
|
EyeClassifier,
|
|
}
|
|
|
|
/// Which numeric form of a model a session was built from.
|
|
///
|
|
/// `Int8` is a different network from `F32` for a detector — it finds a
|
|
/// different set of faces — which is why [`form_suffix`] exists and why a
|
|
/// caller appends it to `model_id`.
|
|
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
|
pub enum Form {
|
|
F32,
|
|
Int8,
|
|
}
|
|
|
|
/// A rung of the ladder (§2). Ordered: a user override names the highest rung
|
|
/// the probe may take, and a compiling rung falls back to the one below it
|
|
/// until its engine exists.
|
|
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
|
|
pub enum Rung {
|
|
/// ONNX Runtime's CPU provider, or tract when no runtime file was found.
|
|
Cpu,
|
|
/// NVIDIA, through the CUDA provider. Desktop only.
|
|
Cuda,
|
|
/// NVIDIA, through a TensorRT engine compiled on this device. Desktop only.
|
|
TensorRt,
|
|
/// Qualcomm's Hexagon NPU through QNN, int8 models only. Android only.
|
|
Hexagon,
|
|
}
|
|
|
|
impl Rung {
|
|
pub fn label(self) -> &'static str {
|
|
match self {
|
|
Rung::Cpu => "CPU",
|
|
Rung::Cuda => "CUDA",
|
|
Rung::TensorRt => "TensorRT",
|
|
Rung::Hexagon => "Hexagon NPU",
|
|
}
|
|
}
|
|
|
|
/// The rung a request lands on while this one's engine is still being
|
|
/// compiled (§6 step 2).
|
|
fn fallback(self) -> Rung {
|
|
match self {
|
|
Rung::TensorRt => Rung::Cuda,
|
|
Rung::Hexagon | Rung::Cuda | Rung::Cpu => Rung::Cpu,
|
|
}
|
|
}
|
|
|
|
/// Whether a session on this rung needs an engine built first.
|
|
fn compiles(self) -> bool {
|
|
matches!(self, Rung::TensorRt | Rung::Hexagon)
|
|
}
|
|
|
|
/// The model form this rung wants for a role.
|
|
fn form(self, role: Role) -> Form {
|
|
match (self, role) {
|
|
(Rung::Hexagon, Role::Embedder) => Form::F32,
|
|
(Rung::Hexagon, _) => Form::Int8,
|
|
_ => Form::F32,
|
|
}
|
|
}
|
|
}
|
|
|
|
/// How long a session outlives its last use unless [`Config::decay`] says
|
|
/// otherwise: long enough for the next click, short enough that a session's
|
|
/// GPU or NPU memory does not sit under the develop view for long.
|
|
pub const DEFAULT_DECAY: Duration = Duration::from_secs(30);
|
|
|
|
/// What [`init`] is told once, at launch.
|
|
#[derive(Clone, Debug, Default)]
|
|
pub struct Config {
|
|
/// Where to look for `libonnxruntime`, in order. An empty path means "the
|
|
/// bare library name through the system loader", which is how the APK's
|
|
/// own copy is found on Android.
|
|
pub runtime_dirs: Vec<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)>,
|
|
/// Models compiled into the binary, for the same reason.
|
|
pub embedded: Vec<(Role, &'static [u8])>,
|
|
/// The highest rung the user allows; `None` is "the best that works".
|
|
pub ceiling: Option<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());
|
|
}
|
|
}
|