The macOS ladder was the CPU provider alone, with CoreML listed as a gap. It is now CoreML, then the CPU, then tract — unmeasured, since nobody here has a Mac, and safe to ship unmeasured because the probe's clock rejects a CoreML slower than the CPU and `attempt` refuses one that crashes. - `Rung::CoreMl`, a compiling rung like TensorRT: an ML Program with every compute unit allowed, falling back to the CPU until each model's program is built. The embedder stays on the CPU, as on the Hexagon (§7). - The cache is one directory per model and runtime version. CoreML keys a model committed from memory on its input and node names, not its weights (ONNX Runtime 1.29, coreml_execution_provider.cc), so two exports of one architecture would otherwise share a program. - The fingerprint on macOS is the chip and the OS release, which ships CoreML. - The desktop looks for the runtime in the bundle's Contents/Frameworks and Homebrew's prefixes; fetch-desktop-runtime.sh on a Mac downloads ONNX Runtime 1.29.0 for Apple silicon, which carries CoreML. docs/dev/macos.md says what exists, how to build it, and which log lines to ask a Mac user for.
714 lines
25 KiB
Rust
714 lines
25 KiB
Rust
//! Which runtime, which provider and which model form — decided once per
|
||
//! device, and the only crate that knows the answer (docs/dev/inference.md).
|
||
//!
|
||
//! Consumers ask for a session by [`Role`] and get `ort`'s `Session` back;
|
||
//! what built it — tract on one core, ONNX Runtime's CPU pool, a TensorRT
|
||
//! engine, a MIGraphX program, the Hexagon — is this crate's business and
|
||
//! shows up in [`status`] for the settings row and nowhere else.
|
||
//!
|
||
//! The shape follows §3 of the spec: `ort` links nothing (`alternative-backend`),
|
||
//! and the first call hands it an API table from either a `libonnxruntime`
|
||
//! found on disk or from `tract`. That choice is once per process, because
|
||
//! `ort::set_api` is; everything after it — which provider, whether an engine
|
||
//! has been compiled yet — is per session and may change between two calls.
|
||
|
||
use std::collections::{BTreeSet, HashMap};
|
||
use std::path::{Path, PathBuf};
|
||
use std::sync::{Arc, Mutex, MutexGuard, OnceLock};
|
||
use std::time::{Duration, Instant};
|
||
|
||
use serde::{Deserialize, Serialize};
|
||
|
||
mod api;
|
||
mod engines;
|
||
mod probe;
|
||
mod session;
|
||
|
||
pub use api::Runtime;
|
||
pub use ort::session::Session;
|
||
|
||
/// What a model is for. The role fixes the precision rule (§7): an embedder
|
||
/// runs in f32 on every rung, a detector may run in fp16 or int8.
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||
pub enum Role {
|
||
Detector,
|
||
Embedder,
|
||
Segmenter,
|
||
Scene,
|
||
/// The dense landmark model behind the eye reading (docs/dev/faces.md §7c).
|
||
Landmarks,
|
||
/// The eye-state and sunglasses classifiers, a few hundred kilobytes.
|
||
EyeClassifier,
|
||
/// XFeat, the panorama keypoint detector (docs/dev/panorama.md).
|
||
Keypoints,
|
||
/// MI-GAN, the panorama border filler (docs/dev/panorama.md §12). Plain
|
||
/// convolutions, so any rung serves it; fp16 on TensorRT and int8 on
|
||
/// the Hexagon are the point of it.
|
||
Inpainter,
|
||
/// The learned demosaic and denoise on the raw mosaic (docs/dev/denoise.md).
|
||
/// fp16 costs it nothing measurable; int8 costs 6–9 dB, because 256
|
||
/// levels cannot hold the shadow steps it exists to recover — so the
|
||
/// Hexagon does not take it.
|
||
Denoiser,
|
||
}
|
||
|
||
/// Which numeric form of a model a session was built from.
|
||
///
|
||
/// `Int8` is a different network from `F32` for a detector — it finds a
|
||
/// different set of faces — which is why [`form_suffix`] exists and why a
|
||
/// caller appends it to `model_id`.
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||
pub enum Form {
|
||
F32,
|
||
Int8,
|
||
}
|
||
|
||
/// A rung of the ladder (§2). Ordered: a user override names the highest rung
|
||
/// the probe may take, and a compiling rung falls back to the one below it
|
||
/// until its engine exists. The order is within a vendor's ladder — a
|
||
/// machine has NVIDIA rungs or an AMD rung, never both — so a ceiling is
|
||
/// read as "no higher than this on whichever ladder the device has".
|
||
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord, Serialize, Deserialize)]
|
||
pub enum Rung {
|
||
/// ONNX Runtime's CPU provider, or tract when no runtime file was found.
|
||
Cpu,
|
||
/// NVIDIA, through the CUDA provider. Desktop only.
|
||
Cuda,
|
||
/// NVIDIA, through a TensorRT engine compiled on this device. Desktop only.
|
||
TensorRt,
|
||
/// AMD, through a MIGraphX program compiled on this device. Desktop
|
||
/// only. ONNX Runtime's ROCm provider, the CUDA provider's twin, was
|
||
/// removed in ONNX Runtime 1.23, so there is no non-compiling AMD rung
|
||
/// to fall back to: this one falls back to the CPU.
|
||
MiGraphX,
|
||
/// Qualcomm's Hexagon NPU through QNN, int8 models only. Android only.
|
||
Hexagon,
|
||
/// Apple, through CoreML: the Neural Engine, the GPU or the CPU, as
|
||
/// CoreML schedules it. macOS only. Compiles an ML Program per model on
|
||
/// first use, so it is a compiling rung with the CPU below it. The
|
||
/// embedder stays on the CPU, as on the Hexagon: the Neural Engine
|
||
/// computes in fp16 (§7).
|
||
CoreMl,
|
||
}
|
||
|
||
impl Rung {
|
||
pub fn label(self) -> &'static str {
|
||
match self {
|
||
Rung::Cpu => "CPU",
|
||
Rung::Cuda => "CUDA",
|
||
Rung::TensorRt => "TensorRT",
|
||
Rung::MiGraphX => "MIGraphX",
|
||
Rung::Hexagon => "Hexagon NPU",
|
||
Rung::CoreMl => "CoreML",
|
||
}
|
||
}
|
||
|
||
/// The rung a request lands on while this one's engine is still being
|
||
/// compiled (§6 step 2).
|
||
fn fallback(self) -> Rung {
|
||
match self {
|
||
Rung::TensorRt => Rung::Cuda,
|
||
Rung::MiGraphX | Rung::Hexagon | Rung::CoreMl | Rung::Cuda | Rung::Cpu => Rung::Cpu,
|
||
}
|
||
}
|
||
|
||
/// Whether a session on this rung needs an engine built first.
|
||
fn compiles(self) -> bool {
|
||
matches!(
|
||
self,
|
||
Rung::TensorRt | Rung::MiGraphX | Rung::Hexagon | Rung::CoreMl
|
||
)
|
||
}
|
||
|
||
/// The model form this rung wants for a role.
|
||
fn form(self, _role: Role) -> Form {
|
||
match self {
|
||
Rung::Hexagon => Form::Int8,
|
||
_ => Form::F32,
|
||
}
|
||
}
|
||
|
||
/// Whether this rung runs `role` at all. The Hexagon takes int8 graphs
|
||
/// only, and the embedder is never int8 (§7) — it runs on the CPU
|
||
/// beside a detector on the NPU, so its vectors compare across devices.
|
||
/// Nor is the denoiser: its int8 form failed the 0.5 dB gate by 6–9 dB
|
||
/// (denoise.md §8), so it runs on the CPU there too. CoreML is kept off
|
||
/// the embedder for the same reason as the Hexagon: the Neural Engine is
|
||
/// fp16, and which unit runs a graph is CoreML's choice.
|
||
fn serves(self, role: Role) -> bool {
|
||
match self {
|
||
Rung::Hexagon => !matches!(role, Role::Embedder | Role::Denoiser),
|
||
Rung::CoreMl => role != Role::Embedder,
|
||
_ => true,
|
||
}
|
||
}
|
||
}
|
||
|
||
/// How long a session outlives its last use unless [`Config::decay`] says
|
||
/// otherwise: long enough for the next click, short enough that a session's
|
||
/// GPU or NPU memory does not sit under the develop view for long.
|
||
pub const DEFAULT_DECAY: Duration = Duration::from_secs(30);
|
||
|
||
/// What [`init`] is told once, at launch.
|
||
#[derive(Clone, Debug, Default)]
|
||
pub struct Config {
|
||
/// Where to look for `libonnxruntime`, in order. An empty path means "the
|
||
/// bare library name through the system loader", which is how the APK's
|
||
/// own copy is found on Android.
|
||
pub runtime_dirs: Vec<PathBuf>,
|
||
/// Probe cache and compiled engines (§4, §5). Disposable.
|
||
pub cache_dir: PathBuf,
|
||
/// The canonical model files on this device, so engines can be compiled
|
||
/// ahead of the first request for them.
|
||
pub models: Vec<(Role, PathBuf)>,
|
||
/// Models compiled into the binary, for the same reason.
|
||
pub embedded: Vec<(Role, &'static [u8])>,
|
||
/// The highest rung the user allows; `None` is "the best that works".
|
||
pub ceiling: Option<Rung>,
|
||
/// ONNX Runtime's intra-op pool; 0 picks from the core count.
|
||
pub threads: usize,
|
||
/// How long an unused session stays loaded. Zero means the default.
|
||
pub decay: Duration,
|
||
}
|
||
|
||
/// One line for the settings row, and the numbers behind the progress row.
|
||
#[derive(Clone, Debug)]
|
||
pub struct Status {
|
||
pub runtime: Runtime,
|
||
/// The rung selected, or the floor while the probe is still running.
|
||
pub rung: Rung,
|
||
/// Why — "probe passed", or the failure that demoted the rung above.
|
||
pub reason: String,
|
||
pub probing: bool,
|
||
/// Engines compiled and engines wanted, for a compiling rung; `(0, 0)`
|
||
/// otherwise.
|
||
pub engines: (usize, usize),
|
||
/// Every rung above the selected one that was tried, and why it lost.
|
||
pub failed: Vec<(Rung, String)>,
|
||
}
|
||
|
||
impl Status {
|
||
/// "Hexagon NPU · int8 · ONNX Runtime 1.29" — the settings row's text.
|
||
pub fn line(&self) -> String {
|
||
let form = match self.rung {
|
||
Rung::Hexagon => " · int8",
|
||
Rung::TensorRt | Rung::MiGraphX => " · fp16",
|
||
_ => "",
|
||
};
|
||
format!("{}{} · {}", self.rung.label(), form, self.runtime.label())
|
||
}
|
||
}
|
||
|
||
/// A model the caller can run, whatever is or is not loaded right now.
|
||
///
|
||
/// Holds the bytes, not a session. [`Model::acquire`] finds the loaded copy
|
||
/// in the registry — shared with every other holder of the same model —
|
||
/// or loads one, and every acquire refreshes the copy's last-used time.
|
||
/// The reaper unloads anything idle for [`Config::decay`]; a scan that runs
|
||
/// the detector on every image never lets it go idle, a click in the
|
||
/// develop view lets the segmenter go after a quiet spell, and a handle
|
||
/// used again after that simply loads again. Nobody states a policy.
|
||
///
|
||
/// The registry key includes the rung, so a reload after a compiled engine
|
||
/// has landed moves up to it by itself (§6 step 4).
|
||
pub struct Model {
|
||
role: Role,
|
||
form: Form,
|
||
bytes: Arc<[u8]>,
|
||
/// `engines::hash` of the bytes, taken once: an acquire per tile of a
|
||
/// border fill must not hash 28 MB each time.
|
||
hash: u64,
|
||
}
|
||
|
||
/// A loaded session, held for one `run` and its output decoding.
|
||
pub struct Acquired {
|
||
entry: Arc<Loaded>,
|
||
}
|
||
|
||
struct Loaded {
|
||
rung: Rung,
|
||
session: Mutex<Session>,
|
||
last_used: Mutex<Instant>,
|
||
}
|
||
|
||
impl Model {
|
||
/// The loaded session, loading it if the reaper took it. Lock it for
|
||
/// one run; a scan and a develop click can want the same detector at
|
||
/// once, and the second waits on the first.
|
||
pub fn acquire(&self) -> Result<Acquired, Error> {
|
||
acquire(self.role, self.form, &self.bytes, self.hash)
|
||
}
|
||
|
||
pub fn form(&self) -> Form {
|
||
self.form
|
||
}
|
||
}
|
||
|
||
impl Acquired {
|
||
pub fn lock(&self) -> MutexGuard<'_, Session> {
|
||
self.entry.session.lock().unwrap_or_else(|e| e.into_inner())
|
||
}
|
||
|
||
/// Where this session runs.
|
||
pub fn rung(&self) -> Rung {
|
||
self.entry.rung
|
||
}
|
||
}
|
||
|
||
impl Drop for Acquired {
|
||
fn drop(&mut self) {
|
||
// The clock starts when the use ends, not when it began: a long run
|
||
// is not idle time.
|
||
*self.entry.last_used.lock().unwrap() = Instant::now();
|
||
}
|
||
}
|
||
|
||
type Registry = HashMap<String, Arc<Loaded>>;
|
||
|
||
static REGISTRY: OnceLock<Mutex<Registry>> = OnceLock::new();
|
||
|
||
fn registry() -> &'static Mutex<Registry> {
|
||
REGISTRY.get_or_init(|| {
|
||
std::thread::Builder::new()
|
||
.name("inference-reaper".into())
|
||
.spawn(|| loop {
|
||
std::thread::sleep(Duration::from_secs(5));
|
||
release_idle();
|
||
})
|
||
.expect("spawn inference reaper");
|
||
Mutex::new(HashMap::new())
|
||
})
|
||
}
|
||
|
||
fn acquire(role: Role, form: Form, bytes: &Arc<[u8]>, hash: u64) -> Result<Acquired, Error> {
|
||
api::ensure_installed();
|
||
let (rung, cfg) = {
|
||
let s = state().lock().unwrap();
|
||
let selected = current_rung(&s);
|
||
(
|
||
effective_rung(&s, selected, role, form, hash),
|
||
s.config.clone(),
|
||
)
|
||
};
|
||
let key = format!("{role:?}:{}", engines::key_of(rung, hash));
|
||
|
||
if let Some(entry) = registry().lock().unwrap().get(&key).cloned() {
|
||
*entry.last_used.lock().unwrap() = Instant::now();
|
||
return Ok(Acquired { entry });
|
||
}
|
||
|
||
// Built outside the registry lock: a TensorRT or MIGraphX engine load
|
||
// is long enough that another role's acquire should not wait on it.
|
||
let session = session::build(rung, role, bytes, &cfg)?;
|
||
log::debug!("inference: {role:?} loaded on {}", rung.label());
|
||
let entry = Arc::new(Loaded {
|
||
rung,
|
||
session: Mutex::new(session),
|
||
last_used: Mutex::new(Instant::now()),
|
||
});
|
||
let mut reg = registry().lock().unwrap();
|
||
// Two acquires raced; keep the first, drop this one.
|
||
let entry = reg.entry(key).or_insert_with(|| entry.clone()).clone();
|
||
Ok(Acquired { entry })
|
||
}
|
||
|
||
/// Unload every session idle for longer than the decay. The reaper does
|
||
/// this every five seconds. A session in use survives until its run ends:
|
||
/// the `Acquired` holds it, the registry merely forgets it.
|
||
pub fn release_idle() {
|
||
let decay = match state().lock().unwrap().config.decay {
|
||
Duration::ZERO => DEFAULT_DECAY,
|
||
d => d,
|
||
};
|
||
let now = Instant::now();
|
||
registry()
|
||
.lock()
|
||
.unwrap()
|
||
.retain(|_, e| now.duration_since(*e.last_used.lock().unwrap()) < decay);
|
||
}
|
||
|
||
/// Unload every session now, decay or not — what a low-memory signal
|
||
/// asks for. Sessions mid-run finish first.
|
||
pub fn release_all() {
|
||
registry().lock().unwrap().clear();
|
||
}
|
||
|
||
/// Unload every session of `role` now — "I am done segmenting".
|
||
pub fn unload(role: Role) {
|
||
let prefix = format!("{role:?}:");
|
||
registry()
|
||
.lock()
|
||
.unwrap()
|
||
.retain(|k, _| !k.starts_with(&prefix));
|
||
}
|
||
|
||
/// How many sessions are loaded, for the settings row and the tests.
|
||
pub fn loaded() -> usize {
|
||
registry().lock().unwrap().len()
|
||
}
|
||
|
||
#[derive(Debug, thiserror::Error)]
|
||
pub enum Error {
|
||
#[error(transparent)]
|
||
Inference(#[from] ort::Error),
|
||
#[error("reading model: {0}")]
|
||
Io(#[from] std::io::Error),
|
||
}
|
||
|
||
/// What the probe writes and the next launch reads (§4 step 3).
|
||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||
struct Cache {
|
||
/// Runtime, driver, hardware and model identity; any change re-probes.
|
||
fingerprint: String,
|
||
rung: Option<Rung>,
|
||
reason: String,
|
||
/// Model hashes whose engine exists on disk, per compiling rung.
|
||
compiled: BTreeSet<String>,
|
||
/// Rungs that failed under this fingerprint, and why. Not retried until
|
||
/// the fingerprint changes: a wedged driver must not cost every launch
|
||
/// thirty seconds.
|
||
failed: Vec<(Rung, String)>,
|
||
/// Engine keys whose compile the process died inside, launch after
|
||
/// launch (`probe::attempt`). Left on the fallback until the
|
||
/// fingerprint changes. Defaulted, so a cache from before this field
|
||
/// still reads.
|
||
#[serde(default)]
|
||
refused: BTreeSet<String>,
|
||
}
|
||
|
||
struct State {
|
||
config: Config,
|
||
cache: Cache,
|
||
probing: bool,
|
||
wanted: usize,
|
||
}
|
||
|
||
static STATE: OnceLock<Mutex<State>> = OnceLock::new();
|
||
|
||
fn state() -> &'static Mutex<State> {
|
||
STATE.get_or_init(|| {
|
||
Mutex::new(State {
|
||
config: Config::default(),
|
||
cache: Cache::default(),
|
||
probing: false,
|
||
wanted: 0,
|
||
})
|
||
})
|
||
}
|
||
|
||
/// Choose the runtime and start the probe. Idempotent; the first call wins.
|
||
///
|
||
/// Returns at once: the probe and any engine compilation run on their own
|
||
/// low-priority thread, and every request meanwhile is served by the floor
|
||
/// (§4). Never blocks the first frame.
|
||
pub fn init(config: Config) {
|
||
let runtime = api::install(&config.runtime_dirs);
|
||
{
|
||
let mut s = state().lock().unwrap();
|
||
if s.probing || s.cache.rung.is_some() {
|
||
return;
|
||
}
|
||
s.config = config;
|
||
s.probing = true;
|
||
}
|
||
log::info!("inference: runtime {}", runtime.label());
|
||
std::thread::Builder::new()
|
||
.name("inference-probe".into())
|
||
.spawn(move || {
|
||
probe::run(runtime);
|
||
engines::run();
|
||
})
|
||
.expect("spawn inference probe");
|
||
}
|
||
|
||
/// Make sure `ort` has an API table, for code that drives `ort` directly.
|
||
/// [`open`] does this itself; only the M1 probe example needs it by name.
|
||
pub fn ensure_runtime() {
|
||
api::ensure_installed();
|
||
}
|
||
|
||
/// The line for the settings row.
|
||
pub fn status() -> Status {
|
||
let s = state().lock().unwrap();
|
||
let rung = current_rung(&s);
|
||
Status {
|
||
runtime: api::runtime(),
|
||
rung,
|
||
reason: s.cache.reason.clone(),
|
||
// Only what explains the selection: on an AMD machine the NVIDIA
|
||
// rungs "not enabled in this build" say nothing about why MIGraphX
|
||
// was taken. With the floor selected, everything tried is above it.
|
||
failed: s
|
||
.cache
|
||
.failed
|
||
.iter()
|
||
.filter(|(r, _)| *r > rung)
|
||
.cloned()
|
||
.collect(),
|
||
probing: s.probing,
|
||
engines: if rung.compiles() {
|
||
(s.cache.compiled.len(), s.wanted)
|
||
} else {
|
||
(0, 0)
|
||
},
|
||
}
|
||
}
|
||
|
||
fn current_rung(s: &State) -> Rung {
|
||
if s.probing {
|
||
Rung::Cpu
|
||
} else {
|
||
s.cache.rung.unwrap_or(Rung::Cpu)
|
||
}
|
||
}
|
||
|
||
/// The file to load for `role` under the current selection, and its form.
|
||
///
|
||
/// A rung that wants int8 gets the `.int8.onnx` sibling of the canonical file
|
||
/// if it exists; otherwise the canonical file, on the rung's fallback. A
|
||
/// caller adds [`form_suffix`] to the `model_id` it records.
|
||
pub fn resolve_model(role: Role, canonical: &Path) -> (PathBuf, Form) {
|
||
let rung = current_rung(&state().lock().unwrap());
|
||
if rung.serves(role) && rung.form(role) == Form::Int8 {
|
||
let sibling = int8_sibling(canonical);
|
||
if sibling.is_file() {
|
||
return (sibling, Form::Int8);
|
||
}
|
||
}
|
||
(canonical.to_path_buf(), Form::F32)
|
||
}
|
||
|
||
fn int8_sibling(canonical: &Path) -> PathBuf {
|
||
let stem = canonical
|
||
.file_stem()
|
||
.map(|s| s.to_string_lossy().into_owned())
|
||
.unwrap_or_default();
|
||
canonical.with_file_name(format!("{stem}.int8.onnx"))
|
||
}
|
||
|
||
/// What a form appends to a detector's `model_id` (§7).
|
||
pub fn form_suffix(form: Form) -> &'static str {
|
||
match form {
|
||
Form::F32 => "",
|
||
Form::Int8 => "_i8",
|
||
}
|
||
}
|
||
|
||
/// A handle on the model `bytes` in `role`.
|
||
///
|
||
/// Loads it once here, so a graph the runtime rejects fails at
|
||
/// construction and not on the first image; what happens to that session
|
||
/// afterwards is the registry's business (see [`Model`]).
|
||
///
|
||
/// Works without [`init`] — a test, or the examples — by installing tract
|
||
/// and using the CPU rung, which is exactly what every consumer did before
|
||
/// this crate existed.
|
||
pub fn open(role: Role, form: Form, bytes: &[u8]) -> Result<Model, Error> {
|
||
let bytes: Arc<[u8]> = Arc::from(bytes);
|
||
let hash = engines::hash(&bytes);
|
||
acquire(role, form, &bytes, hash)?;
|
||
Ok(Model {
|
||
role,
|
||
form,
|
||
bytes,
|
||
hash,
|
||
})
|
||
}
|
||
|
||
/// Where a request lands: the selected rung unless the role's precision rule,
|
||
/// the form on offer, or a missing engine says one lower (§6 step 4).
|
||
fn effective_rung(s: &State, selected: Rung, role: Role, form: Form, hash: u64) -> Rung {
|
||
let mut rung = selected;
|
||
if !rung.serves(role) || rung.form(role) != form {
|
||
// The embedder on a Hexagon device, or an f32 detector where the int8
|
||
// sibling was missing: neither can go to the NPU.
|
||
rung = rung.fallback();
|
||
}
|
||
if rung.compiles() && !s.cache.compiled.contains(&engines::key_of(rung, hash)) {
|
||
rung = rung.fallback();
|
||
}
|
||
rung
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
/// The registry is one per process, so these run one at a time.
|
||
static SERIAL: Mutex<()> = Mutex::new(());
|
||
fn serial() -> MutexGuard<'static, ()> {
|
||
SERIAL.lock().unwrap_or_else(|e| e.into_inner())
|
||
}
|
||
|
||
/// The smallest shipped graph, if this checkout has the weights; a test
|
||
/// suite that needs a research-licensed download is one that does not
|
||
/// run in CI (docs/dev/faces.md §3), so absence is a skip.
|
||
fn probe_bytes() -> Option<Vec<u8>> {
|
||
let path = concat!(
|
||
env!("CARGO_MANIFEST_DIR"),
|
||
"/../../models/face/scrfd_500m_640.onnx"
|
||
);
|
||
let bytes = std::fs::read(path).ok()?;
|
||
(bytes.len() > 100_000).then_some(bytes)
|
||
}
|
||
|
||
#[test]
|
||
fn two_handles_on_one_model_share_one_session() {
|
||
let _serial = serial();
|
||
let Some(bytes) = probe_bytes() else { return };
|
||
release_all();
|
||
let a = open(Role::Detector, Form::F32, &bytes).unwrap();
|
||
let b = open(Role::Detector, Form::F32, &bytes).unwrap();
|
||
assert_eq!(loaded(), 1);
|
||
let (x, y) = (a.acquire().unwrap(), b.acquire().unwrap());
|
||
assert!(Arc::ptr_eq(&x.entry, &y.entry));
|
||
}
|
||
|
||
#[test]
|
||
fn a_released_model_reloads_on_its_next_use() {
|
||
let _serial = serial();
|
||
let Some(bytes) = probe_bytes() else { return };
|
||
release_all();
|
||
let model = open(Role::Detector, Form::F32, &bytes).unwrap();
|
||
assert_eq!(loaded(), 1);
|
||
release_all();
|
||
assert_eq!(loaded(), 0);
|
||
let acquired = model.acquire().unwrap();
|
||
assert_eq!(loaded(), 1);
|
||
assert_eq!(acquired.lock().inputs().len(), 1);
|
||
}
|
||
|
||
#[test]
|
||
fn an_idle_session_decays_and_a_used_one_does_not() {
|
||
let _serial = serial();
|
||
let Some(bytes) = probe_bytes() else { return };
|
||
release_all();
|
||
state().lock().unwrap().config.decay = Duration::from_millis(50);
|
||
let model = open(Role::Detector, Form::F32, &bytes).unwrap();
|
||
// Used within the decay: stays.
|
||
std::thread::sleep(Duration::from_millis(30));
|
||
drop(model.acquire().unwrap());
|
||
release_idle();
|
||
assert_eq!(loaded(), 1);
|
||
// Idle past it: goes.
|
||
std::thread::sleep(Duration::from_millis(80));
|
||
release_idle();
|
||
assert_eq!(loaded(), 0);
|
||
state().lock().unwrap().config.decay = Duration::ZERO;
|
||
}
|
||
|
||
#[test]
|
||
fn unload_by_role_leaves_the_other_roles() {
|
||
let _serial = serial();
|
||
let Some(bytes) = probe_bytes() else { return };
|
||
release_all();
|
||
let _d = open(Role::Detector, Form::F32, &bytes).unwrap();
|
||
let _s = open(Role::Segmenter, Form::F32, &bytes).unwrap();
|
||
assert_eq!(loaded(), 2);
|
||
unload(Role::Segmenter);
|
||
assert_eq!(loaded(), 1);
|
||
}
|
||
|
||
#[test]
|
||
fn the_hexagon_never_takes_the_embedder() {
|
||
assert!(!Rung::Hexagon.serves(Role::Embedder));
|
||
assert!(!Rung::Hexagon.serves(Role::Denoiser));
|
||
assert!(Rung::Hexagon.serves(Role::Detector));
|
||
assert_eq!(Rung::Hexagon.form(Role::Detector), Form::Int8);
|
||
// A detector offered in f32 on a Hexagon device lands on the CPU.
|
||
let s = State {
|
||
config: Config::default(),
|
||
cache: Cache {
|
||
rung: Some(Rung::Hexagon),
|
||
..Cache::default()
|
||
},
|
||
probing: false,
|
||
wanted: 0,
|
||
};
|
||
assert_eq!(
|
||
effective_rung(
|
||
&s,
|
||
Rung::Hexagon,
|
||
Role::Embedder,
|
||
Form::F32,
|
||
engines::hash(b"")
|
||
),
|
||
Rung::Cpu
|
||
);
|
||
assert_eq!(
|
||
effective_rung(
|
||
&s,
|
||
Rung::Hexagon,
|
||
Role::Detector,
|
||
Form::F32,
|
||
engines::hash(b"")
|
||
),
|
||
Rung::Cpu
|
||
);
|
||
// An int8 detector whose context is not compiled yet: also the CPU.
|
||
assert_eq!(
|
||
effective_rung(
|
||
&s,
|
||
Rung::Hexagon,
|
||
Role::Detector,
|
||
Form::Int8,
|
||
engines::hash(b"")
|
||
),
|
||
Rung::Cpu
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn coreml_takes_a_compiled_detector_and_never_the_embedder() {
|
||
let hash = engines::hash(b"detector");
|
||
let mut s = State {
|
||
config: Config::default(),
|
||
cache: Cache {
|
||
rung: Some(Rung::CoreMl),
|
||
..Cache::default()
|
||
},
|
||
probing: false,
|
||
wanted: 0,
|
||
};
|
||
let on = |s: &State, role| effective_rung(s, Rung::CoreMl, role, Form::F32, hash);
|
||
// Before its program is compiled the detector waits on the CPU.
|
||
assert_eq!(on(&s, Role::Detector), Rung::Cpu);
|
||
s.cache.compiled.insert(engines::key_of(Rung::CoreMl, hash));
|
||
assert_eq!(on(&s, Role::Detector), Rung::CoreMl);
|
||
// The embedder does not move, compiled or not (§7).
|
||
assert_eq!(on(&s, Role::Embedder), Rung::Cpu);
|
||
}
|
||
|
||
#[test]
|
||
fn the_status_reports_only_the_rungs_above_the_selection() {
|
||
let _serial = serial();
|
||
let failed = vec![
|
||
(Rung::TensorRt, "not enabled".to_string()),
|
||
(Rung::Cuda, "not enabled".to_string()),
|
||
];
|
||
let before = state().lock().unwrap().cache.clone();
|
||
state().lock().unwrap().cache = Cache {
|
||
rung: Some(Rung::MiGraphX),
|
||
failed: failed.clone(),
|
||
..Cache::default()
|
||
};
|
||
// An AMD desktop: the NVIDIA rungs below MIGraphX are not the story.
|
||
assert!(status().failed.is_empty());
|
||
// An NVIDIA desktop on the CUDA provider: TensorRT's failure is.
|
||
state().lock().unwrap().cache.rung = Some(Rung::Cuda);
|
||
assert_eq!(status().failed, vec![failed[0].clone()]);
|
||
// The floor: everything tried explains it.
|
||
state().lock().unwrap().cache.rung = Some(Rung::Cpu);
|
||
assert_eq!(status().failed.len(), 2);
|
||
state().lock().unwrap().cache = before;
|
||
}
|
||
|
||
#[test]
|
||
fn the_status_line_reads_as_the_floor_before_init() {
|
||
let _serial = serial();
|
||
let s = status();
|
||
assert_eq!(s.rung, Rung::Cpu);
|
||
assert!(s.line().starts_with("CPU"), "{}", s.line());
|
||
}
|
||
}
|