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:
2026-09-19 16:02:37 +02:00
parent caf21bea64
commit d15c41e699
16 changed files with 1368 additions and 90 deletions
+6 -5
View File
@@ -9,10 +9,11 @@ license.workspace = true
thiserror.workspace = true
log.workspace = true
# Inference. `ort` is the API; **tract is the engine** — see the workspace
# manifest, and docs/faces.md §3, for why the C++ ONNX Runtime is not linked.
# Inference. `ort` is the API; **what runs it is `dr-inference-engine`'s
# 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-tract = { workspace = true, optional = true }
dr-inference-engine = { workspace = true, optional = true }
ndarray = { workspace = true, optional = true }
[dev-dependencies]
@@ -20,7 +21,7 @@ zune-jpeg.workspace = true
env_logger.workspace = true
# The M1 probe drives `ort` directly so it can print the raw load error.
ort = { workspace = true }
ort-tract = { workspace = true }
dr-inference-engine = { workspace = true }
[[example]]
name = "probe"
@@ -48,4 +49,4 @@ default = []
# 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
# 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
View File
@@ -14,7 +14,8 @@
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.
pub const INPUT_EDGE: usize = 640;
@@ -135,7 +136,10 @@ impl Detection {
/// A loaded SCRFD graph.
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}.
///
/// Discovered from the output count rather than assumed, because both
@@ -145,18 +149,29 @@ pub struct Detector {
}
impl Detector {
pub fn from_path(path: impl AsRef<std::path::Path>) -> Result<Self, FaceError> {
let bytes = std::fs::read(path).map_err(FaceError::ModelRead)?;
Self::from_bytes(&bytes)
/// Which form this detector was loaded from.
pub fn form(&self) -> Form {
self.form
}
pub fn from_bytes(bytes: &[u8]) -> Result<Self, FaceError> {
install_backend();
/// Load the canonical f32 file at `path`, or the form the device's
/// 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()
.map_err(FaceError::Inference)?
.commit_from_memory(bytes)
.map_err(FaceError::Inference)?;
/// An f32 graph from memory.
pub fn from_bytes(bytes: &[u8]) -> Result<Self, FaceError> {
Self::from_bytes_in(bytes, Form::F32)
}
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();
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.
@@ -223,8 +244,9 @@ impl Detector {
let lb = Letterbox::fit(width as f32, height as f32);
let input = lb.sample(rgb, width, height);
let outputs = self
.session
let acquired = self.session.acquire()?;
let mut session = acquired.lock();
let outputs = session
.run(ort::inputs![
ort::value::Tensor::from_array(input).map_err(FaceError::Inference)?
])
+17 -11
View File
@@ -14,7 +14,8 @@ use ndarray::Array4;
use crate::align::{Aligned112, ALIGNED_EDGE};
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.
///
@@ -53,7 +54,7 @@ impl Embedded {
/// A loaded ArcFace graph.
pub struct Embedder {
session: ort::session::Session,
session: Model,
model: ModelId,
}
@@ -64,12 +65,11 @@ impl Embedder {
}
pub fn from_bytes(bytes: &[u8], model: ModelId) -> Result<Self, FaceError> {
install_backend();
let session = ort::session::Session::builder()
.map_err(FaceError::Inference)?
.commit_from_memory(bytes)
.map_err(FaceError::Inference)?;
// Always the f32 form: an embedding must compare across devices
// (docs/inference.md §7), and the engine pins this role to it.
let loaded = dr_inference_engine::open(Role::Embedder, Form::F32, bytes)?;
let acquired = loaded.acquire()?;
let session = acquired.lock();
// One output, `[1, 512]`. Checked because an ArcFace variant with a
// 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 {
@@ -111,8 +116,9 @@ impl Embedder {
}
}
let outputs = self
.session
let acquired = self.session.acquire()?;
let mut session = acquired.lock();
let outputs = session
.run(ort::inputs![
ort::value::Tensor::from_array(input).map_err(FaceError::Inference)?
])
+11 -15
View File
@@ -124,25 +124,21 @@ pub enum FaceError {
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")]
pub(crate) fn install_backend() {
use std::sync::Once;
static ONCE: Once = Once::new();
ONCE.call_once(|| {
let _ = ort::set_api(ort_tract::api());
});
impl From<dr_inference_engine::Error> for FaceError {
fn from(e: dr_inference_engine::Error) -> Self {
match e {
dr_inference_engine::Error::Inference(e) => FaceError::Inference(e),
dr_inference_engine::Error::Io(e) => FaceError::ModelRead(e),
}
}
}
/// [`install_backend`] for the M1 probe example, which drives `ort` directly
/// rather than through [`detect::Detector`] so it can report the raw error.
/// Make sure `ort` has a backend, for the M1 probe example, which drives
/// `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")]
#[doc(hidden)]
pub fn install_backend_for_probe() {
install_backend();
dr_inference_engine::ensure_runtime();
}