Detect, align and embed faces with SCRFD and MobileFaceNet
Ports the pipeline from the C++ reference in ../scene-actor-extraction (MIT, same author). End to end on real portraits it separates identities the way the reference's fitted calibration says it should: 0.596 between distinct photographs of one person, 0.05 between different people, either side of MBF's 0.267 boundary. Three things are structural rather than incidental: Aligned112 can only be built by align::warp, so Embedder::embed cannot be handed an unaligned bounding-box crop. That mistake yields 512 plausible unit-norm numbers and no error, so the type system refuses it instead. Embedding carries its ModelId and cosine() returns None across models, because a cross-model similarity is the one mistake that produces plausible garbage rather than a failure. The model-free half -- alignment, embedding arithmetic, f16 storage -- sits outside the inference feature and is covered by 11 tests that need no weights on the machine. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,104 @@
|
||||
//! 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::{install_backend, FaceError};
|
||||
|
||||
/// A loaded ArcFace graph.
|
||||
pub struct Embedder {
|
||||
session: ort::session::Session,
|
||||
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> {
|
||||
install_backend();
|
||||
|
||||
let session = ort::session::Session::builder()
|
||||
.map_err(FaceError::Inference)?
|
||||
.commit_from_memory(bytes)
|
||||
.map_err(FaceError::Inference)?;
|
||||
|
||||
// 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),
|
||||
});
|
||||
}
|
||||
|
||||
Ok(Self { session, model })
|
||||
}
|
||||
|
||||
pub fn model(&self) -> &ModelId {
|
||||
&self.model
|
||||
}
|
||||
|
||||
/// Embed one aligned face.
|
||||
pub fn embed(&mut self, face: &Aligned112) -> Result<Embedding, 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 outputs = self
|
||||
.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]);
|
||||
normalise(&mut v);
|
||||
|
||||
Ok(Embedding {
|
||||
model: self.model.clone(),
|
||||
v,
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user