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:
+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)?
|
||||
])
|
||||
|
||||
Reference in New Issue
Block a user