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:
@@ -26,6 +26,10 @@ ort-tract = { workspace = true }
|
||||
name = "probe"
|
||||
required-features = ["inference"]
|
||||
|
||||
[[example]]
|
||||
name = "faces"
|
||||
required-features = ["inference"]
|
||||
|
||||
[features]
|
||||
# Nothing on by default, and in particular **no `embedded-model`**: the weights
|
||||
# are not a build input and never become one (docs/faces.md §2.2). A feature
|
||||
|
||||
@@ -0,0 +1,124 @@
|
||||
//! Detect, align and embed the faces in a JPEG.
|
||||
//!
|
||||
//! The thing worth looking at is whether the landmarks land on a real
|
||||
//! photograph — the same reason `dr-segment` has `examples/detect.rs`.
|
||||
//!
|
||||
//! cargo run -p dr-face --features inference --example faces -- \
|
||||
//! DET.onnx EMB.onnx photo.jpg [photo.jpg ...]
|
||||
//!
|
||||
//! The models must have had their input dims frozen first; see
|
||||
//! `tools/fix-face-model-shapes.sh` and docs/faces.md §12 M1.
|
||||
|
||||
use std::time::Instant;
|
||||
|
||||
use dr_face::{align, DetectOptions, Detector, Embedder, ModelId};
|
||||
|
||||
fn main() {
|
||||
env_logger::init();
|
||||
|
||||
let args: Vec<String> = std::env::args().skip(1).collect();
|
||||
if args.len() < 3 {
|
||||
eprintln!("usage: faces DET.onnx EMB.onnx IMAGE.jpg [IMAGE.jpg ...]");
|
||||
std::process::exit(2);
|
||||
}
|
||||
|
||||
let t = Instant::now();
|
||||
let mut detector = Detector::from_path(&args[0]).expect("load detector");
|
||||
let mut embedder =
|
||||
Embedder::from_path(&args[1], ModelId::new("w600k_mbf")).expect("load embedder");
|
||||
println!(
|
||||
"loaded both models in {:?} (strides {:?})",
|
||||
t.elapsed(),
|
||||
detector.strides()
|
||||
);
|
||||
|
||||
let opts = DetectOptions::default();
|
||||
let mut all = Vec::new();
|
||||
|
||||
for path in &args[2..] {
|
||||
let (rgb, w, h) = match load_jpeg(path) {
|
||||
Ok(v) => v,
|
||||
Err(e) => {
|
||||
println!("{path}: {e}");
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let t = Instant::now();
|
||||
let dets = detector.detect(&rgb, w, h, &opts).expect("detect");
|
||||
let detect_ms = t.elapsed().as_secs_f64() * 1e3;
|
||||
|
||||
println!("\n{path} ({w}×{h}) {} face(s) in {detect_ms:.0} ms", dets.len());
|
||||
|
||||
for (i, d) in dets.iter().enumerate() {
|
||||
let Some(aligned) = align::warp(&rgb, w, h, &d.landmarks) else {
|
||||
println!(" [{i}] degenerate landmarks, skipped");
|
||||
continue;
|
||||
};
|
||||
let t = Instant::now();
|
||||
let emb = embedder.embed(&aligned).expect("embed");
|
||||
let embed_ms = t.elapsed().as_secs_f64() * 1e3;
|
||||
|
||||
println!(
|
||||
" [{i}] conf {:.3} box {:.0},{:.0} {:.0}×{:.0} crop_px {:.0} embed {embed_ms:.0} ms",
|
||||
d.confidence,
|
||||
d.bbox.0,
|
||||
d.bbox.1,
|
||||
d.width(),
|
||||
d.height(),
|
||||
aligned.source_px(),
|
||||
);
|
||||
all.push((path.clone(), i, emb));
|
||||
}
|
||||
}
|
||||
|
||||
// Every pair, so the numbers can be eyeballed against the expectation that
|
||||
// faces from one identity's folder score high and everything else low.
|
||||
if all.len() > 1 {
|
||||
println!("\ncosine similarity");
|
||||
for i in 0..all.len() {
|
||||
for j in i + 1..all.len() {
|
||||
let cos = all[i].2.cosine(&all[j].2).expect("same model");
|
||||
println!(
|
||||
" {:.4} {}#{} vs {}#{}",
|
||||
cos,
|
||||
short(&all[i].0),
|
||||
all[i].1,
|
||||
short(&all[j].0),
|
||||
all[j].1
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn short(path: &str) -> String {
|
||||
let p = std::path::Path::new(path);
|
||||
let file = p.file_name().unwrap_or_default().to_string_lossy();
|
||||
match p.parent().and_then(|d| d.file_name()) {
|
||||
Some(dir) => format!("{}/{file}", dir.to_string_lossy()),
|
||||
None => file.into_owned(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Decode to the tightly packed `f32` RGB `0.0..=1.0` the crate expects.
|
||||
fn load_jpeg(path: &str) -> Result<(Vec<f32>, usize, usize), String> {
|
||||
let bytes = std::fs::read(path).map_err(|e| e.to_string())?;
|
||||
let mut dec = zune_jpeg::JpegDecoder::new(&bytes);
|
||||
let px = dec.decode().map_err(|e| e.to_string())?;
|
||||
let info = dec.info().ok_or("no jpeg header")?;
|
||||
let (w, h) = (info.width as usize, info.height as usize);
|
||||
|
||||
let rgb: Vec<f32> = match px.len() / (w * h) {
|
||||
3 => px.iter().map(|&v| v as f32 / 255.0).collect(),
|
||||
1 => px
|
||||
.iter()
|
||||
.flat_map(|&v| {
|
||||
let g = v as f32 / 255.0;
|
||||
[g, g, g]
|
||||
})
|
||||
.collect(),
|
||||
n => return Err(format!("{n} components per pixel, expected 1 or 3")),
|
||||
};
|
||||
Ok((rgb, w, h))
|
||||
}
|
||||
@@ -0,0 +1,360 @@
|
||||
//! Five-point face alignment (docs/faces.md §5).
|
||||
//!
|
||||
//! ArcFace embeddings are trained on faces warped to a canonical 112×112
|
||||
//! arrangement. Feeding the model a plain bounding-box crop *works* — it
|
||||
//! produces 512 numbers, they are unit-norm, and cosine similarities between
|
||||
//! them look entirely reasonable. They are just much worse, and nothing in the
|
||||
//! system reports it.
|
||||
//!
|
||||
//! That is the whole reason this module exists, and the reason [`Aligned112`]
|
||||
//! is a newtype only [`warp`] can construct: the mistake is not one a reviewer
|
||||
//! catches, so the type system catches it instead.
|
||||
//!
|
||||
//! Model-free, so it builds and tests without the `inference` feature.
|
||||
|
||||
/// Canonical landmark positions for a 112×112 ArcFace crop.
|
||||
///
|
||||
/// # The naming is a trap; the order is not
|
||||
///
|
||||
/// Point 0 sits at x=38 on a 112-wide canvas — left of centre *in the image*,
|
||||
/// which is the subject's **right** eye. Both namings are in circulation and
|
||||
/// they are opposite, so the array is written in the detector's order and the
|
||||
/// comment says whose left is whose:
|
||||
///
|
||||
/// ```text
|
||||
/// 0 subject's right eye (image-left)
|
||||
/// 1 subject's left eye (image-right)
|
||||
/// 2 nose tip
|
||||
/// 3 subject's right mouth corner
|
||||
/// 4 subject's left mouth corner
|
||||
/// ```
|
||||
///
|
||||
/// SCRFD emits its five points in this same order, so the correct amount of
|
||||
/// reordering between detector and template is **none**. A detector with a
|
||||
/// different order carries its own permutation beside its model id rather than
|
||||
/// this constant growing an assumption.
|
||||
pub const ARCFACE_TEMPLATE: [(f32, f32); 5] = [
|
||||
(38.2946, 51.6963),
|
||||
(73.5318, 51.5014),
|
||||
(56.0252, 71.7366),
|
||||
(41.5493, 92.3655),
|
||||
(70.7299, 92.2041),
|
||||
];
|
||||
|
||||
/// Edge of the aligned crop, in pixels. Fixed by the embedder's input.
|
||||
pub const ALIGNED_EDGE: usize = 112;
|
||||
|
||||
/// A face warped to [`ARCFACE_TEMPLATE`], ready for the embedder.
|
||||
///
|
||||
/// Constructible only by [`warp`]. That is the point: an `Embedder` that took
|
||||
/// a plain `&[f32]` would accept an unaligned bounding-box crop and silently
|
||||
/// return worse embeddings, which is a failure no test of the embedder itself
|
||||
/// would catch.
|
||||
pub struct Aligned112 {
|
||||
/// `112 × 112 × 3`, row-major RGB in `0.0..=1.0`.
|
||||
pixels: Vec<f32>,
|
||||
/// Source pixels across the crop before warping — `crop_px` in the catalog.
|
||||
///
|
||||
/// Carried here rather than recomputed later because the scale factor is
|
||||
/// known exactly at warp time and only approximately from the box
|
||||
/// afterwards. §7: it is the honest quality signal, and a feature in the
|
||||
/// calibration.
|
||||
source_px: f32,
|
||||
}
|
||||
|
||||
impl Aligned112 {
|
||||
pub fn pixels(&self) -> &[f32] {
|
||||
&self.pixels
|
||||
}
|
||||
|
||||
/// Source pixels spanned by the 112-pixel crop.
|
||||
///
|
||||
/// Below ~112 the face was upsampled to reach the embedder and the
|
||||
/// embedding is correspondingly weaker; above it, downsampled and healthy.
|
||||
pub fn source_px(&self) -> f32 {
|
||||
self.source_px
|
||||
}
|
||||
}
|
||||
|
||||
/// A similarity transform: rotation, uniform scale, translation.
|
||||
///
|
||||
/// Stored as the four independent parameters rather than a 2×3 matrix so that
|
||||
/// [`Similarity::scale`] is readable without a decomposition.
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub struct Similarity {
|
||||
a: f32,
|
||||
b: f32,
|
||||
tx: f32,
|
||||
ty: f32,
|
||||
}
|
||||
|
||||
impl Similarity {
|
||||
/// `x' = a·x − b·y + tx`, `y' = b·x + a·y + ty`.
|
||||
pub fn apply(&self, x: f32, y: f32) -> (f32, f32) {
|
||||
(
|
||||
self.a * x - self.b * y + self.tx,
|
||||
self.b * x + self.a * y + self.ty,
|
||||
)
|
||||
}
|
||||
|
||||
/// Uniform scale factor — destination pixels per source pixel.
|
||||
pub fn scale(&self) -> f32 {
|
||||
(self.a * self.a + self.b * self.b).sqrt()
|
||||
}
|
||||
|
||||
fn invert(&self, u: f32, v: f32) -> (f32, f32) {
|
||||
let det = self.a * self.a + self.b * self.b;
|
||||
let du = u - self.tx;
|
||||
let dv = v - self.ty;
|
||||
((self.a * du + self.b * dv) / det, (-self.b * du + self.a * dv) / det)
|
||||
}
|
||||
}
|
||||
|
||||
/// Least-squares similarity transform from `src` onto `dst`.
|
||||
///
|
||||
/// # Why least squares and not RANSAC
|
||||
///
|
||||
/// The reference C++ implementation (docs/faces.md §1.1) fits this with
|
||||
/// OpenCV's `estimateAffinePartial2D` under RANSAC. RANSAC over five points is
|
||||
/// a strange fit: the minimal sample for a similarity is two, so it can discard
|
||||
/// landmarks it judges outliers and solve from a subset — and on a profile face
|
||||
/// the "outlier" is as likely to be the correct geometry as the wrong one.
|
||||
/// InsightFace's own pipeline uses plain least squares over all five points,
|
||||
/// which cannot silently drop anything, and that is what this is.
|
||||
///
|
||||
/// # The closed form
|
||||
///
|
||||
/// A 2-D similarity is linear in its four parameters:
|
||||
///
|
||||
/// ```text
|
||||
/// x' = a·x − b·y + tx
|
||||
/// y' = b·x + a·y + ty
|
||||
/// ```
|
||||
///
|
||||
/// so this is an ordinary linear least-squares problem, not an SVD one.
|
||||
/// Centring both point sets kills `tx`/`ty` from the normal equations and
|
||||
/// leaves `a` and `b` as two dot products over a common denominator — which is
|
||||
/// why there is no matrix decomposition anywhere in this function.
|
||||
///
|
||||
/// Returns `None` when the source points are degenerate (coincident or
|
||||
/// collinear to within f32), which does happen: a detector firing on a
|
||||
/// motion-blurred profile can put all five landmarks on a line.
|
||||
pub fn fit_similarity(src: &[(f32, f32); 5], dst: &[(f32, f32); 5]) -> Option<Similarity> {
|
||||
let n = 5.0_f32;
|
||||
let (mut sx, mut sy, mut dx, mut dy) = (0.0, 0.0, 0.0, 0.0);
|
||||
for i in 0..5 {
|
||||
sx += src[i].0;
|
||||
sy += src[i].1;
|
||||
dx += dst[i].0;
|
||||
dy += dst[i].1;
|
||||
}
|
||||
let (sx, sy, dx, dy) = (sx / n, sy / n, dx / n, dy / n);
|
||||
|
||||
let mut var = 0.0_f32;
|
||||
let mut num_a = 0.0_f32;
|
||||
let mut num_b = 0.0_f32;
|
||||
for i in 0..5 {
|
||||
let (px, py) = (src[i].0 - sx, src[i].1 - sy);
|
||||
let (qx, qy) = (dst[i].0 - dx, dst[i].1 - dy);
|
||||
var += px * px + py * py;
|
||||
num_a += px * qx + py * qy;
|
||||
num_b += px * qy - py * qx;
|
||||
}
|
||||
|
||||
// Degenerate: every landmark on one point. Collinear input still solves,
|
||||
// but with a scale that can be absurd, so the caller's sanity check on
|
||||
// `scale()` is what catches that case.
|
||||
if var <= f32::EPSILON {
|
||||
return None;
|
||||
}
|
||||
|
||||
let a = num_a / var;
|
||||
let b = num_b / var;
|
||||
if !a.is_finite() || !b.is_finite() || (a * a + b * b) <= f32::EPSILON {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(Similarity {
|
||||
a,
|
||||
b,
|
||||
tx: dx - (a * sx - b * sy),
|
||||
ty: dy - (b * sx + a * sy),
|
||||
})
|
||||
}
|
||||
|
||||
/// Warp a face onto the canonical 112×112 arrangement.
|
||||
///
|
||||
/// `rgb` is tightly packed `f32` RGB in `0.0..=1.0`, row-major — the same
|
||||
/// convention `dr-segment` uses, so both read the same proxy.
|
||||
///
|
||||
/// Sampling is bilinear **from the source in one step**: never crop-then-warp,
|
||||
/// which resamples twice and throws away detail the warp could have used.
|
||||
/// Pixels falling outside the source read as black.
|
||||
pub fn warp(
|
||||
rgb: &[f32],
|
||||
width: usize,
|
||||
height: usize,
|
||||
landmarks: &[(f32, f32); 5],
|
||||
) -> Option<Aligned112> {
|
||||
if rgb.len() != width * height * 3 {
|
||||
return None;
|
||||
}
|
||||
let m = fit_similarity(landmarks, &ARCFACE_TEMPLATE)?;
|
||||
|
||||
let e = ALIGNED_EDGE;
|
||||
let mut pixels = vec![0.0_f32; e * e * 3];
|
||||
for v in 0..e {
|
||||
for u in 0..e {
|
||||
// Pixel centres, so the transform is not off by half a pixel —
|
||||
// which is small enough to survive review and large enough to
|
||||
// matter on a 40-pixel face.
|
||||
let (x, y) = m.invert(u as f32 + 0.5, v as f32 + 0.5);
|
||||
let (x, y) = (x - 0.5, y - 0.5);
|
||||
let out = (v * e + u) * 3;
|
||||
sample_bilinear(rgb, width, height, x, y, &mut pixels[out..out + 3]);
|
||||
}
|
||||
}
|
||||
|
||||
Some(Aligned112 {
|
||||
pixels,
|
||||
// The warp maps `scale` source pixels to one destination pixel, so the
|
||||
// crop spans 112/scale of the source.
|
||||
source_px: ALIGNED_EDGE as f32 / m.scale(),
|
||||
})
|
||||
}
|
||||
|
||||
fn sample_bilinear(rgb: &[f32], w: usize, h: usize, x: f32, y: f32, out: &mut [f32]) {
|
||||
let x0 = x.floor();
|
||||
let y0 = y.floor();
|
||||
let fx = x - x0;
|
||||
let fy = y - y0;
|
||||
let x0 = x0 as isize;
|
||||
let y0 = y0 as isize;
|
||||
|
||||
for (c, o) in out.iter_mut().enumerate() {
|
||||
let get = |xi: isize, yi: isize| -> f32 {
|
||||
if xi < 0 || yi < 0 || xi >= w as isize || yi >= h as isize {
|
||||
0.0
|
||||
} else {
|
||||
rgb[(yi as usize * w + xi as usize) * 3 + c]
|
||||
}
|
||||
};
|
||||
let top = get(x0, y0) * (1.0 - fx) + get(x0 + 1, y0) * fx;
|
||||
let bot = get(x0, y0 + 1) * (1.0 - fx) + get(x0 + 1, y0 + 1) * fx;
|
||||
*o = top * (1.0 - fy) + bot * fy;
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn shifted_scaled(scale: f32, dx: f32, dy: f32, rot: f32) -> [(f32, f32); 5] {
|
||||
let (s, c) = (rot.sin(), rot.cos());
|
||||
let mut out = [(0.0, 0.0); 5];
|
||||
for (i, &(x, y)) in ARCFACE_TEMPLATE.iter().enumerate() {
|
||||
out[i] = (
|
||||
scale * (c * x - s * y) + dx,
|
||||
scale * (s * x + c * y) + dy,
|
||||
);
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn template_onto_itself_is_the_identity() {
|
||||
let m = fit_similarity(&ARCFACE_TEMPLATE, &ARCFACE_TEMPLATE).unwrap();
|
||||
for &(x, y) in &ARCFACE_TEMPLATE {
|
||||
let (u, v) = m.apply(x, y);
|
||||
assert!((u - x).abs() < 1e-3, "{u} vs {x}");
|
||||
assert!((v - y).abs() < 1e-3, "{v} vs {y}");
|
||||
}
|
||||
assert!((m.scale() - 1.0).abs() < 1e-4);
|
||||
}
|
||||
|
||||
/// The property that matters: whatever similarity the face was seen under,
|
||||
/// the fit must undo it and land the landmarks back on the template. This
|
||||
/// is the test that fails if the transform is ever "simplified" into an
|
||||
/// affine or a bare scale-and-translate.
|
||||
#[test]
|
||||
fn any_similarity_of_the_template_maps_back_onto_it() {
|
||||
for &(scale, dx, dy, rot) in &[
|
||||
(1.0_f32, 0.0_f32, 0.0_f32, 0.0_f32),
|
||||
(2.5, 100.0, -40.0, 0.0),
|
||||
(0.4, -12.0, 300.0, 0.6),
|
||||
(1.7, 5.0, 5.0, -1.2),
|
||||
] {
|
||||
let observed = shifted_scaled(scale, dx, dy, rot);
|
||||
let m = fit_similarity(&observed, &ARCFACE_TEMPLATE).unwrap();
|
||||
for (i, &(tx, ty)) in ARCFACE_TEMPLATE.iter().enumerate() {
|
||||
let (u, v) = m.apply(observed[i].0, observed[i].1);
|
||||
assert!(
|
||||
(u - tx).abs() < 1e-2 && (v - ty).abs() < 1e-2,
|
||||
"scale={scale} rot={rot}: point {i} landed at ({u}, {v}), want ({tx}, {ty})"
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
(m.scale() - 1.0 / scale).abs() < 1e-3,
|
||||
"scale {} should invert {scale}",
|
||||
m.scale()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn coincident_landmarks_are_rejected_rather_than_producing_a_crop() {
|
||||
let degenerate = [(50.0, 50.0); 5];
|
||||
assert!(fit_similarity(°enerate, &ARCFACE_TEMPLATE).is_none());
|
||||
let rgb = vec![0.5_f32; 64 * 64 * 3];
|
||||
assert!(warp(&rgb, 64, 64, °enerate).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn source_px_reports_the_face_size_the_embedder_actually_saw() {
|
||||
let rgb = vec![0.5_f32; 400 * 400 * 3];
|
||||
// A face twice the template's size spans 224 source pixels.
|
||||
let big = shifted_scaled(2.0, 80.0, 80.0, 0.0);
|
||||
let a = warp(&rgb, 400, 400, &big).unwrap();
|
||||
assert!((a.source_px() - 224.0).abs() < 0.5, "{}", a.source_px());
|
||||
|
||||
// Half-size: 56 source pixels upsampled to 112, which §7 calls the
|
||||
// degraded bucket.
|
||||
let small = shifted_scaled(0.5, 10.0, 10.0, 0.0);
|
||||
let a = warp(&rgb, 400, 400, &small).unwrap();
|
||||
assert!((a.source_px() - 56.0).abs() < 0.5, "{}", a.source_px());
|
||||
}
|
||||
|
||||
/// A white square on black, warped by a transform that should centre it:
|
||||
/// checks the sampler's geometry rather than the fit's algebra.
|
||||
#[test]
|
||||
fn warp_resamples_the_right_pixels() {
|
||||
let (w, h) = (224, 224);
|
||||
let mut rgb = vec![0.0_f32; w * h * 3];
|
||||
for y in 0..h {
|
||||
for x in 0..w {
|
||||
if x >= 56 && x < 168 && y >= 56 && y < 168 {
|
||||
for c in 0..3 {
|
||||
rgb[(y * w + x) * 3 + c] = 1.0;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// Landmarks placed so the fit is a pure translation of (56, 56):
|
||||
// the white square maps exactly onto the 112×112 output.
|
||||
let lm = shifted_scaled(1.0, 56.0, 56.0, 0.0);
|
||||
let a = warp(&rgb, w, h, &lm).unwrap();
|
||||
let px = a.pixels();
|
||||
for (i, v) in px.iter().enumerate() {
|
||||
assert!((v - 1.0).abs() < 1e-3, "pixel {i} is {v}, expected white");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn out_of_bounds_samples_read_black_rather_than_wrapping() {
|
||||
let rgb = vec![1.0_f32; 32 * 32 * 3];
|
||||
// Face far outside the image: every sample is out of bounds.
|
||||
let lm = shifted_scaled(1.0, 5000.0, 5000.0, 0.0);
|
||||
let a = warp(&rgb, 32, 32, &lm).unwrap();
|
||||
assert!(a.pixels().iter().all(|&v| v == 0.0));
|
||||
}
|
||||
}
|
||||
+333
-31
@@ -1,28 +1,96 @@
|
||||
//! SCRFD face detection (docs/faces.md §4).
|
||||
//!
|
||||
//! For now: enough of the loader to answer M1 — whether tract will parse these
|
||||
//! graphs at all — plus the load-time shape validation that keeps a YuNet file
|
||||
//! from being decoded as an SCRFD one.
|
||||
//! One forward pass produces a box, a confidence and **five landmarks** per
|
||||
//! face — the landmarks being the reason for this detector rather than a
|
||||
//! general one, since [`crate::align`] cannot work without them.
|
||||
//!
|
||||
//! # The graph must have fixed input dimensions
|
||||
//!
|
||||
//! InsightFace ships `det_500m.onnx` with a dynamic H/W input, and **tract
|
||||
//! cannot parse it in that form** — it fails at node #0. The same file run
|
||||
//! through `tools/fix-face-model-shapes.sh` loads cleanly. Its outputs were
|
||||
//! already static at 640, so 640 is not a choice made here: it is the shape
|
||||
//! the export was always going to run at.
|
||||
|
||||
use ndarray::Array4;
|
||||
|
||||
use crate::{install_backend, FaceError};
|
||||
|
||||
/// The graph's input edge, in pixels. See the module note: not configurable.
|
||||
pub const INPUT_EDGE: usize = 640;
|
||||
|
||||
/// Strides, in the order SCRFD emits them.
|
||||
const ALL_STRIDES: [usize; 4] = [8, 16, 32, 64];
|
||||
|
||||
/// Anchors per feature-map location.
|
||||
const ANCHORS: usize = 2;
|
||||
|
||||
/// How detection is tuned.
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub struct DetectOptions {
|
||||
/// Minimum detector confidence.
|
||||
///
|
||||
/// Deliberately *not* the low threshold `dr-segment` chose. There a false
|
||||
/// positive costs one spurious row in a list the user is picking from;
|
||||
/// here it costs a face in the People view to reject and — worse — a
|
||||
/// garbage embedding that can bridge two real clusters into one. A false
|
||||
/// negative is recoverable by re-indexing with a better model; a polluted
|
||||
/// cluster graph, once the user has confirmed faces inside it, is not.
|
||||
pub confidence: f32,
|
||||
/// Box IoU above which two detections are judged to be the same face.
|
||||
pub nms_iou: f32,
|
||||
/// Smallest face to keep, in source pixels on the shorter box edge.
|
||||
///
|
||||
/// Below this there is not enough face left to align reliably — see
|
||||
/// docs/faces.md §7 for what the embedder actually receives at each size.
|
||||
pub min_face_px: f32,
|
||||
}
|
||||
|
||||
impl Default for DetectOptions {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
confidence: 0.5,
|
||||
nms_iou: 0.4,
|
||||
min_face_px: 40.0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// One detected face, in **source image pixels**.
|
||||
///
|
||||
/// Pixels rather than the normalised form the catalog stores, because the
|
||||
/// caller still has to crop from this image. Normalisation happens at the
|
||||
/// storage boundary, where the long edge is known to be the right divisor.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct Detection {
|
||||
/// `(x0, y0, x1, y1)`.
|
||||
pub bbox: (f32, f32, f32, f32),
|
||||
/// Five points in the detector's own order — see [`crate::align`], which
|
||||
/// consumes them without reordering.
|
||||
pub landmarks: [(f32, f32); 5],
|
||||
pub confidence: f32,
|
||||
}
|
||||
|
||||
impl Detection {
|
||||
pub fn width(&self) -> f32 {
|
||||
self.bbox.2 - self.bbox.0
|
||||
}
|
||||
pub fn height(&self) -> f32 {
|
||||
self.bbox.3 - self.bbox.1
|
||||
}
|
||||
}
|
||||
|
||||
/// A loaded SCRFD graph.
|
||||
pub struct Detector {
|
||||
session: ort::session::Session,
|
||||
/// 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
|
||||
/// exports exist and hardcoding 3 silently ignores the largest faces in a
|
||||
/// four-stride model.
|
||||
strides: usize,
|
||||
/// exports exist and hardcoding 3 silently ignores the largest faces a
|
||||
/// four-stride model finds.
|
||||
fmc: usize,
|
||||
}
|
||||
|
||||
/// Strides, in the order SCRFD emits them.
|
||||
pub const ALL_STRIDES: [usize; 4] = [8, 16, 32, 64];
|
||||
|
||||
/// Anchors per feature-map location.
|
||||
pub const ANCHORS: usize = 2;
|
||||
|
||||
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)?;
|
||||
@@ -44,14 +112,16 @@ impl Detector {
|
||||
detail: format!("expected 9 or 12 outputs, got {n_out}"),
|
||||
});
|
||||
}
|
||||
let strides = n_out / 3;
|
||||
let fmc = n_out / 3;
|
||||
|
||||
// The check that actually distinguishes the models: score, box and
|
||||
// landmark groups end in 1, 4 and 10 respectively. YuNet also has
|
||||
// twelve outputs, so the count alone proves nothing.
|
||||
// The check that actually distinguishes the models. YuNet also has
|
||||
// twelve outputs in three strides, so the count proves nothing — its
|
||||
// groups are cls/obj/bbox/kps where SCRFD's are score/bbox/kps, and
|
||||
// decoding one as the other yields a page of plausible numbers rather
|
||||
// than an error. The last dimension is what separates them.
|
||||
for (group, expected_last) in [1_i64, 4, 10].into_iter().enumerate() {
|
||||
for s in 0..strides {
|
||||
let idx = group * strides + s;
|
||||
for s in 0..fmc {
|
||||
let idx = group * fmc + s;
|
||||
let out = &session.outputs()[idx];
|
||||
let last: Option<i64> =
|
||||
out.dtype().tensor_shape().and_then(|d| d.last().copied());
|
||||
@@ -59,29 +129,261 @@ impl Detector {
|
||||
return Err(FaceError::WrongModel {
|
||||
expected: "InsightFace SCRFD",
|
||||
detail: format!(
|
||||
"output '{}' last dim is {:?}, expected {expected_last}",
|
||||
out.name(), last
|
||||
"output '{}' last dim is {:?}, expected {expected_last} \
|
||||
(a YuNet export fails exactly here)",
|
||||
out.name(),
|
||||
last
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Self { session, strides })
|
||||
Ok(Self { session, fmc })
|
||||
}
|
||||
|
||||
/// Number of stride levels this graph emits.
|
||||
/// Stride levels this graph emits.
|
||||
pub fn strides(&self) -> &'static [usize] {
|
||||
&ALL_STRIDES[..self.strides]
|
||||
&ALL_STRIDES[..self.fmc]
|
||||
}
|
||||
|
||||
/// The graph's declared input shape, for diagnosing a dynamic export.
|
||||
pub fn input_shape(&self) -> Option<Vec<i64>> {
|
||||
self.session
|
||||
.inputs()
|
||||
.first()?
|
||||
.dtype()
|
||||
.tensor_shape()
|
||||
.map(|s| s.to_vec())
|
||||
/// Find the faces in an image.
|
||||
///
|
||||
/// `rgb` is tightly packed `f32` RGB in `0.0..=1.0`, row-major — the same
|
||||
/// convention `dr-segment` and [`crate::align`] use.
|
||||
pub fn detect(
|
||||
&mut self,
|
||||
rgb: &[f32],
|
||||
width: usize,
|
||||
height: usize,
|
||||
options: &DetectOptions,
|
||||
) -> Result<Vec<Detection>, FaceError> {
|
||||
if width == 0 || height == 0 {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
if rgb.len() != width * height * 3 {
|
||||
return Err(FaceError::ImageShape {
|
||||
expected: width * height * 3,
|
||||
got: rgb.len(),
|
||||
});
|
||||
}
|
||||
|
||||
let lb = Letterbox::fit(width as f32, height as f32);
|
||||
let input = lb.sample(rgb, width, height);
|
||||
|
||||
let outputs = self
|
||||
.session
|
||||
.run(ort::inputs![
|
||||
ort::value::Tensor::from_array(input).map_err(FaceError::Inference)?
|
||||
])
|
||||
.map_err(FaceError::Inference)?;
|
||||
|
||||
let mut raw: Vec<Detection> = Vec::new();
|
||||
|
||||
for (si, &stride) in ALL_STRIDES[..self.fmc].iter().enumerate() {
|
||||
let (_, scores) = outputs[si]
|
||||
.try_extract_tensor::<f32>()
|
||||
.map_err(FaceError::Inference)?;
|
||||
let (_, boxes) = outputs[self.fmc + si]
|
||||
.try_extract_tensor::<f32>()
|
||||
.map_err(FaceError::Inference)?;
|
||||
let (_, kps) = outputs[self.fmc * 2 + si]
|
||||
.try_extract_tensor::<f32>()
|
||||
.map_err(FaceError::Inference)?;
|
||||
|
||||
let fw = INPUT_EDGE / stride;
|
||||
let fh = INPUT_EDGE / stride;
|
||||
let s = stride as f32;
|
||||
|
||||
for r in 0..fh {
|
||||
for c in 0..fw {
|
||||
for a in 0..ANCHORS {
|
||||
let idx = (r * fw + c) * ANCHORS + a;
|
||||
let score = scores[idx];
|
||||
if score < options.confidence {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Anchor centre in input space, then distance-to-box
|
||||
// decoding: the four regressed values are distances
|
||||
// left/top/right/bottom in units of the stride.
|
||||
let (cx, cy) = ((c * stride) as f32, (r * stride) as f32);
|
||||
let b = &boxes[idx * 4..idx * 4 + 4];
|
||||
let (x0, y0) = lb.to_source(cx - b[0] * s, cy - b[1] * s);
|
||||
let (x1, y1) = lb.to_source(cx + b[2] * s, cy + b[3] * s);
|
||||
|
||||
let k = &kps[idx * 10..idx * 10 + 10];
|
||||
let mut landmarks = [(0.0_f32, 0.0_f32); 5];
|
||||
for (p, lm) in landmarks.iter_mut().enumerate() {
|
||||
*lm = lb.to_source(cx + k[p * 2] * s, cy + k[p * 2 + 1] * s);
|
||||
}
|
||||
|
||||
raw.push(Detection {
|
||||
bbox: (x0, y0, x1, y1),
|
||||
landmarks,
|
||||
confidence: score,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut kept = non_max_suppress(raw, options.nms_iou);
|
||||
|
||||
// Size floor last, on the *merged* boxes: a face that only clears the
|
||||
// floor once NMS has picked the best of its overlapping detections
|
||||
// should be kept.
|
||||
kept.retain(|d| d.width().min(d.height()) >= options.min_face_px);
|
||||
|
||||
// No cap on the count. The reference implementation keeps the ten
|
||||
// largest, which is right for a film frame where background extras are
|
||||
// noise; it is wrong for a photo library, where a group shot with
|
||||
// thirty faces is precisely the picture worth indexing.
|
||||
Ok(kept)
|
||||
}
|
||||
}
|
||||
|
||||
/// Greedy NMS across all strides together.
|
||||
fn non_max_suppress(mut dets: Vec<Detection>, iou_threshold: f32) -> Vec<Detection> {
|
||||
dets.sort_by(|a, b| b.confidence.total_cmp(&a.confidence));
|
||||
let mut kept: Vec<Detection> = Vec::new();
|
||||
for d in dets {
|
||||
if kept.iter().all(|k| iou(&k.bbox, &d.bbox) <= iou_threshold) {
|
||||
kept.push(d);
|
||||
}
|
||||
}
|
||||
kept
|
||||
}
|
||||
|
||||
fn iou(a: &(f32, f32, f32, f32), b: &(f32, f32, f32, f32)) -> f32 {
|
||||
let ix = (a.2.min(b.2) - a.0.max(b.0)).max(0.0);
|
||||
let iy = (a.3.min(b.3) - a.1.max(b.1)).max(0.0);
|
||||
let inter = ix * iy;
|
||||
let area_a = (a.2 - a.0).max(0.0) * (a.3 - a.1).max(0.0);
|
||||
let area_b = (b.2 - b.0).max(0.0) * (b.3 - b.1).max(0.0);
|
||||
let union = area_a + area_b - inter;
|
||||
if union <= 0.0 {
|
||||
0.0
|
||||
} else {
|
||||
inter / union
|
||||
}
|
||||
}
|
||||
|
||||
/// How the image is fitted into the graph's fixed square input.
|
||||
///
|
||||
/// The forward and inverse mappings live in one struct on purpose:
|
||||
/// docs/faces.md §4.1 notes that what matters is not *where* the padding goes
|
||||
/// but that the two agree. A mismatch offsets every box and landmark by the
|
||||
/// padding, producing detections that look plausible and embeddings that
|
||||
/// quietly cluster badly three stages later.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct Letterbox {
|
||||
/// Input pixels per source pixel.
|
||||
scale: f32,
|
||||
pad_x: f32,
|
||||
pad_y: f32,
|
||||
}
|
||||
|
||||
impl Letterbox {
|
||||
fn fit(w: f32, h: f32) -> Self {
|
||||
let scale = (INPUT_EDGE as f32 / w).min(INPUT_EDGE as f32 / h);
|
||||
Self {
|
||||
scale,
|
||||
pad_x: (INPUT_EDGE as f32 - w * scale) * 0.5,
|
||||
pad_y: (INPUT_EDGE as f32 - h * scale) * 0.5,
|
||||
}
|
||||
}
|
||||
|
||||
/// Resample into `[1, 3, 640, 640]`, normalised as the weights expect.
|
||||
///
|
||||
/// `(x·255 − 127.5) / 128` — note `/128`, not `/127.5`. The reference
|
||||
/// implementation this is ported from uses `/128` for both models, and
|
||||
/// every measured number in docs/faces.md §1 came from it.
|
||||
///
|
||||
/// Padding is grey, matching the reference's `114`: the value the network
|
||||
/// reads least as an edge, where black would draw a hard border across the
|
||||
/// frame and invite a detection along it.
|
||||
fn sample(&self, rgb: &[f32], width: usize, height: usize) -> Array4<f32> {
|
||||
const PAD: f32 = 114.0;
|
||||
let norm = |v: f32| (v * 255.0 - 127.5) / 128.0;
|
||||
|
||||
let mut input =
|
||||
Array4::<f32>::from_elem((1, 3, INPUT_EDGE, INPUT_EDGE), (PAD - 127.5) / 128.0);
|
||||
|
||||
for iy in 0..INPUT_EDGE {
|
||||
let sy = (iy as f32 + 0.5 - self.pad_y) / self.scale - 0.5;
|
||||
if sy < -0.5 || sy > height as f32 - 0.5 {
|
||||
continue;
|
||||
}
|
||||
for ix in 0..INPUT_EDGE {
|
||||
let sx = (ix as f32 + 0.5 - self.pad_x) / self.scale - 0.5;
|
||||
if sx < -0.5 || sx > width as f32 - 0.5 {
|
||||
continue;
|
||||
}
|
||||
let (x0f, y0f) = (sx.floor(), sy.floor());
|
||||
let (fx, fy) = (sx - x0f, sy - y0f);
|
||||
let x0 = (x0f as isize).clamp(0, width as isize - 1) as usize;
|
||||
let y0 = (y0f as isize).clamp(0, height as isize - 1) as usize;
|
||||
let x1 = (x0 + 1).min(width - 1);
|
||||
let y1 = (y0 + 1).min(height - 1);
|
||||
|
||||
for c in 0..3 {
|
||||
let at = |x: usize, y: usize| rgb[(y * width + x) * 3 + c];
|
||||
let top = at(x0, y0) * (1.0 - fx) + at(x1, y0) * fx;
|
||||
let bot = at(x0, y1) * (1.0 - fx) + at(x1, y1) * fx;
|
||||
input[[0, c, iy, ix]] = norm(top * (1.0 - fy) + bot * fy);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
input
|
||||
}
|
||||
|
||||
/// Input-space point back to source pixels.
|
||||
fn to_source(&self, x: f32, y: f32) -> (f32, f32) {
|
||||
((x - self.pad_x) / self.scale, (y - self.pad_y) / self.scale)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn letterbox_round_trips_a_point() {
|
||||
let lb = Letterbox::fit(1024.0, 683.0);
|
||||
for &(x, y) in &[(0.0_f32, 0.0_f32), (512.0, 341.0), (1023.0, 682.0)] {
|
||||
let (bx, by) = lb.to_source(x * lb.scale + lb.pad_x, y * lb.scale + lb.pad_y);
|
||||
assert!((bx - x).abs() < 1e-2, "{bx} vs {x}");
|
||||
assert!((by - y).abs() < 1e-2, "{by} vs {y}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn letterbox_centres_the_short_axis() {
|
||||
let lb = Letterbox::fit(640.0, 320.0);
|
||||
assert!((lb.scale - 1.0).abs() < 1e-6);
|
||||
assert!(lb.pad_x.abs() < 1e-6);
|
||||
assert!((lb.pad_y - 160.0).abs() < 1e-6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn nms_keeps_the_confident_box_and_drops_its_duplicate() {
|
||||
let d = |x: f32, conf: f32| Detection {
|
||||
bbox: (x, 0.0, x + 100.0, 100.0),
|
||||
landmarks: [(0.0, 0.0); 5],
|
||||
confidence: conf,
|
||||
};
|
||||
let kept = non_max_suppress(vec![d(0.0, 0.8), d(5.0, 0.9), d(500.0, 0.7)], 0.4);
|
||||
assert_eq!(kept.len(), 2);
|
||||
assert!((kept[0].confidence - 0.9).abs() < 1e-6);
|
||||
assert!((kept[1].bbox.0 - 500.0).abs() < 1e-6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn iou_of_a_box_with_itself_is_one_and_with_a_disjoint_box_is_zero() {
|
||||
let a = (0.0, 0.0, 10.0, 10.0);
|
||||
assert!((iou(&a, &a) - 1.0).abs() < 1e-6);
|
||||
assert!(iou(&a, &(100.0, 100.0, 110.0, 110.0)) < 1e-6);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
//! What an embedder produces, and how it is stored (docs/faces.md §6).
|
||||
//!
|
||||
//! Deliberately **model-free**: the vector, its identity, its comparison and
|
||||
//! its storage encoding are arithmetic, and `calibrate` and `cluster` are built
|
||||
//! on them. Keeping them out of the `inference` feature is what lets the part
|
||||
//! of this subsystem most likely to be subtly wrong be tested on a machine with
|
||||
//! no weights on it.
|
||||
//!
|
||||
//! [`crate::embed::Embedder`] is the thing that needs a model, and it lives
|
||||
//! behind the feature.
|
||||
|
||||
/// Embedding dimensionality. Fixed by the model family, not a parameter.
|
||||
pub const EMBEDDING_DIM: usize = 512;
|
||||
|
||||
/// Which model produced an embedding.
|
||||
///
|
||||
/// Embeddings from different models are not comparable, and this is the one
|
||||
/// mistake that produces plausible-looking garbage rather than an error — so
|
||||
/// the id travels *with* the vector rather than beside it, and
|
||||
/// [`Embedding::cosine`] refuses a cross-model comparison.
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||
pub struct ModelId(pub std::sync::Arc<str>);
|
||||
|
||||
impl ModelId {
|
||||
pub fn new(s: impl Into<std::sync::Arc<str>>) -> Self {
|
||||
Self(s.into())
|
||||
}
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ModelId {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.write_str(&self.0)
|
||||
}
|
||||
}
|
||||
|
||||
/// A 512-d L2-normalised face embedding.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct Embedding {
|
||||
pub model: ModelId,
|
||||
pub v: Box<[f32; EMBEDDING_DIM]>,
|
||||
}
|
||||
|
||||
impl Embedding {
|
||||
/// Cosine similarity, which for unit vectors is the plain dot product.
|
||||
///
|
||||
/// `None` when the two came from different models. That is a real
|
||||
/// possibility in a library indexed across a model upgrade, and the
|
||||
/// alternative — returning a number — is the failure mode
|
||||
/// `faces.model_id` exists to prevent.
|
||||
pub fn cosine(&self, other: &Embedding) -> Option<f32> {
|
||||
if self.model != other.model {
|
||||
return None;
|
||||
}
|
||||
Some(dot(&self.v, &other.v))
|
||||
}
|
||||
|
||||
/// Storage form: `512 × f16`, 1 KB per face (catalog.md §10.1).
|
||||
pub fn to_f16_bytes(&self) -> Vec<u8> {
|
||||
let mut out = Vec::with_capacity(EMBEDDING_DIM * 2);
|
||||
for &x in self.v.iter() {
|
||||
out.extend_from_slice(&f32_to_f16_bits(x).to_le_bytes());
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
/// Read back from storage, re-normalising.
|
||||
///
|
||||
/// The f16 round-trip perturbs a unit vector by ~1e-3 in cosine — three
|
||||
/// orders below the separation between a match and a non-match — but the
|
||||
/// drift is free to remove and invisible if left, so it is removed here
|
||||
/// rather than remembered at every call site.
|
||||
pub fn from_f16_bytes(model: ModelId, bytes: &[u8]) -> Option<Self> {
|
||||
if bytes.len() != EMBEDDING_DIM * 2 {
|
||||
return None;
|
||||
}
|
||||
let mut v = Box::new([0.0_f32; EMBEDDING_DIM]);
|
||||
for (i, chunk) in bytes.chunks_exact(2).enumerate() {
|
||||
v[i] = f16_bits_to_f32(u16::from_le_bytes([chunk[0], chunk[1]]));
|
||||
}
|
||||
normalise(&mut v);
|
||||
Some(Self { model, v })
|
||||
}
|
||||
}
|
||||
|
||||
fn dot(a: &[f32; EMBEDDING_DIM], b: &[f32; EMBEDDING_DIM]) -> f32 {
|
||||
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
|
||||
}
|
||||
|
||||
pub(crate) fn normalise(v: &mut [f32; EMBEDDING_DIM]) {
|
||||
// Clamped rather than checked: a zero-norm embedding is a broken model,
|
||||
// not a runtime condition worth an error path, and dividing by 1e-6 keeps
|
||||
// the NaN out of the catalog.
|
||||
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt().max(1e-6);
|
||||
for x in v.iter_mut() {
|
||||
*x /= norm;
|
||||
}
|
||||
}
|
||||
|
||||
// ── f16 ───────────────────────────────────────────────────────────────────
|
||||
//
|
||||
// Hand-rolled rather than pulling in `half`: two functions over a format that
|
||||
// has not changed since 2008, used at exactly one boundary. The dependency
|
||||
// policy (D13, D1) makes the bar for a new crate high, and this is well under
|
||||
// it.
|
||||
|
||||
fn f32_to_f16_bits(x: f32) -> u16 {
|
||||
let bits = x.to_bits();
|
||||
let sign = ((bits >> 16) & 0x8000) as u16;
|
||||
let exp = ((bits >> 23) & 0xff) as i32 - 127 + 15;
|
||||
let mant = bits & 0x007f_ffff;
|
||||
|
||||
if exp >= 0x1f {
|
||||
// Overflow, inf, or NaN. Embeddings are unit-norm so this is the
|
||||
// broken-model path; infinity is the honest answer, not a clamp that
|
||||
// hides it.
|
||||
return sign | 0x7c00 | if mant != 0 && exp == 0x1f + 112 { 0x200 } else { 0 };
|
||||
}
|
||||
if exp <= 0 {
|
||||
// Subnormal or underflow. A component of a unit 512-vector is ~0.04,
|
||||
// nowhere near here, so this branch exists for correctness rather than
|
||||
// for traffic.
|
||||
if exp < -10 {
|
||||
return sign;
|
||||
}
|
||||
let mant = mant | 0x0080_0000;
|
||||
let shift = (14 - exp) as u32;
|
||||
let half = (mant >> shift) as u16;
|
||||
// Round to nearest, ties to even.
|
||||
let rem = mant & ((1 << shift) - 1);
|
||||
let tie = 1 << (shift - 1);
|
||||
let round = u16::from(rem > tie || (rem == tie && (half & 1) == 1));
|
||||
return sign | (half + round);
|
||||
}
|
||||
|
||||
let half = ((exp as u16) << 10) | (mant >> 13) as u16;
|
||||
let rem = mant & 0x1fff;
|
||||
let round = u16::from(rem > 0x1000 || (rem == 0x1000 && (half & 1) == 1));
|
||||
sign | (half + round)
|
||||
}
|
||||
|
||||
fn f16_bits_to_f32(h: u16) -> f32 {
|
||||
let sign = ((h & 0x8000) as u32) << 16;
|
||||
let exp = ((h >> 10) & 0x1f) as u32;
|
||||
let mant = (h & 0x03ff) as u32;
|
||||
|
||||
if exp == 0 {
|
||||
if mant == 0 {
|
||||
return f32::from_bits(sign);
|
||||
}
|
||||
// Subnormal: renormalise into f32's range.
|
||||
let mut e = -1_i32;
|
||||
let mut m = mant;
|
||||
while m & 0x0400 == 0 {
|
||||
m <<= 1;
|
||||
e -= 1;
|
||||
}
|
||||
let m = m & 0x03ff;
|
||||
return f32::from_bits(sign | (((127 - 15 + 1 + e) as u32) << 23) | (m << 13));
|
||||
}
|
||||
if exp == 0x1f {
|
||||
return f32::from_bits(sign | 0x7f80_0000 | (mant << 13));
|
||||
}
|
||||
f32::from_bits(sign | ((exp + 127 - 15) << 23) | (mant << 13))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn unit(seed: u32) -> Embedding {
|
||||
let mut v = Box::new([0.0_f32; EMBEDDING_DIM]);
|
||||
let mut s = seed.wrapping_mul(2_654_435_761).wrapping_add(1);
|
||||
for x in v.iter_mut() {
|
||||
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
|
||||
*x = (s >> 8) as f32 / (1u32 << 23) as f32 - 0.5;
|
||||
}
|
||||
normalise(&mut v);
|
||||
Embedding {
|
||||
model: ModelId::new("test"),
|
||||
v,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_normalised_embedding_has_cosine_one_with_itself() {
|
||||
let e = unit(7);
|
||||
assert!((e.cosine(&e).unwrap() - 1.0).abs() < 1e-5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn embeddings_from_different_models_do_not_compare() {
|
||||
let a = unit(1);
|
||||
let mut b = unit(1);
|
||||
b.model = ModelId::new("other");
|
||||
assert_eq!(a.cosine(&b), None, "a cross-model cosine must not be a number");
|
||||
}
|
||||
|
||||
/// The claim docs/faces.md §6 makes about the storage format: the f16
|
||||
/// round-trip costs ~1e-3 of cosine, three orders below the separation
|
||||
/// between a match and a non-match.
|
||||
#[test]
|
||||
fn f16_round_trip_preserves_the_embedding() {
|
||||
for seed in 0..16 {
|
||||
let e = unit(seed);
|
||||
let back = Embedding::from_f16_bytes(e.model.clone(), &e.to_f16_bytes()).unwrap();
|
||||
let cos = e.cosine(&back).unwrap();
|
||||
assert!(cos > 0.9999, "seed {seed}: round-trip cosine {cos}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn f16_round_trip_rejects_a_wrong_length_blob() {
|
||||
assert!(Embedding::from_f16_bytes(ModelId::new("m"), &[0u8; 100]).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn f16_handles_the_values_an_embedding_actually_contains() {
|
||||
// Components of a unit 512-vector cluster around ±1/sqrt(512) ≈ 0.044.
|
||||
for &x in &[0.0_f32, 1.0, -1.0, 0.044_194_17, -0.044_194_17, 1e-3, -7e-4] {
|
||||
let back = f16_bits_to_f32(f32_to_f16_bits(x));
|
||||
assert!(
|
||||
(back - x).abs() <= 1e-3 * x.abs().max(1e-3),
|
||||
"{x} round-tripped to {back}"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -30,8 +30,19 @@
|
||||
//! machine with no weights on it — which is what lets CI cover the part most
|
||||
//! likely to be subtly wrong.
|
||||
|
||||
pub mod align;
|
||||
pub mod embedding;
|
||||
#[cfg(feature = "inference")]
|
||||
pub mod detect;
|
||||
#[cfg(feature = "inference")]
|
||||
pub mod embed;
|
||||
|
||||
pub use align::{warp, Aligned112, Similarity, ALIGNED_EDGE, ARCFACE_TEMPLATE};
|
||||
pub use embedding::{Embedding, ModelId, EMBEDDING_DIM};
|
||||
#[cfg(feature = "inference")]
|
||||
pub use detect::{DetectOptions, Detection, Detector};
|
||||
#[cfg(feature = "inference")]
|
||||
pub use embed::Embedder;
|
||||
|
||||
/// What can go wrong between an image and a face.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
|
||||
Reference in New Issue
Block a user