Files
DarkRoom/core/dr-face/src/embed.rs
T
dtourolle d15c41e699 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.
2026-09-19 16:02:37 +02:00

150 lines
5.4 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! ArcFace / MobileFaceNet inference (docs/faces.md §6).
//!
//! Takes an aligned crop and returns 512 L2-normalised floats. The alignment is
//! not optional and cannot be skipped by accident: [`Embedder::embed`] takes an
//! [`Aligned112`], which only [`crate::align::warp`] can construct.
//!
//! # The graph must have a fixed batch
//!
//! `w600k_mbf.onnx` declares its batch dimension as the literal `dim_param`
//! `"None"`, and tract fails to analyse the first Conv because of it. Pinned to
//! 1 by `tools/fix-face-model-shapes.sh`, it loads and runs.
use ndarray::Array4;
use crate::align::{Aligned112, ALIGNED_EDGE};
use crate::embedding::{normalise, Embedding, ModelId, EMBEDDING_DIM};
use crate::FaceError;
use dr_inference_engine::{Form, Model, Role};
/// What one pass of the embedder produces: the direction, and the length.
///
/// Two fields rather than a `quality` on [`Embedding`], because every other
/// holder of an `Embedding` relies on it being unit length and compares by
/// dot product; the length is a separate fact about the same face, and it is
/// stored separately too.
#[derive(Debug, Clone, PartialEq)]
pub struct Embedded {
pub embedding: Embedding,
/// L2 norm of the raw model output.
///
/// The model's own opinion of how recognisable the crop was — see
/// [`crate::embedding::MIN_GALLERY_QUALITY`] for what it means and where
/// it is used.
pub quality: f32,
}
impl Embedded {
/// Storage form: the **raw** vector, `512 × f16`.
///
/// Not the unit vector. The length is the quality, and a store that held
/// only the direction would have thrown it away at the one moment it could
/// be known — which is what this crate used to do. Readers re-normalise
/// ([`Embedding::from_f16_bytes`]), so every comparison is still a dot
/// product, and [`crate::embedding::read_f16_bytes`] gives the length back
/// to a reader that wants it.
///
/// f16 costs nothing extra at this scale: its precision is relative, so a
/// component of a vector of length 20 is kept to the same three figures as
/// the same component scaled to length 1.
pub fn to_f16_bytes(&self) -> Vec<u8> {
self.embedding.to_f16_bytes_scaled(self.quality)
}
}
/// A loaded ArcFace graph.
pub struct Embedder {
session: Model,
model: ModelId,
}
impl Embedder {
pub fn from_path(path: impl AsRef<std::path::Path>, model: ModelId) -> Result<Self, FaceError> {
let bytes = std::fs::read(path).map_err(FaceError::ModelRead)?;
Self::from_bytes(&bytes, model)
}
pub fn from_bytes(bytes: &[u8], model: ModelId) -> Result<Self, FaceError> {
// 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
// one, and 512 is baked into the catalog's BLOB width.
let out = session.outputs().first().ok_or(FaceError::WrongModel {
expected: "ArcFace",
detail: "model has no outputs".into(),
})?;
let last = out.dtype().tensor_shape().and_then(|d| d.last().copied());
if last != Some(EMBEDDING_DIM as i64) {
return Err(FaceError::WrongModel {
expected: "ArcFace",
detail: format!(
"output '{}' is {:?}-wide, expected {EMBEDDING_DIM}",
out.name(),
last
),
});
}
drop(session);
drop(acquired);
Ok(Self {
session: loaded,
model,
})
}
pub fn model(&self) -> &ModelId {
&self.model
}
/// Embed one aligned face.
pub fn embed(&mut self, face: &Aligned112) -> Result<Embedded, FaceError> {
// `(x·255 − 127.5) / 128` — see the `/128` note in `detect::Letterbox`.
let px = face.pixels();
let mut input = Array4::<f32>::zeros((1, 3, ALIGNED_EDGE, ALIGNED_EDGE));
for y in 0..ALIGNED_EDGE {
for x in 0..ALIGNED_EDGE {
for c in 0..3 {
let v = px[(y * ALIGNED_EDGE + x) * 3 + c];
input[[0, c, y, x]] = (v * 255.0 - 127.5) / 128.0;
}
}
}
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)?
])
.map_err(FaceError::Inference)?;
let (_, data) = outputs[0]
.try_extract_tensor::<f32>()
.map_err(FaceError::Inference)?;
if data.len() < EMBEDDING_DIM {
return Err(FaceError::WrongModel {
expected: "ArcFace",
detail: format!("got {} values, expected {EMBEDDING_DIM}", data.len()),
});
}
let mut v = Box::new([0.0_f32; EMBEDDING_DIM]);
v.copy_from_slice(&data[..EMBEDDING_DIM]);
let quality = normalise(&mut v);
Ok(Embedded {
embedding: Embedding {
model: self.model.clone(),
v,
},
quality,
})
}
}