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:
2026-08-26 19:57:56 +02:00
co-authored by Claude Opus 5
parent 72410f39c6
commit 19981c1033
9 changed files with 1196 additions and 32 deletions
+124
View File
@@ -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))
}