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,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<Rung>) -> Vec<Rung> {
|
||||
#[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>), 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<usize> = 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<f64, String> {
|
||||
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::<f32>()
|
||||
.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<Cache> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user