Files
DarkRoom/core/dr-inference-engine/src/probe.rs
T
dtourolle 87c405eb46 Add OpenVINO and WebGPU rungs to the inference ladder
OpenVINO is the Intel rung: the integrated or Arc GPU, fp16 for every
role but the embedder, a compiled program per model kept in a directory
per model, precision and runtime version. On the Iris Xe it beats ONNX
Runtime's CPU provider on every shipped model — scrfd_10g 23 ms against
58, the scene model 17 against 57, MI-GAN 57 against 330, a denoise tile
40 against 158.

WebGPU is the generic rung for a GPU no vendor rung covers. It was slower
than the CPU on the Iris Xe, the RTX 3050 and the Adreno, so it is on the
ladder for the GPUs it has not been timed on, behind the probe's clock.

MIGraphX's registration becomes one generic key/value helper that all
three share, with option names read from each runtime's own source.
2026-10-04 21:00:15 -04:00

582 lines
21 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Walk the ladder, once, by building real sessions (docs/dev/inference.md §4).
//!
//! A rung is taken when a 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 — or quietly hand most
//! of it back to the CPU — 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> {
// WebGPU is the generic rung (§2): it is reached only on a runtime that
// carries it, which `api` loads where no vendor's runtime fits the
// device, and kept only where it beats the CPU.
#[cfg(target_os = "android")]
let all = [Rung::Hexagon, Rung::WebGpu];
// Unmeasured (§2 ⁵): it is on the ladder because the probe's clock and
// `attempt` make a wrong guess cost one slow or failed probe, not a
// slow or crashing app.
#[cfg(target_os = "macos")]
let all = [Rung::CoreMl];
// A runtime carries one vendor's providers, chosen for this device's
// GPU (`api`); the others are "not enabled in this build", and that
// answer arrives in milliseconds.
#[cfg(not(any(target_os = "android", target_os = "macos")))]
let all = [
Rung::TensorRt,
Rung::Cuda,
Rung::MiGraphX,
Rung::OpenVino,
Rung::WebGpu,
];
all.into_iter()
.filter(|r| ceiling.is_none_or(|c| *r <= c))
.collect()
}
/// Probes under one fingerprint that may end on the CPU after an
/// accelerator failed or lost, before that answer is kept.
const RETRIES: u32 = 3;
/// What a cached probe result is good for.
#[derive(Debug, PartialEq)]
enum Reuse {
/// Use it as it is.
Keep,
/// Probe again: it fell back to the CPU after this many probes.
Again(u32),
/// Another device, runtime or model set: probe from the start.
Fresh,
}
/// The CPU because an accelerator failed or lost is asked again on the next
/// launches, a few times: a failure can be a moment's (QNN could not create
/// its device on 0.22.0's first launch after the update), and keeping it for
/// good left the tablet's every model on the CPU. Bounded, so a wedged
/// driver costs a few launches, not all.
fn reuse(cached: &Cache, fingerprint: &str) -> Reuse {
if cached.fingerprint != fingerprint || cached.rung.is_none() {
return Reuse::Fresh;
}
let fell_back = cached.rung == Some(Rung::Cpu) && !cached.failed.is_empty();
if fell_back && cached.attempts < RETRIES {
Reuse::Again(cached.attempts)
} else {
Reuse::Keep
}
}
/// 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);
let mut attempts = 0;
if let Some(cached) = read_cache(&cfg) {
match reuse(&cached, &fingerprint) {
Reuse::Keep => {
log::info!(
"inference: cached selection {} ({})",
cached.rung.map_or("?", |r| r.label()),
cached.reason
);
finish(cached);
return;
}
Reuse::Again(n) => {
attempts = n;
log::info!(
"inference: probing again after falling back to the CPU ({}), attempt {} of {RETRIES}",
cached.reason,
n + 1
);
}
Reuse::Fresh => {}
}
}
let mut cache = Cache {
fingerprint,
attempts: attempts + 1,
..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) {
let timed = attempt(&cfg, &format!("probe {}", rung.label()), || {
time_rung(rung, role, &canonical, &cfg)
})
.and_then(|timed| timed);
match timed {
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;
}
/// How many launches in a row may die inside one attempt before it is
/// refused. Two, not one: quitting the app while TensorRT spends forty
/// seconds on an engine leaves the same trace as a provider that aborted.
const STRIKES: u32 = 2;
/// Run `f` — a session build on a provider — with `what` written down
/// first, so that if the provider takes the process with it the next launch
/// knows what to stop trying.
///
/// A provider can fail by aborting rather than by returning an error:
/// XNNPACK did on SCRFD (§2), and a C++ exception or a panic across the C
/// API is an abort. The probe runs in the app's own process, so a rung that
/// does this once would do it on every launch, before the first photograph
/// is on screen. The file (`attempt` in the cache directory) holds the
/// attempt and how many launches have started it without finishing;
/// finishing, by success or by error, removes it. After [`STRIKES`] the
/// attempt is refused, and the caller records the refusal in the cache,
/// where it lasts until the fingerprint changes like any other failure.
pub fn attempt<T>(cfg: &Config, what: &str, f: impl FnOnce() -> T) -> Result<T, String> {
if cfg.cache_dir.as_os_str().is_empty() {
return Ok(f());
}
let path = cfg.cache_dir.join("attempt");
let died = std::fs::read_to_string(&path)
.ok()
.and_then(|s| {
let (w, n) = s.split_once('\t')?;
(w == what).then(|| n.trim().parse::<u32>().ok())?
})
.unwrap_or(0);
if died >= STRIKES {
log::error!("inference: the app died during `{what}` on the last {died} launches; not trying it again");
return Err(format!(
"the app died while trying this on {died} launches in a row"
));
}
if died > 0 {
log::warn!("inference: the last launch died during `{what}`; trying it once more");
}
let _ = std::fs::create_dir_all(&cfg.cache_dir);
let _ = std::fs::write(&path, format!("{what}\t{}", died + 1));
let out = f();
let _ = std::fs::remove_file(&path);
Ok(out)
}
/// The smallest detector, or the smallest model of any role if there is
/// none. A ~2 MB detector is the cheapest real test of a provider, and
/// every rung serves it — the eye classifiers are smaller still, but the
/// Hexagon does not take them, and a probe with one would fail it for that.
fn probe_model(cfg: &Config) -> Option<(Role, PathBuf)> {
let smallest = |want: Option<Role>| {
cfg.models
.iter()
.filter(|(role, _)| want.is_none_or(|w| *role == w))
.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))
};
smallest(Some(Role::Detector)).or_else(|| smallest(None))
}
/// Build, 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 = crate::form_sibling(canonical, want);
if want != Form::F32 && !path.is_file() {
return Err(format!(
"no {} form of {}",
want.file_tag().unwrap_or("f32"),
canonical.display()
));
}
let bytes = std::fs::read(&path).map_err(|e| e.to_string())?;
let started = Instant::now();
let mut session = crate::session::build_probe(rung, role, &bytes, cfg)
.map_err(|e| first_line(&e.to_string()))?;
log::info!(
"inference: {} session built in {:.1} s",
rung.label(),
started.elapsed().as_secs_f64()
);
// Zeros for every input the model declares, by name — the denoiser
// takes two (mosaic and σ), and a probe that fed only the first failed
// every rung and left it on the CPU.
let feeds: Vec<(String, Vec<usize>)> = session
.inputs()
.iter()
.map(|i| {
let shape = i
.dtype()
.tensor_shape()
.ok_or("model input is not a tensor")?
.iter()
.map(|&d| if d > 0 { d as usize } else { 1 })
.collect();
Ok((i.name().to_string(), shape))
})
.collect::<Result<_, &str>>()?;
let run = |session: &mut ort::session::Session| -> Result<f64, String> {
let mut inputs: Vec<(String, ort::session::SessionInputValue)> = Vec::new();
for (name, shape) in &feeds {
let zeros = vec![0f32; shape.iter().product()];
let t = ort::value::Tensor::from_array((shape.clone(), zeros))
.map_err(|e| e.to_string())?;
inputs.push((name.clone(), t.into()));
}
let t = Instant::now();
let out = session.run(inputs).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))
}
/// The part of a provider's error a person can act on. ONNX Runtime's
/// begin with a source path and a C++ template signature; the words —
/// "CUDA failure 999: unknown error", "FAIL : Failed to load library" —
/// come after, and the settings row has room for one line of them.
fn first_line(s: &str) -> String {
let line = s.lines().next().unwrap_or("");
let start = ["failure", "FAIL :", "Error:", "error:"]
.iter()
.filter_map(|m| line.find(m))
.min()
.unwrap_or(0);
line[start..].chars().take(200).collect()
}
/// Everything a change of which should re-probe: the runtime, where it
/// came from and which providers sit beside it, 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(),
providers_beside(path)
)
}
},
device_identity(),
];
for (role, form, bytes) in &cfg.embedded {
parts.push(format!(
"{role:?} embedded {form:?} {:016x}",
crate::engines::hash(bytes)
));
}
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}"));
for form in [Form::Int8, Form::A16W8, Form::A16W16] {
if let Ok(b) = std::fs::read(crate::form_sibling(path, form)) {
parts.push(format!(
"{role:?} {form:?} {:016x}",
crate::engines::hash(&b)
));
}
}
}
parts.join("\n")
}
/// The `libonnxruntime_providers_*.so` files in the runtime's directory.
/// A distribution's CPU-only and ROCm builds are the same version at the
/// same path; the provider libraries beside them are what differs.
fn providers_beside(runtime: &Path) -> String {
let Some(dir) = runtime.parent() else {
return String::new();
};
let mut names: Vec<String> = std::fs::read_dir(dir)
.into_iter()
.flatten()
.filter_map(|e| e.ok())
.filter_map(|e| e.file_name().into_string().ok())
.filter(|n| {
n.starts_with("libonnxruntime_providers_") || n.starts_with("onnxruntime_providers_")
})
.collect();
names.sort();
names.join(" ")
}
#[cfg(target_os = "linux")]
fn device_identity() -> String {
// The NVIDIA driver's version line, or the ROCm release the AMD stack
// came from (`rocm-core` writes it; the kernel driver has no version
// of its own). Absent means neither.
if let Some(line) = std::fs::read_to_string("/proc/driver/nvidia/version")
.ok()
.and_then(|s| s.lines().next().map(str::to_string))
{
return line;
}
if let Ok(rocm) = std::fs::read_to_string("/opt/rocm/.info/version") {
return format!("rocm {}", rocm.trim());
}
"no nvidia driver, no rocm".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(target_os = "macos")]
fn device_identity() -> String {
// The chip, and the OS release: CoreML ships with the OS, so a macOS
// update is a new provider as surely as a new driver is on Linux.
format!(
"{} macOS {}",
sysctl("machdep.cpu.brand_string"),
sysctl("kern.osproductversion")
)
}
#[cfg(target_os = "macos")]
fn sysctl(name: &str) -> String {
extern "C" {
fn sysctlbyname(
name: *const std::ffi::c_char,
oldp: *mut std::ffi::c_void,
oldlenp: *mut usize,
newp: *mut std::ffi::c_void,
newlen: usize,
) -> i32;
}
let name = std::ffi::CString::new(name).unwrap();
let mut buf = [0u8; 256];
let mut len = buf.len();
// SAFETY: libSystem's documented call; `len` is the buffer's size in and
// the string's length, with its terminator, out.
let rc = unsafe {
sysctlbyname(
name.as_ptr(),
buf.as_mut_ptr().cast(),
&mut len,
std::ptr::null_mut(),
0,
)
};
if rc != 0 {
return String::new();
}
let s = &buf[..len.min(buf.len())];
String::from_utf8_lossy(s.strip_suffix(&[0]).unwrap_or(s)).into_owned()
}
#[cfg(not(any(target_os = "linux", target_os = "android", target_os = "macos")))]
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);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn a_cache_dir(name: &str) -> Config {
let dir = std::env::temp_dir().join(format!("dr-attempt-{}-{name}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
Config {
cache_dir: dir,
..Config::default()
}
}
/// What a launch that died inside `what` leaves behind.
fn died_inside(cfg: &Config, what: &str, launches: u32) {
std::fs::create_dir_all(&cfg.cache_dir).unwrap();
std::fs::write(cfg.cache_dir.join("attempt"), format!("{what}\t{launches}")).unwrap();
}
#[test]
fn a_finished_attempt_leaves_no_trace() {
let cfg = a_cache_dir("finished");
assert_eq!(attempt(&cfg, "probe CoreML", || 7), Ok(7));
assert!(!cfg.cache_dir.join("attempt").exists());
}
#[test]
fn one_death_is_forgiven_and_two_are_not() {
let cfg = a_cache_dir("strikes");
died_inside(&cfg, "probe CoreML", 1);
assert_eq!(attempt(&cfg, "probe CoreML", || 7), Ok(7));
died_inside(&cfg, "probe CoreML", 2);
let mut ran = false;
assert!(attempt(&cfg, "probe CoreML", || ran = true).is_err());
assert!(!ran, "a refused attempt must not run");
}
#[test]
fn another_attempts_deaths_do_not_count() {
let cfg = a_cache_dir("other");
died_inside(&cfg, "probe TensorRT", 2);
assert_eq!(attempt(&cfg, "probe CUDA", || 7), Ok(7));
}
/// The tablet's cache after 0.22.0's first launch, as 0.22.0 wrote it:
/// no `attempts`, the Hexagon "rejected", the CPU selected.
const TABLET: &str = r#"{"fingerprint":"f","rung":"Cpu","reason":"Hexagon NPU 28.5 ms, slower than the CPU's 19.4 ms","compiled":[],"failed":[["Hexagon","28.5 ms, slower than the CPU's 19.4 ms"]]}"#;
#[test]
fn a_fall_back_to_the_cpu_is_probed_again_a_few_times() {
let mut cache: Cache = serde_json::from_str(TABLET).unwrap();
assert_eq!(cache.attempts, 0, "a 0.22.0 cache reads as never retried");
assert_eq!(reuse(&cache, "f"), Reuse::Again(0));
cache.attempts = RETRIES - 1;
assert_eq!(reuse(&cache, "f"), Reuse::Again(RETRIES - 1));
cache.attempts = RETRIES;
assert_eq!(reuse(&cache, "f"), Reuse::Keep, "then it is kept");
}
#[test]
fn an_accelerator_chosen_or_a_cpu_only_device_is_kept() {
let mut cache: Cache = serde_json::from_str(TABLET).unwrap();
cache.rung = Some(Rung::Hexagon);
assert_eq!(reuse(&cache, "f"), Reuse::Keep);
cache.rung = Some(Rung::Cpu);
cache.failed.clear();
assert_eq!(
reuse(&cache, "f"),
Reuse::Keep,
"nothing failed: the only rung"
);
assert_eq!(reuse(&cache, "other"), Reuse::Fresh);
}
}