One warm-up run and the median of three: on the Iris Xe the OpenVINO rung lost to the CPU on the smallest detector in two probes of three, because an idle integrated GPU takes a few runs to raise its clock — warm, it is 5.8 ms against 9.5. Three warm-ups and the median of seven took it in five probes of five (5.6–7.2 ms against 7.7–13.1). The extra runs cost tens of milliseconds, once per fingerprint.
616 lines
22 KiB
Rust
616 lines
22 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> {
|
||
// 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);
|
||
// The rung that tried and lost, not the first one the runtime was
|
||
// never built with: "WebGPU 150 ms, slower than the CPU" says why
|
||
// this device is on the CPU, "TensorRT not enabled" does not.
|
||
let tried = cache
|
||
.failed
|
||
.iter()
|
||
.rev()
|
||
.find(|(_, why)| !why.contains("in this build"))
|
||
.or(cache.failed.first());
|
||
cache.reason = match tried {
|
||
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, warm up, then time seven runs; the median in milliseconds and,
|
||
/// for a compiling rung, the cache key of the engine this just built.
|
||
///
|
||
/// Three warm-ups, not one: an idle integrated GPU takes a few runs to
|
||
/// raise its clock. With one, the Iris Xe's OpenVINO lost to the CPU on
|
||
/// the smallest detector in two probes of three, where warm it is 5.8 ms
|
||
/// against 9.5 (§1.6). The smallest detector is a GPU's worst case; the
|
||
/// clock must not also be.
|
||
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)
|
||
};
|
||
for _ in 0..3 {
|
||
run(&mut session)?;
|
||
}
|
||
let mut times = (0..7)
|
||
.map(|_| run(&mut session))
|
||
.collect::<Result<Vec<_>, _>>()?;
|
||
times.sort_by(|a, b| a.partial_cmp(b).unwrap());
|
||
let key = rung.compiles().then(|| crate::engines::key(rung, &bytes));
|
||
Ok((times[times.len() / 2], 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, the ROCm release the AMD stack came
|
||
// from (`rocm-core` writes it; the kernel driver has no version of its
|
||
// own), and the OpenCL drivers registered — OpenVINO reaches the GPU
|
||
// through one, and installing Intel's is what makes the Iris Xe a rung.
|
||
let mut parts = Vec::new();
|
||
if let Some(line) = std::fs::read_to_string("/proc/driver/nvidia/version")
|
||
.ok()
|
||
.and_then(|s| s.lines().next().map(str::to_string))
|
||
{
|
||
parts.push(line);
|
||
}
|
||
if let Ok(rocm) = std::fs::read_to_string("/opt/rocm/.info/version") {
|
||
parts.push(format!("rocm {}", rocm.trim()));
|
||
}
|
||
let mut icds: Vec<String> = std::fs::read_dir("/etc/OpenCL/vendors")
|
||
.into_iter()
|
||
.flatten()
|
||
.filter_map(|e| e.ok())
|
||
.map(|e| e.file_name().to_string_lossy().into_owned())
|
||
.collect();
|
||
icds.sort();
|
||
if !icds.is_empty() {
|
||
parts.push(format!("opencl {}", icds.join(" ")));
|
||
}
|
||
if parts.is_empty() {
|
||
"no nvidia driver, no rocm, no opencl".into()
|
||
} else {
|
||
parts.join("; ")
|
||
}
|
||
}
|
||
|
||
#[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")]
|
||
pub(crate) 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);
|
||
}
|
||
}
|