Add dr-inference-engine and route every model session through it
One crate names the runtime, the providers and the devices; dr-face and dr-segment ask it for a session by role. It hands ort an API table once per process — from a libonnxruntime it dlopens when the app names a directory holding one, otherwise from tract — so the Rust build stays free of C on every target and a package can install the runtime as a file (docs/inference.md §3). Sessions live in a registry behind a Model handle that holds the bytes, not the session: every use refreshes a timestamp and a reaper unloads whatever sat idle past the decay. A scan that runs the detector on each image never lets it go idle; a click in the develop view lets the segmenter go after thirty seconds; a handle used after that reloads, and reloads on a higher rung if a compiled engine has landed meanwhile. The probe walks the platform's ladder by building strict sessions and timing them against the CPU provider, caches the choice against a fingerprint of the runtime, driver, hardware and models, and compiles engines for the selected rung in the background, smallest model first. Nothing in this commit turns the native path on: the apps still run on tract until they call init with a runtime directory.
This commit is contained in:
Generated
+16
-2
@@ -1475,11 +1475,11 @@ dependencies = [
|
|||||||
name = "dr-face"
|
name = "dr-face"
|
||||||
version = "0.12.2"
|
version = "0.12.2"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"dr-inference-engine",
|
||||||
"env_logger",
|
"env_logger",
|
||||||
"log",
|
"log",
|
||||||
"ndarray",
|
"ndarray",
|
||||||
"ort",
|
"ort",
|
||||||
"ort-tract",
|
|
||||||
"thiserror 2.0.20",
|
"thiserror 2.0.20",
|
||||||
"zune-jpeg 0.4.21",
|
"zune-jpeg 0.4.21",
|
||||||
]
|
]
|
||||||
@@ -1511,6 +1511,20 @@ dependencies = [
|
|||||||
"wgpu",
|
"wgpu",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "dr-inference-engine"
|
||||||
|
version = "0.12.2"
|
||||||
|
dependencies = [
|
||||||
|
"libloading",
|
||||||
|
"log",
|
||||||
|
"ort",
|
||||||
|
"ort-sys",
|
||||||
|
"ort-tract",
|
||||||
|
"serde",
|
||||||
|
"serde_json",
|
||||||
|
"thiserror 2.0.20",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "dr-ingest"
|
name = "dr-ingest"
|
||||||
version = "0.12.2"
|
version = "0.12.2"
|
||||||
@@ -1584,11 +1598,11 @@ dependencies = [
|
|||||||
name = "dr-segment"
|
name = "dr-segment"
|
||||||
version = "0.12.2"
|
version = "0.12.2"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"dr-inference-engine",
|
||||||
"env_logger",
|
"env_logger",
|
||||||
"log",
|
"log",
|
||||||
"ndarray",
|
"ndarray",
|
||||||
"ort",
|
"ort",
|
||||||
"ort-tract",
|
|
||||||
"thiserror 2.0.20",
|
"thiserror 2.0.20",
|
||||||
"zune-jpeg 0.4.21",
|
"zune-jpeg 0.4.21",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ members = [
|
|||||||
"core/dr-export",
|
"core/dr-export",
|
||||||
"core/dr-face",
|
"core/dr-face",
|
||||||
"core/dr-film",
|
"core/dr-film",
|
||||||
|
"core/dr-inference-engine",
|
||||||
"core/dr-ingest",
|
"core/dr-ingest",
|
||||||
"core/dr-gpu",
|
"core/dr-gpu",
|
||||||
"core/dr-lens",
|
"core/dr-lens",
|
||||||
@@ -46,6 +47,9 @@ dr-export = { path = "core/dr-export" }
|
|||||||
# `features = ["inference"]`.
|
# `features = ["inference"]`.
|
||||||
dr-face = { path = "core/dr-face", default-features = false }
|
dr-face = { path = "core/dr-face", default-features = false }
|
||||||
dr-film = { path = "core/dr-film" }
|
dr-film = { path = "core/dr-film" }
|
||||||
|
# `tract` on by default so a test binary can open a session with nothing
|
||||||
|
# installed; the apps add `native` to look for a runtime file (docs/inference.md §3).
|
||||||
|
dr-inference-engine = { path = "core/dr-inference-engine" }
|
||||||
dr-ingest = { path = "core/dr-ingest" }
|
dr-ingest = { path = "core/dr-ingest" }
|
||||||
dr-gpu = { path = "core/dr-gpu" }
|
dr-gpu = { path = "core/dr-gpu" }
|
||||||
dr-lens = { path = "core/dr-lens" }
|
dr-lens = { path = "core/dr-lens" }
|
||||||
|
|||||||
@@ -9,10 +9,11 @@ license.workspace = true
|
|||||||
thiserror.workspace = true
|
thiserror.workspace = true
|
||||||
log.workspace = true
|
log.workspace = true
|
||||||
|
|
||||||
# Inference. `ort` is the API; **tract is the engine** — see the workspace
|
# Inference. `ort` is the API; **what runs it is `dr-inference-engine`'s
|
||||||
# manifest, and docs/faces.md §3, for why the C++ ONNX Runtime is not linked.
|
# business** — tract, or an ONNX Runtime the app found on disk, on whichever
|
||||||
|
# provider the device has (docs/inference.md). This crate never names either.
|
||||||
ort = { workspace = true, optional = true }
|
ort = { workspace = true, optional = true }
|
||||||
ort-tract = { workspace = true, optional = true }
|
dr-inference-engine = { workspace = true, optional = true }
|
||||||
ndarray = { workspace = true, optional = true }
|
ndarray = { workspace = true, optional = true }
|
||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
@@ -20,7 +21,7 @@ zune-jpeg.workspace = true
|
|||||||
env_logger.workspace = true
|
env_logger.workspace = true
|
||||||
# The M1 probe drives `ort` directly so it can print the raw load error.
|
# The M1 probe drives `ort` directly so it can print the raw load error.
|
||||||
ort = { workspace = true }
|
ort = { workspace = true }
|
||||||
ort-tract = { workspace = true }
|
dr-inference-engine = { workspace = true }
|
||||||
|
|
||||||
[[example]]
|
[[example]]
|
||||||
name = "probe"
|
name = "probe"
|
||||||
@@ -48,4 +49,4 @@ default = []
|
|||||||
# must be testable against synthetic embeddings on a machine with no weights on
|
# must be testable against synthetic embeddings on a machine with no weights on
|
||||||
# it — a test suite that needs a research-licensed download is a test suite
|
# it — a test suite that needs a research-licensed download is a test suite
|
||||||
# that does not run in CI.
|
# that does not run in CI.
|
||||||
inference = ["dep:ort", "dep:ort-tract", "dep:ndarray"]
|
inference = ["dep:ort", "dep:dr-inference-engine", "dep:ndarray"]
|
||||||
|
|||||||
+36
-14
@@ -14,7 +14,8 @@
|
|||||||
|
|
||||||
use ndarray::Array4;
|
use ndarray::Array4;
|
||||||
|
|
||||||
use crate::{install_backend, FaceError};
|
use crate::FaceError;
|
||||||
|
use dr_inference_engine::{Form, Model, Role};
|
||||||
|
|
||||||
/// The graph's input edge, in pixels. See the module note: not configurable.
|
/// The graph's input edge, in pixels. See the module note: not configurable.
|
||||||
pub const INPUT_EDGE: usize = 640;
|
pub const INPUT_EDGE: usize = 640;
|
||||||
@@ -135,7 +136,10 @@ impl Detection {
|
|||||||
|
|
||||||
/// A loaded SCRFD graph.
|
/// A loaded SCRFD graph.
|
||||||
pub struct Detector {
|
pub struct Detector {
|
||||||
session: ort::session::Session,
|
session: Model,
|
||||||
|
/// f32 or int8 — the int8 form finds a different set of faces and is a
|
||||||
|
/// different detector in `model_id` (docs/inference.md §7).
|
||||||
|
form: Form,
|
||||||
/// Feature-map count: 3 for strides {8,16,32}, 4 for {8,16,32,64}.
|
/// Feature-map count: 3 for strides {8,16,32}, 4 for {8,16,32,64}.
|
||||||
///
|
///
|
||||||
/// Discovered from the output count rather than assumed, because both
|
/// Discovered from the output count rather than assumed, because both
|
||||||
@@ -145,18 +149,29 @@ pub struct Detector {
|
|||||||
}
|
}
|
||||||
|
|
||||||
impl Detector {
|
impl Detector {
|
||||||
pub fn from_path(path: impl AsRef<std::path::Path>) -> Result<Self, FaceError> {
|
/// Which form this detector was loaded from.
|
||||||
let bytes = std::fs::read(path).map_err(FaceError::ModelRead)?;
|
pub fn form(&self) -> Form {
|
||||||
Self::from_bytes(&bytes)
|
self.form
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn from_bytes(bytes: &[u8]) -> Result<Self, FaceError> {
|
/// Load the canonical f32 file at `path`, or the form the device's
|
||||||
install_backend();
|
/// backend wants instead — the `.int8.onnx` beside it on a Hexagon —
|
||||||
|
/// which [`Detector::form`] then reports.
|
||||||
|
pub fn from_path(path: impl AsRef<std::path::Path>) -> Result<Self, FaceError> {
|
||||||
|
let (path, form) = dr_inference_engine::resolve_model(Role::Detector, path.as_ref());
|
||||||
|
let bytes = std::fs::read(path).map_err(FaceError::ModelRead)?;
|
||||||
|
Self::from_bytes_in(&bytes, form)
|
||||||
|
}
|
||||||
|
|
||||||
let session = ort::session::Session::builder()
|
/// An f32 graph from memory.
|
||||||
.map_err(FaceError::Inference)?
|
pub fn from_bytes(bytes: &[u8]) -> Result<Self, FaceError> {
|
||||||
.commit_from_memory(bytes)
|
Self::from_bytes_in(bytes, Form::F32)
|
||||||
.map_err(FaceError::Inference)?;
|
}
|
||||||
|
|
||||||
|
fn from_bytes_in(bytes: &[u8], form: Form) -> Result<Self, FaceError> {
|
||||||
|
let model = dr_inference_engine::open(Role::Detector, form, bytes)?;
|
||||||
|
let acquired = model.acquire()?;
|
||||||
|
let session = acquired.lock();
|
||||||
|
|
||||||
let n_out = session.outputs().len();
|
let n_out = session.outputs().len();
|
||||||
if n_out % 3 != 0 || !(9..=12).contains(&n_out) {
|
if n_out % 3 != 0 || !(9..=12).contains(&n_out) {
|
||||||
@@ -191,7 +206,13 @@ impl Detector {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(Self { session, fmc })
|
drop(session);
|
||||||
|
drop(acquired);
|
||||||
|
Ok(Self {
|
||||||
|
session: model,
|
||||||
|
form,
|
||||||
|
fmc,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Stride levels this graph emits.
|
/// Stride levels this graph emits.
|
||||||
@@ -223,8 +244,9 @@ impl Detector {
|
|||||||
let lb = Letterbox::fit(width as f32, height as f32);
|
let lb = Letterbox::fit(width as f32, height as f32);
|
||||||
let input = lb.sample(rgb, width, height);
|
let input = lb.sample(rgb, width, height);
|
||||||
|
|
||||||
let outputs = self
|
let acquired = self.session.acquire()?;
|
||||||
.session
|
let mut session = acquired.lock();
|
||||||
|
let outputs = session
|
||||||
.run(ort::inputs![
|
.run(ort::inputs![
|
||||||
ort::value::Tensor::from_array(input).map_err(FaceError::Inference)?
|
ort::value::Tensor::from_array(input).map_err(FaceError::Inference)?
|
||||||
])
|
])
|
||||||
|
|||||||
+17
-11
@@ -14,7 +14,8 @@ use ndarray::Array4;
|
|||||||
|
|
||||||
use crate::align::{Aligned112, ALIGNED_EDGE};
|
use crate::align::{Aligned112, ALIGNED_EDGE};
|
||||||
use crate::embedding::{normalise, Embedding, ModelId, EMBEDDING_DIM};
|
use crate::embedding::{normalise, Embedding, ModelId, EMBEDDING_DIM};
|
||||||
use crate::{install_backend, FaceError};
|
use crate::FaceError;
|
||||||
|
use dr_inference_engine::{Form, Model, Role};
|
||||||
|
|
||||||
/// What one pass of the embedder produces: the direction, and the length.
|
/// What one pass of the embedder produces: the direction, and the length.
|
||||||
///
|
///
|
||||||
@@ -53,7 +54,7 @@ impl Embedded {
|
|||||||
|
|
||||||
/// A loaded ArcFace graph.
|
/// A loaded ArcFace graph.
|
||||||
pub struct Embedder {
|
pub struct Embedder {
|
||||||
session: ort::session::Session,
|
session: Model,
|
||||||
model: ModelId,
|
model: ModelId,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -64,12 +65,11 @@ impl Embedder {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn from_bytes(bytes: &[u8], model: ModelId) -> Result<Self, FaceError> {
|
pub fn from_bytes(bytes: &[u8], model: ModelId) -> Result<Self, FaceError> {
|
||||||
install_backend();
|
// Always the f32 form: an embedding must compare across devices
|
||||||
|
// (docs/inference.md §7), and the engine pins this role to it.
|
||||||
let session = ort::session::Session::builder()
|
let loaded = dr_inference_engine::open(Role::Embedder, Form::F32, bytes)?;
|
||||||
.map_err(FaceError::Inference)?
|
let acquired = loaded.acquire()?;
|
||||||
.commit_from_memory(bytes)
|
let session = acquired.lock();
|
||||||
.map_err(FaceError::Inference)?;
|
|
||||||
|
|
||||||
// One output, `[1, 512]`. Checked because an ArcFace variant with a
|
// One output, `[1, 512]`. Checked because an ArcFace variant with a
|
||||||
// different embedding width would otherwise be read as a truncated
|
// different embedding width would otherwise be read as a truncated
|
||||||
@@ -90,7 +90,12 @@ impl Embedder {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(Self { session, model })
|
drop(session);
|
||||||
|
drop(acquired);
|
||||||
|
Ok(Self {
|
||||||
|
session: loaded,
|
||||||
|
model,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn model(&self) -> &ModelId {
|
pub fn model(&self) -> &ModelId {
|
||||||
@@ -111,8 +116,9 @@ impl Embedder {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let outputs = self
|
let acquired = self.session.acquire()?;
|
||||||
.session
|
let mut session = acquired.lock();
|
||||||
|
let outputs = session
|
||||||
.run(ort::inputs![
|
.run(ort::inputs![
|
||||||
ort::value::Tensor::from_array(input).map_err(FaceError::Inference)?
|
ort::value::Tensor::from_array(input).map_err(FaceError::Inference)?
|
||||||
])
|
])
|
||||||
|
|||||||
+11
-15
@@ -124,25 +124,21 @@ pub enum FaceError {
|
|||||||
ImageShape { expected: usize, got: usize },
|
ImageShape { expected: usize, got: usize },
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Install tract as `ort`'s backend.
|
|
||||||
///
|
|
||||||
/// Idempotent, and it must happen before any other `ort` call: with
|
|
||||||
/// `alternative-backend` there is no linked runtime to fall back on, so an
|
|
||||||
/// un-set API is a panic rather than a slow path. Same helper as
|
|
||||||
/// `dr-segment::semantic`, for the same reason.
|
|
||||||
#[cfg(feature = "inference")]
|
#[cfg(feature = "inference")]
|
||||||
pub(crate) fn install_backend() {
|
impl From<dr_inference_engine::Error> for FaceError {
|
||||||
use std::sync::Once;
|
fn from(e: dr_inference_engine::Error) -> Self {
|
||||||
static ONCE: Once = Once::new();
|
match e {
|
||||||
ONCE.call_once(|| {
|
dr_inference_engine::Error::Inference(e) => FaceError::Inference(e),
|
||||||
let _ = ort::set_api(ort_tract::api());
|
dr_inference_engine::Error::Io(e) => FaceError::ModelRead(e),
|
||||||
});
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// [`install_backend`] for the M1 probe example, which drives `ort` directly
|
/// Make sure `ort` has a backend, for the M1 probe example, which drives
|
||||||
/// rather than through [`detect::Detector`] so it can report the raw error.
|
/// `ort` directly rather than through [`detect::Detector`] so it can report
|
||||||
|
/// the raw error. Every other path goes through `dr-inference-engine`.
|
||||||
#[cfg(feature = "inference")]
|
#[cfg(feature = "inference")]
|
||||||
#[doc(hidden)]
|
#[doc(hidden)]
|
||||||
pub fn install_backend_for_probe() {
|
pub fn install_backend_for_probe() {
|
||||||
install_backend();
|
dr_inference_engine::ensure_runtime();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,44 @@
|
|||||||
|
[package]
|
||||||
|
name = "dr-inference-engine"
|
||||||
|
version.workspace = true
|
||||||
|
edition.workspace = true
|
||||||
|
rust-version.workspace = true
|
||||||
|
license.workspace = true
|
||||||
|
|
||||||
|
# The one crate that names a runtime, a provider, a vendor library or a
|
||||||
|
# device (docs/inference.md §8). `dr-face` and `dr-segment` ask it for a
|
||||||
|
# session by role and never see which of these answered.
|
||||||
|
|
||||||
|
[dependencies]
|
||||||
|
thiserror.workspace = true
|
||||||
|
log.workspace = true
|
||||||
|
serde.workspace = true
|
||||||
|
serde_json.workspace = true
|
||||||
|
|
||||||
|
# `ort` is the API; what supplies it is decided once per process (§3):
|
||||||
|
# `libonnxruntime` found on disk, or `tract`. Both are behind
|
||||||
|
# `alternative-backend`, so nothing here links C on any target.
|
||||||
|
ort = { workspace = true }
|
||||||
|
ort-tract = { workspace = true, optional = true }
|
||||||
|
# dlopen, and the C types of the table it fetches. Both pure Rust;
|
||||||
|
# `libloading` is already in the tree through wgpu.
|
||||||
|
libloading = { version = "0.8", optional = true }
|
||||||
|
ort-sys = { version = "2.0.0-rc.13", default-features = false, features = ["disable-linking"], optional = true }
|
||||||
|
|
||||||
|
# The NVIDIA rungs exist on the desktop only. These features add `ort`'s
|
||||||
|
# option builders and nothing else — no linking under `alternative-backend` —
|
||||||
|
# but an Android binary has no business carrying even the option names, and
|
||||||
|
# the packaging must never be tempted to (§2, §3.1).
|
||||||
|
[target.'cfg(not(target_os = "android"))'.dependencies]
|
||||||
|
ort = { workspace = true, features = ["cuda", "tensorrt"] }
|
||||||
|
|
||||||
|
[target.'cfg(target_os = "android")'.dependencies]
|
||||||
|
ort = { workspace = true, features = ["qnn"] }
|
||||||
|
|
||||||
|
[features]
|
||||||
|
# The floor: `tract` supplies the API table when no runtime file is found, or
|
||||||
|
# always, in a build without `native`. Tests want this and nothing else.
|
||||||
|
default = ["tract"]
|
||||||
|
tract = ["dep:ort-tract"]
|
||||||
|
# Look for `libonnxruntime` on disk and hand its table to `ort`.
|
||||||
|
native = ["dep:libloading", "dep:ort-sys"]
|
||||||
@@ -0,0 +1,142 @@
|
|||||||
|
//! The API table `ort` runs on, chosen once (docs/inference.md §3).
|
||||||
|
//!
|
||||||
|
//! `ort` with `alternative-backend` links no runtime and asks, on first use,
|
||||||
|
//! for an `OrtApi` — a struct of function pointers. Two things can fill it:
|
||||||
|
//! a `libonnxruntime` this module `dlopen`s, or `ort-tract`. The Rust build
|
||||||
|
//! is identical either way; the difference is whether a file was found.
|
||||||
|
|
||||||
|
use std::path::PathBuf;
|
||||||
|
use std::sync::OnceLock;
|
||||||
|
|
||||||
|
/// What supplied the table.
|
||||||
|
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||||
|
pub enum Runtime {
|
||||||
|
/// Pure Rust, one core, every operator these graphs use. The floor.
|
||||||
|
Tract,
|
||||||
|
/// The C++ ONNX Runtime, loaded from `path`.
|
||||||
|
OnnxRuntime { path: PathBuf, version: String },
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Runtime {
|
||||||
|
pub fn label(&self) -> String {
|
||||||
|
match self {
|
||||||
|
Runtime::Tract => "tract".into(),
|
||||||
|
Runtime::OnnxRuntime { version, .. } => format!("ONNX Runtime {version}"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn is_native(&self) -> bool {
|
||||||
|
matches!(self, Runtime::OnnxRuntime { .. })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
static RUNTIME: OnceLock<Runtime> = OnceLock::new();
|
||||||
|
|
||||||
|
/// The runtime in use; tract until something installs another.
|
||||||
|
pub fn runtime() -> Runtime {
|
||||||
|
RUNTIME.get().cloned().unwrap_or(Runtime::Tract)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Install a table if none is installed yet — tract, since no directories
|
||||||
|
/// were named. What a test or an example gets.
|
||||||
|
pub fn ensure_installed() {
|
||||||
|
if RUNTIME.get().is_none() {
|
||||||
|
install(&[]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Look for `libonnxruntime` in `dirs`, in order, and hand `ort` the first
|
||||||
|
/// table that loads; otherwise tract. Once per process.
|
||||||
|
pub fn install(dirs: &[PathBuf]) -> Runtime {
|
||||||
|
RUNTIME
|
||||||
|
.get_or_init(|| {
|
||||||
|
#[cfg(feature = "native")]
|
||||||
|
for dir in dirs {
|
||||||
|
match load_native(dir) {
|
||||||
|
Ok(rt) => return rt,
|
||||||
|
Err(e) => log::info!("inference: no runtime in {}: {e}", dir.display()),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
#[cfg(not(feature = "native"))]
|
||||||
|
let _ = dirs;
|
||||||
|
install_tract()
|
||||||
|
})
|
||||||
|
.clone()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(feature = "tract")]
|
||||||
|
fn install_tract() -> Runtime {
|
||||||
|
let _ = ort::set_api(ort_tract::api());
|
||||||
|
Runtime::Tract
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(not(feature = "tract"))]
|
||||||
|
fn install_tract() -> Runtime {
|
||||||
|
// A build with neither tract nor a runtime file has nothing to run
|
||||||
|
// models on; every `open` will report the un-set API rather than panic
|
||||||
|
// somewhere deeper.
|
||||||
|
log::error!("inference: no ONNX Runtime found and tract is not compiled in");
|
||||||
|
Runtime::Tract
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(feature = "native")]
|
||||||
|
fn load_native(dir: &std::path::Path) -> Result<Runtime, String> {
|
||||||
|
let name = if cfg!(target_os = "windows") {
|
||||||
|
"onnxruntime.dll"
|
||||||
|
} else if cfg!(any(target_os = "macos", target_os = "ios")) {
|
||||||
|
"libonnxruntime.dylib"
|
||||||
|
} else {
|
||||||
|
"libonnxruntime.so"
|
||||||
|
};
|
||||||
|
// An empty dir means the bare name: the system loader's search, which on
|
||||||
|
// Android includes the APK's own native libraries.
|
||||||
|
let path = if dir.as_os_str().is_empty() {
|
||||||
|
PathBuf::from(name)
|
||||||
|
} else {
|
||||||
|
let p = dir.join(name);
|
||||||
|
if !p.is_file() {
|
||||||
|
return Err("not present".into());
|
||||||
|
}
|
||||||
|
p
|
||||||
|
};
|
||||||
|
|
||||||
|
// SAFETY: the library's initialisers are ONNX Runtime's own; the symbol
|
||||||
|
// is the documented entry point with the documented signature; the table
|
||||||
|
// is copied out and the library handle is leaked, so every pointer in
|
||||||
|
// the copy stays valid for the life of the process.
|
||||||
|
unsafe {
|
||||||
|
let lib = libloading::Library::new(&path).map_err(|e| e.to_string())?;
|
||||||
|
let get_base: libloading::Symbol<
|
||||||
|
unsafe extern "system" fn() -> *const ort_sys::OrtApiBase,
|
||||||
|
> = lib.get(b"OrtGetApiBase\0").map_err(|e| e.to_string())?;
|
||||||
|
let base = get_base();
|
||||||
|
if base.is_null() {
|
||||||
|
return Err("OrtGetApiBase returned null".into());
|
||||||
|
}
|
||||||
|
let version = std::ffi::CStr::from_ptr(((*base).GetVersionString)())
|
||||||
|
.to_string_lossy()
|
||||||
|
.into_owned();
|
||||||
|
let api = ((*base).GetApi)(ort_sys::ORT_API_VERSION);
|
||||||
|
if api.is_null() {
|
||||||
|
return Err(format!(
|
||||||
|
"ONNX Runtime {version} is older than API version {}",
|
||||||
|
ort_sys::ORT_API_VERSION
|
||||||
|
));
|
||||||
|
}
|
||||||
|
if !ort::set_api((*api).clone()) {
|
||||||
|
return Err("an API table was already installed".into());
|
||||||
|
}
|
||||||
|
std::mem::forget(lib);
|
||||||
|
|
||||||
|
// Qualcomm's DSP loader finds the Hexagon skel through this variable,
|
||||||
|
// and only through it; the runtime's own directory is where the APK
|
||||||
|
// put it. Harmless anywhere else.
|
||||||
|
#[cfg(target_os = "android")]
|
||||||
|
if !dir.as_os_str().is_empty() {
|
||||||
|
std::env::set_var("ADSP_LIBRARY_PATH", dir);
|
||||||
|
}
|
||||||
|
|
||||||
|
log::info!("inference: ONNX Runtime {version} from {}", path.display());
|
||||||
|
Ok(Runtime::OnnxRuntime { path, version })
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,104 @@
|
|||||||
|
//! Compiled engines: what a rung builds once per device, and the thread that
|
||||||
|
//! builds them before anyone asks (docs/inference.md §5, §6).
|
||||||
|
//!
|
||||||
|
//! TensorRT keeps its own engine cache keyed by graph hash; QNN writes a
|
||||||
|
//! context model. Both are opaque to this crate, which tracks only *that* a
|
||||||
|
//! model compiled — by the hash of its bytes — so [`crate::open`] can tell a
|
||||||
|
//! request whether to expect the rung or its fallback.
|
||||||
|
|
||||||
|
use std::path::PathBuf;
|
||||||
|
|
||||||
|
use crate::{state, Config, Form, Rung};
|
||||||
|
|
||||||
|
/// 64-bit FNV-1a. A cache key, not a checksum: two model files that collide
|
||||||
|
/// here would have to also be the same size and the same role, and the cost
|
||||||
|
/// of that is a rebuilt engine.
|
||||||
|
pub fn hash(bytes: &[u8]) -> u64 {
|
||||||
|
let mut h = 0xcbf2_9ce4_8422_2325u64;
|
||||||
|
for &b in bytes {
|
||||||
|
h ^= b as u64;
|
||||||
|
h = h.wrapping_mul(0x0000_0100_0000_01b3);
|
||||||
|
}
|
||||||
|
h
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The cache entry for `bytes` compiled on `rung`.
|
||||||
|
pub fn key(rung: Rung, bytes: &[u8]) -> String {
|
||||||
|
format!("{}:{:016x}", rung.label(), hash(bytes))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Where QNN's compiled context for `bytes` lives.
|
||||||
|
pub fn context_path(cfg: &Config, bytes: &[u8]) -> PathBuf {
|
||||||
|
cfg.cache_dir
|
||||||
|
.join("qnn")
|
||||||
|
.join(format!("{:016x}_ctx.onnx", hash(bytes)))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// After the probe: compile every configured model the selected rung can
|
||||||
|
/// take, smallest first, recording each as it lands.
|
||||||
|
pub fn run() {
|
||||||
|
let (rung, cfg) = {
|
||||||
|
let s = state().lock().unwrap();
|
||||||
|
(crate::current_rung(&s), s.config.clone())
|
||||||
|
};
|
||||||
|
if !rung.compiles() {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Smallest first, so the detector — the one that runs per image — is
|
||||||
|
// ready soonest (§6 step 3).
|
||||||
|
let mut jobs: Vec<(crate::Role, PathBuf, u64)> = cfg
|
||||||
|
.models
|
||||||
|
.iter()
|
||||||
|
.filter(|(role, _)| rung.form(*role) != Form::F32 || rung != Rung::Hexagon)
|
||||||
|
.filter_map(|(role, path)| {
|
||||||
|
let (path, form) = crate::resolve_model(*role, path);
|
||||||
|
(form == rung.form(*role)).then(|| {
|
||||||
|
let size = std::fs::metadata(&path).map(|m| m.len()).unwrap_or(0);
|
||||||
|
(*role, path, size)
|
||||||
|
})
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
jobs.sort_by_key(|j| j.2);
|
||||||
|
state().lock().unwrap().wanted = jobs.len();
|
||||||
|
|
||||||
|
for (role, path, _) in jobs {
|
||||||
|
let Ok(bytes) = std::fs::read(&path) else {
|
||||||
|
continue;
|
||||||
|
};
|
||||||
|
let key = key(rung, &bytes);
|
||||||
|
if state().lock().unwrap().cache.compiled.contains(&key) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
log::info!(
|
||||||
|
"inference: compiling {} for {}",
|
||||||
|
path.display(),
|
||||||
|
rung.label()
|
||||||
|
);
|
||||||
|
let started = std::time::Instant::now();
|
||||||
|
match crate::session::build(rung, role, &bytes, &cfg, false) {
|
||||||
|
Ok(session) => {
|
||||||
|
drop(session);
|
||||||
|
let mut s = state().lock().unwrap();
|
||||||
|
s.cache.compiled.insert(key);
|
||||||
|
crate::probe::write_cache(&s.config, &s.cache);
|
||||||
|
log::info!(
|
||||||
|
"inference: {} ready on {} in {:.1} s",
|
||||||
|
path.display(),
|
||||||
|
rung.label(),
|
||||||
|
started.elapsed().as_secs_f64()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
Err(e) => {
|
||||||
|
// This model stays on the fallback; the others still get
|
||||||
|
// their engine. A corrected model file changes the hash and
|
||||||
|
// is retried.
|
||||||
|
log::warn!(
|
||||||
|
"inference: {} will not compile for {}: {e}",
|
||||||
|
path.display(),
|
||||||
|
rung.label()
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,542 @@
|
|||||||
|
//! Which runtime, which provider and which model form — decided once per
|
||||||
|
//! device, and the only crate that knows the answer (docs/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, 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,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 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.
|
||||||
|
#[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,
|
||||||
|
/// Qualcomm's Hexagon NPU through QNN, int8 models only. Android only.
|
||||||
|
Hexagon,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Rung {
|
||||||
|
pub fn label(self) -> &'static str {
|
||||||
|
match self {
|
||||||
|
Rung::Cpu => "CPU",
|
||||||
|
Rung::Cuda => "CUDA",
|
||||||
|
Rung::TensorRt => "TensorRT",
|
||||||
|
Rung::Hexagon => "Hexagon NPU",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 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::Hexagon | 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::Hexagon)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The model form this rung wants for a role.
|
||||||
|
fn form(self, role: Role) -> Form {
|
||||||
|
match (self, role) {
|
||||||
|
(Rung::Hexagon, Role::Embedder) => Form::F32,
|
||||||
|
(Rung::Hexagon, _) => Form::Int8,
|
||||||
|
_ => Form::F32,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 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)>,
|
||||||
|
/// 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),
|
||||||
|
}
|
||||||
|
|
||||||
|
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 => " · 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]>,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 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)
|
||||||
|
}
|
||||||
|
|
||||||
|
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]>) -> 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, bytes),
|
||||||
|
s.config.clone(),
|
||||||
|
)
|
||||||
|
};
|
||||||
|
let key = format!("{role:?}:{}", engines::key(rung, bytes));
|
||||||
|
|
||||||
|
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 engine load is long enough
|
||||||
|
// that another role's acquire should not wait on it.
|
||||||
|
let session = session::build(rung, role, bytes, &cfg, false)?;
|
||||||
|
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)>,
|
||||||
|
}
|
||||||
|
|
||||||
|
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(),
|
||||||
|
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.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);
|
||||||
|
acquire(role, form, &bytes)?;
|
||||||
|
Ok(Model { role, form, bytes })
|
||||||
|
}
|
||||||
|
|
||||||
|
/// 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, bytes: &[u8]) -> Rung {
|
||||||
|
let mut rung = selected;
|
||||||
|
if 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(rung, bytes)) {
|
||||||
|
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/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_status_line_reads_as_the_floor_before_init() {
|
||||||
|
let s = status();
|
||||||
|
assert_eq!(s.rung, Rung::Cpu);
|
||||||
|
assert!(s.line().starts_with("CPU"), "{}", s.line());
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,279 @@
|
|||||||
|
//! Walk the ladder, once, by building real sessions (docs/inference.md §4).
|
||||||
|
//!
|
||||||
|
//! A rung is taken when a strict 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 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];
|
||||||
|
#[cfg(not(target_os = "android"))]
|
||||||
|
let all = [Rung::TensorRt, Rung::Cuda];
|
||||||
|
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 configured model: the detector on every device shipped
|
||||||
|
/// today, and a ~2 MB graph is the cheapest real test of a provider.
|
||||||
|
fn probe_model(cfg: &Config) -> Option<(Role, PathBuf)> {
|
||||||
|
cfg.models
|
||||||
|
.iter()
|
||||||
|
.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))
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Build strictly, 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, true)
|
||||||
|
.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))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn first_line(s: &str) -> String {
|
||||||
|
s.lines().next().unwrap_or("").chars().take(160).collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Everything a change of which should re-probe: the runtime and where it
|
||||||
|
/// came from, 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()),
|
||||||
|
},
|
||||||
|
device_identity(),
|
||||||
|
];
|
||||||
|
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")
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(target_os = "linux")]
|
||||||
|
fn device_identity() -> String {
|
||||||
|
// The NVIDIA driver's version line; absent means no NVIDIA driver.
|
||||||
|
std::fs::read_to_string("/proc/driver/nvidia/version")
|
||||||
|
.ok()
|
||||||
|
.and_then(|s| s.lines().next().map(str::to_string))
|
||||||
|
.unwrap_or_else(|| "no nvidia driver".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);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
//! One session builder per rung (docs/inference.md §2, §7, §9).
|
||||||
|
|
||||||
|
use ort::session::builder::GraphOptimizationLevel;
|
||||||
|
use ort::session::Session;
|
||||||
|
|
||||||
|
use crate::{Config, Role, Rung};
|
||||||
|
|
||||||
|
/// Build a session for `bytes` on `rung`.
|
||||||
|
///
|
||||||
|
/// `strict` is the probe's flag: with it, a provider that would hand any
|
||||||
|
/// node to the CPU fails the build instead, so "the session built" means
|
||||||
|
/// "the provider took the graph" and not "the provider registered" (§4).
|
||||||
|
pub fn build(
|
||||||
|
rung: Rung,
|
||||||
|
role: Role,
|
||||||
|
bytes: &[u8],
|
||||||
|
cfg: &Config,
|
||||||
|
strict: bool,
|
||||||
|
) -> ort::Result<Session> {
|
||||||
|
let mut b = Session::builder()?
|
||||||
|
.with_optimization_level(GraphOptimizationLevel::Level3)?
|
||||||
|
.with_intra_threads(threads(cfg))?;
|
||||||
|
if strict {
|
||||||
|
b = b.with_config_entry("session.disable_cpu_ep_fallback", "1")?;
|
||||||
|
}
|
||||||
|
// A Hexagon session loads the compiled context when there is one and
|
||||||
|
// compiles it from the model when there is not; the engine thread is
|
||||||
|
// what makes the second case rare (§6).
|
||||||
|
let context = (rung == Rung::Hexagon).then(|| crate::engines::context_path(cfg, bytes));
|
||||||
|
let ready = context.as_ref().is_some_and(|p| p.is_file());
|
||||||
|
b = providers(
|
||||||
|
b,
|
||||||
|
rung,
|
||||||
|
role,
|
||||||
|
cfg,
|
||||||
|
if ready { None } else { context.as_deref() },
|
||||||
|
)?;
|
||||||
|
match (ready, context) {
|
||||||
|
(true, Some(path)) => b.commit_from_file(path),
|
||||||
|
_ => b.commit_from_memory(bytes),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The intra-op pool: what the config says, else the cores less two for
|
||||||
|
/// the compositor and the decoder (§9). tract ignores it.
|
||||||
|
fn threads(cfg: &Config) -> usize {
|
||||||
|
if cfg.threads > 0 {
|
||||||
|
return cfg.threads;
|
||||||
|
}
|
||||||
|
std::thread::available_parallelism()
|
||||||
|
.map(|n| n.get().saturating_sub(2).max(1))
|
||||||
|
.unwrap_or(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(not(target_os = "android"))]
|
||||||
|
fn providers(
|
||||||
|
b: ort::session::builder::SessionBuilder,
|
||||||
|
rung: Rung,
|
||||||
|
role: Role,
|
||||||
|
cfg: &Config,
|
||||||
|
_generate_context: Option<&std::path::Path>,
|
||||||
|
) -> ort::Result<ort::session::builder::SessionBuilder> {
|
||||||
|
use ort::ep;
|
||||||
|
match rung {
|
||||||
|
Rung::Cpu => Ok(b),
|
||||||
|
Rung::Cuda => {
|
||||||
|
Ok(b.with_execution_providers([ep::CUDA::default().build().error_on_failure()])?)
|
||||||
|
}
|
||||||
|
Rung::TensorRt => {
|
||||||
|
let cache = cfg.cache_dir.join("tensorrt");
|
||||||
|
let _ = std::fs::create_dir_all(&cache);
|
||||||
|
let cache = cache.to_string_lossy().into_owned();
|
||||||
|
// fp16 for everything but the embedder, whose comparability
|
||||||
|
// across devices is worth more than its 0.2 ms (§7). The
|
||||||
|
// workspace cap keeps the develop view's tiles on the card
|
||||||
|
// (NFR-RES-2). CUDA behind it takes any node TensorRT declines.
|
||||||
|
Ok(b.with_execution_providers([
|
||||||
|
ep::TensorRT::default()
|
||||||
|
.with_fp16(role != Role::Embedder)
|
||||||
|
.with_engine_cache(true)
|
||||||
|
.with_engine_cache_path(&cache)
|
||||||
|
.with_timing_cache(true)
|
||||||
|
.with_timing_cache_path(&cache)
|
||||||
|
.with_max_workspace_size(512 << 20)
|
||||||
|
.build()
|
||||||
|
.error_on_failure(),
|
||||||
|
ep::CUDA::default().build(),
|
||||||
|
])?)
|
||||||
|
}
|
||||||
|
Rung::Hexagon => unreachable!("the Hexagon rung is not on a desktop ladder"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(target_os = "android")]
|
||||||
|
fn providers(
|
||||||
|
b: ort::session::builder::SessionBuilder,
|
||||||
|
rung: Rung,
|
||||||
|
_role: Role,
|
||||||
|
_cfg: &Config,
|
||||||
|
generate_context: Option<&std::path::Path>,
|
||||||
|
) -> ort::Result<ort::session::builder::SessionBuilder> {
|
||||||
|
use ort::ep;
|
||||||
|
match rung {
|
||||||
|
Rung::Cpu => Ok(b),
|
||||||
|
Rung::Hexagon => {
|
||||||
|
// The HTP compiles the graph once per device (0.8–1.7 s here).
|
||||||
|
// With `ep.context_enable` ONNX Runtime writes the compiled
|
||||||
|
// context beside the probe cache; the next session loads that
|
||||||
|
// file as its model and skips the compile (§5).
|
||||||
|
let mut b = b;
|
||||||
|
if let Some(ctx) = generate_context {
|
||||||
|
let _ = std::fs::create_dir_all(ctx.parent().unwrap());
|
||||||
|
b = b
|
||||||
|
.with_config_entry("ep.context_enable", "1")?
|
||||||
|
.with_config_entry("ep.context_file_path", ctx.to_string_lossy())?
|
||||||
|
.with_config_entry("ep.context_embed_mode", "0")?;
|
||||||
|
}
|
||||||
|
// Quantise/dequantise at the graph's edges stay on the NPU too,
|
||||||
|
// so a strict build is a whole-graph build.
|
||||||
|
Ok(b.with_execution_providers([ep::QNN::default()
|
||||||
|
.with_backend_path("libQnnHtp.so")
|
||||||
|
.with_performance_mode(ep::qnn::PerformanceMode::Burst)
|
||||||
|
.with_offload_graph_io_quantization(false)
|
||||||
|
.build()
|
||||||
|
.error_on_failure()])?)
|
||||||
|
}
|
||||||
|
Rung::Cuda | Rung::TensorRt => unreachable!("no NVIDIA rung on Android"),
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -11,10 +11,11 @@ build = "build.rs"
|
|||||||
thiserror.workspace = true
|
thiserror.workspace = true
|
||||||
log.workspace = true
|
log.workspace = true
|
||||||
|
|
||||||
# Inference. `ort` is the API; **tract is the engine** — see the workspace
|
# Inference. `ort` is the API; **what runs it is `dr-inference-engine`'s
|
||||||
# manifest for why the C++ ONNX Runtime is not linked here.
|
# business** — tract, or an ONNX Runtime the app found on disk, on whichever
|
||||||
|
# provider the device has (docs/inference.md). This crate never names either.
|
||||||
ort = { workspace = true, optional = true }
|
ort = { workspace = true, optional = true }
|
||||||
ort-tract = { workspace = true, optional = true }
|
dr-inference-engine = { workspace = true, optional = true }
|
||||||
ndarray = { workspace = true, optional = true }
|
ndarray = { workspace = true, optional = true }
|
||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
@@ -35,7 +36,7 @@ default = ["semantic", "embedded-model"]
|
|||||||
# Separable because the watershed half is genuinely independent of it: with
|
# Separable because the watershed half is genuinely independent of it: with
|
||||||
# this off, `dr-segment` is a pure-CPU graph algorithm crate with no model to
|
# this off, `dr-segment` is a pure-CPU graph algorithm crate with no model to
|
||||||
# carry, which is what the headless hierarchy tests want.
|
# carry, which is what the headless hierarchy tests want.
|
||||||
semantic = ["dep:ort", "dep:ort-tract", "dep:ndarray"]
|
semantic = ["dep:ort", "dep:dr-inference-engine", "dep:ndarray"]
|
||||||
|
|
||||||
# Compile the weights into the binary.
|
# Compile the weights into the binary.
|
||||||
#
|
#
|
||||||
|
|||||||
@@ -98,3 +98,13 @@ pub enum SegmentError {
|
|||||||
#[error("category descriptor: {0}")]
|
#[error("category descriptor: {0}")]
|
||||||
CategoryDescriptor(String),
|
CategoryDescriptor(String),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(feature = "semantic")]
|
||||||
|
impl From<dr_inference_engine::Error> for SegmentError {
|
||||||
|
fn from(e: dr_inference_engine::Error) -> Self {
|
||||||
|
match e {
|
||||||
|
dr_inference_engine::Error::Inference(e) => SegmentError::Inference(e),
|
||||||
|
dr_inference_engine::Error::Io(e) => SegmentError::ModelRead(e),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ use ndarray::ArrayView3;
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
use crate::semantic::INPUT_EDGE;
|
use crate::semantic::INPUT_EDGE;
|
||||||
use crate::semantic::{install_backend, Letterbox, Window};
|
use crate::semantic::{Letterbox, Window};
|
||||||
use crate::SegmentError;
|
use crate::SegmentError;
|
||||||
|
|
||||||
/// Classes in the ADE20K vocabulary the scene model was trained on.
|
/// Classes in the ADE20K vocabulary the scene model was trained on.
|
||||||
@@ -84,7 +84,7 @@ pub struct Category {
|
|||||||
|
|
||||||
/// The scene model, and the categories it has been told to report.
|
/// The scene model, and the categories it has been told to report.
|
||||||
pub struct SceneModel {
|
pub struct SceneModel {
|
||||||
session: ort::session::Session,
|
session: dr_inference_engine::Model,
|
||||||
categories: Vec<Category>,
|
categories: Vec<Category>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -132,12 +132,12 @@ impl SceneModel {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn from_bytes(bytes: &[u8], categories: Vec<Category>) -> Result<Self, SegmentError> {
|
pub fn from_bytes(bytes: &[u8], categories: Vec<Category>) -> Result<Self, SegmentError> {
|
||||||
install_backend();
|
// f32, as for `SemanticModel`; see there.
|
||||||
|
let session = dr_inference_engine::open(
|
||||||
let session = ort::session::Session::builder()
|
dr_inference_engine::Role::Scene,
|
||||||
.map_err(SegmentError::Inference)?
|
dr_inference_engine::Form::F32,
|
||||||
.commit_from_memory(bytes)
|
bytes,
|
||||||
.map_err(SegmentError::Inference)?;
|
)?;
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
session,
|
session,
|
||||||
@@ -170,13 +170,9 @@ impl SceneModel {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
// Split the borrow: `run` needs the session mutably while
|
let categories = &self.categories;
|
||||||
// `marginalise` needs the categories, and going through `self` for
|
let acquired = self.session.acquire()?;
|
||||||
// both at once is what the borrow checker objects to.
|
let mut session = acquired.lock();
|
||||||
let Self {
|
|
||||||
session,
|
|
||||||
categories,
|
|
||||||
} = self;
|
|
||||||
|
|
||||||
let window = Window {
|
let window = Window {
|
||||||
x: 0.0,
|
x: 0.0,
|
||||||
|
|||||||
@@ -194,7 +194,7 @@ impl Instance {
|
|||||||
/// Holds an `ort` session, so it is neither `Clone` nor cheap to build —
|
/// Holds an `ort` session, so it is neither `Clone` nor cheap to build —
|
||||||
/// construct once and keep it. Loading is ~50 ms.
|
/// construct once and keep it. Loading is ~50 ms.
|
||||||
pub struct SemanticModel {
|
pub struct SemanticModel {
|
||||||
session: ort::session::Session,
|
session: dr_inference_engine::Model,
|
||||||
classes: Vec<Arc<str>>,
|
classes: Vec<Arc<str>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -229,15 +229,14 @@ impl SemanticModel {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub fn from_bytes(bytes: &[u8], classes: Vec<Arc<str>>) -> Result<Self, SegmentError> {
|
pub fn from_bytes(bytes: &[u8], classes: Vec<Arc<str>>) -> Result<Self, SegmentError> {
|
||||||
// Idempotent, and it must happen before any other `ort` call: with
|
// The f32 graph on whatever the device's backend is. An int8 form
|
||||||
// `alternative-backend` there is no linked runtime to fall back on, so
|
// for the Hexagon waits on docs/inference.md §10 M7 — the mask
|
||||||
// an un-set API is a panic rather than a slow path.
|
// boundary has to be measured before it moves.
|
||||||
install_backend();
|
let session = dr_inference_engine::open(
|
||||||
|
dr_inference_engine::Role::Segmenter,
|
||||||
let session = ort::session::Session::builder()
|
dr_inference_engine::Form::F32,
|
||||||
.map_err(SegmentError::Inference)?
|
bytes,
|
||||||
.commit_from_memory(bytes)
|
)?;
|
||||||
.map_err(SegmentError::Inference)?;
|
|
||||||
|
|
||||||
Ok(Self { session, classes })
|
Ok(Self { session, classes })
|
||||||
}
|
}
|
||||||
@@ -332,8 +331,9 @@ impl SemanticModel {
|
|||||||
let letterbox = Letterbox::fit(window.w, window.h);
|
let letterbox = Letterbox::fit(window.w, window.h);
|
||||||
let input = letterbox.sample(rgb, width, height, window);
|
let input = letterbox.sample(rgb, width, height, window);
|
||||||
|
|
||||||
let outputs = self
|
let acquired = self.session.acquire()?;
|
||||||
.session
|
let mut session = acquired.lock();
|
||||||
|
let outputs = session
|
||||||
.run(ort::inputs![
|
.run(ort::inputs![
|
||||||
ort::value::Tensor::from_array(input).map_err(SegmentError::Inference)?
|
ort::value::Tensor::from_array(input).map_err(SegmentError::Inference)?
|
||||||
])
|
])
|
||||||
@@ -674,18 +674,6 @@ fn steps(extent: f32, edge: f32, stride: f32) -> usize {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Point `ort` at tract, exactly once per process.
|
|
||||||
pub(crate) fn install_backend() {
|
|
||||||
use std::sync::Once;
|
|
||||||
static ONCE: Once = Once::new();
|
|
||||||
ONCE.call_once(|| {
|
|
||||||
// Returns false if an API was already installed, which is not an error
|
|
||||||
// — it means something else got here first, and there is only one
|
|
||||||
// backend compiled in for it to have chosen.
|
|
||||||
let _ = ort::set_api(ort_tract::api());
|
|
||||||
});
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Read the class list written beside the model by `tools/export-seg-model.sh`.
|
/// Read the class list written beside the model by `tools/export-seg-model.sh`.
|
||||||
///
|
///
|
||||||
/// A deliberately small hand-rolled reader for a flat array of strings, rather
|
/// A deliberately small hand-rolled reader for a flat array of strings, rather
|
||||||
|
|||||||
Reference in New Issue
Block a user