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:
+36
-14
@@ -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
@@ -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
@@ -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();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user