//! 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) -> Vec { #[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| { 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> { 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 = 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 { 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::() .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 = 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 { 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); } } }