Files
DarkRoom/core/dr-inference-engine/src/probe.rs
T
dtourolle 84fade99ec Put the developer docs under docs/dev and index the folder for users first
docs/ had 26 developer documents flat beside the manual, and the two
audiences are very differently sized: most readers want the manual and
the gesture reference, a few want the register, the designs and the
measurements. The manual and gestures.md stay at the top; everything for
someone changing the code moves to docs/dev/, and the two documents that
name their own successors — the v0.1 milestone and the UI-refinement plan
— go to docs/dev/archive/ rather than being deleted, since both are still
cited. docs/README.md is the index, users first.

Every reference follows: code comments, Cargo manifests, the workflows,
the pre-commit hook, the bench and traceability tools (which locate the
repo root by docs/dev/requirements.md now), packaging, the Docker READMEs,
CLAUDE.md, CONTRIBUTING.md and the README. The matrix links one level
deeper and is regenerated. Links out of the moved documents into the tree
gain a level; a link checker over every Markdown file finds none broken.
2026-09-20 21:16:03 +02:00

341 lines
12 KiB
Rust

//! 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> {
#[cfg(target_os = "android")]
let all = [Rung::Hexagon];
// A desktop has one vendor's GPU; the other vendor's providers are
// "not enabled in this build" or a library that fails to load, and
// either answer arrives in milliseconds.
#[cfg(not(target_os = "android"))]
let all = [Rung::TensorRt, Rung::Cuda, Rung::MiGraphX];
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 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 the
/// detector is the role the int8 forms exist for — the eye classifiers are
/// smaller still, and a Hexagon probed with one would fail for want of a
/// form nobody ships.
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 = 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).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))
}
/// 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, bytes) in &cfg.embedded {
parts.push(format!(
"{role:?} embedded {: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}"));
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")
}
/// 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(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);
}
}
}