12ede47b2c
Implements the last Phase 2 piece: a working face detector/embedder behind
the new `faces-onnx` cargo feature (mirrors how `plugins` gates wasmtime).
Inert by default — the default build is unchanged and ships the no-op
analyzer.
Pipeline (InsightFace/immich pattern): SCRFD detection with 5-point
landmarks → least-squares similarity alignment to the canonical 112×112
template → ArcFace embedding → L2-normalized 512-d vector.
- face_geometry.rs (always compiled, unit-tested): SCRFD anchor/distance
decode, NMS, the closed-form (complex-number) similarity transform,
bilinear affine warp, NCHW normalization, L2-norm, Laplacian sharpness.
11 unit tests cover the error-prone math with no model needed.
- onnx_face_analyzer.rs (feature `faces-onnx`): wires the geometry to ONNX
Runtime via `ort` (load-dynamic, so libonnxruntime is dlopen'd at runtime
and the crate builds without it). Inference runs on spawn_blocking; each
session is serialized behind a Mutex. Loads via `ort::init_from` (fallible)
not ORT's lazy loader, which would panic under `panic = "abort"`.
- config: FacesConfig + OXICLOUD_FACES_{ORT_DYLIB,DETECTOR_MODEL,
EMBEDDER_MODEL,DET_SIZE,DET_THRESHOLD,NMS_THRESHOLD,INTRA_THREADS}.
- di: build_face_analyzer() loads the real analyzer when the feature is
compiled in and runtime+models are configured; any missing piece or load
failure degrades to the no-op analyzer (logged) so startup never fails.
- ort/ndarray added as optional deps; example.env documents the setup.
Models and the ONNX Runtime dylib are operator-provided at runtime and are
never committed. Cannot be exercised in CI (no models/dylib); the geometry
is unit-tested and the ONNX seam is isolated.
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01JW6ghFMDtnRYuYNzZhb47M
341 lines
12 KiB
Rust
341 lines
12 KiB
Rust
//! ONNX-backed face analyzer (SCRFD detector + ArcFace embedder).
|
||
//!
|
||
//! Compiled only with the `faces-onnx` cargo feature. Mirrors the
|
||
//! immich/InsightFace pipeline: detect faces + 5-point landmarks (SCRFD),
|
||
//! similarity-align each face to the canonical 112×112 template, then embed
|
||
//! (ArcFace) into an L2-normalized 512-d vector. All inference runs on a
|
||
//! blocking thread (`spawn_blocking`) so it never stalls a Tokio worker, and
|
||
//! each ONNX session is serialized behind a `Mutex` (ORT's `run` needs `&mut`).
|
||
//!
|
||
//! The heavy numerical post-processing lives in [`super::face_geometry`] (plain
|
||
//! Rust, unit-tested); this module only wires it to ONNX Runtime.
|
||
//!
|
||
//! **Models are operator-provided at runtime, never committed.** `load` returns
|
||
//! an error (→ caller falls back to the no-op analyzer) if the ONNX Runtime
|
||
//! dylib or either model file is missing or incompatible — the server still
|
||
//! boots. The dylib is loaded via [`ort::init_from`] (a fallible path) rather
|
||
//! than ORT's lazy loader, which would `panic` on a missing library (fatal
|
||
//! under `panic = "abort"`).
|
||
|
||
use std::path::Path;
|
||
use std::sync::{Arc, Mutex};
|
||
|
||
use async_trait::async_trait;
|
||
use image::RgbImage;
|
||
use ort::session::Session;
|
||
use ort::value::Tensor;
|
||
|
||
use super::face_geometry as geom;
|
||
use crate::application::ports::face_ports::FaceAnalyzerPort;
|
||
use crate::common::errors::DomainError;
|
||
use crate::domain::entities::face::{BoundingBox, DetectedFace, EMBEDDING_DIM};
|
||
|
||
/// SCRFD pyramid strides for the 3- and 5-level model variants.
|
||
const STRIDES_3: [u32; 3] = [8, 16, 32];
|
||
const STRIDES_5: [u32; 5] = [8, 16, 32, 64, 128];
|
||
|
||
/// Discard faces smaller than this (original-image pixels) — embeddings of tiny
|
||
/// faces are unreliable.
|
||
const MIN_FACE_PX: f32 = 24.0;
|
||
/// Hard cap on faces processed per image (bounds work on crowd shots).
|
||
const MAX_FACES: usize = 64;
|
||
|
||
/// Output layout of an InsightFace SCRFD model, inferred from its output count.
|
||
#[derive(Clone, Copy)]
|
||
struct ScrfdLayout {
|
||
/// Feature-map count per output kind (3 for strides 8/16/32, 5 with 64/128).
|
||
fmc: usize,
|
||
num_anchors: u32,
|
||
use_kps: bool,
|
||
}
|
||
|
||
impl ScrfdLayout {
|
||
fn from_num_outputs(n: usize) -> Option<Self> {
|
||
match n {
|
||
6 => Some(Self {
|
||
fmc: 3,
|
||
num_anchors: 2,
|
||
use_kps: false,
|
||
}),
|
||
9 => Some(Self {
|
||
fmc: 3,
|
||
num_anchors: 2,
|
||
use_kps: true,
|
||
}),
|
||
10 => Some(Self {
|
||
fmc: 5,
|
||
num_anchors: 1,
|
||
use_kps: false,
|
||
}),
|
||
15 => Some(Self {
|
||
fmc: 5,
|
||
num_anchors: 1,
|
||
use_kps: true,
|
||
}),
|
||
_ => None,
|
||
}
|
||
}
|
||
|
||
fn strides(&self) -> &'static [u32] {
|
||
if self.fmc == 3 {
|
||
&STRIDES_3
|
||
} else {
|
||
&STRIDES_5
|
||
}
|
||
}
|
||
}
|
||
|
||
/// Where to find the runtime + models, plus detector knobs. Borrowed paths;
|
||
/// nothing is retained after [`OnnxFaceAnalyzer::load`].
|
||
pub struct OnnxLoadConfig<'a> {
|
||
/// Path to `libonnxruntime.{so,dylib,dll}`.
|
||
pub dylib: &'a Path,
|
||
/// SCRFD detector `.onnx`.
|
||
pub detector: &'a Path,
|
||
/// ArcFace embedder `.onnx`.
|
||
pub embedder: &'a Path,
|
||
pub det_size: u32,
|
||
pub det_threshold: f32,
|
||
pub nms_threshold: f32,
|
||
/// ORT intra-op threads (0 = let ONNX Runtime decide).
|
||
pub intra_threads: usize,
|
||
}
|
||
|
||
struct Inner {
|
||
detector: Mutex<Session>,
|
||
embedder: Mutex<Session>,
|
||
layout: ScrfdLayout,
|
||
det_size: u32,
|
||
det_threshold: f32,
|
||
nms_threshold: f32,
|
||
}
|
||
|
||
/// Real face analyzer. Cheap to clone (`Arc` inside).
|
||
#[derive(Clone)]
|
||
pub struct OnnxFaceAnalyzer {
|
||
inner: Arc<Inner>,
|
||
}
|
||
|
||
fn dom(e: impl std::fmt::Display) -> DomainError {
|
||
DomainError::internal_error("Faces", e.to_string())
|
||
}
|
||
|
||
fn build_session(path: &Path, intra_threads: usize) -> Result<Session, DomainError> {
|
||
let mut builder = Session::builder().map_err(dom)?;
|
||
if intra_threads > 0 {
|
||
builder = builder.with_intra_threads(intra_threads).map_err(dom)?;
|
||
}
|
||
builder.commit_from_file(path).map_err(dom)
|
||
}
|
||
|
||
impl OnnxFaceAnalyzer {
|
||
/// Load the ONNX Runtime dylib and both models. Returns an error (caller
|
||
/// falls back to the no-op analyzer) on any missing/incompatible artifact.
|
||
pub fn load(cfg: &OnnxLoadConfig<'_>) -> Result<Self, DomainError> {
|
||
// Fallible dylib load — populates ORT's global handle so later calls
|
||
// never hit the panicking lazy loader.
|
||
ort::init_from(cfg.dylib)
|
||
.map_err(|e| dom(format!("ONNX Runtime dylib: {e}")))?
|
||
.commit();
|
||
|
||
let detector = build_session(cfg.detector, cfg.intra_threads)?;
|
||
let embedder = build_session(cfg.embedder, cfg.intra_threads)?;
|
||
|
||
let n_out = detector.outputs().len();
|
||
let layout = ScrfdLayout::from_num_outputs(n_out).ok_or_else(|| {
|
||
dom(format!(
|
||
"detector has {n_out} outputs; expected an SCRFD model (6/9/10/15)"
|
||
))
|
||
})?;
|
||
if !layout.use_kps {
|
||
tracing::warn!(
|
||
target: "oxicloud::faces",
|
||
"SCRFD model has no landmark outputs; face alignment will be approximate"
|
||
);
|
||
}
|
||
|
||
tracing::info!(
|
||
target: "oxicloud::faces",
|
||
"ONNX face analyzer ready (detector {} outputs, embedder loaded, det_size={})",
|
||
n_out, cfg.det_size
|
||
);
|
||
|
||
Ok(Self {
|
||
inner: Arc::new(Inner {
|
||
detector: Mutex::new(detector),
|
||
embedder: Mutex::new(embedder),
|
||
layout,
|
||
det_size: cfg.det_size,
|
||
det_threshold: cfg.det_threshold,
|
||
nms_threshold: cfg.nms_threshold,
|
||
}),
|
||
})
|
||
}
|
||
}
|
||
|
||
impl Inner {
|
||
/// Full synchronous pipeline for one encoded image.
|
||
fn analyze_blocking(&self, image_bytes: &[u8]) -> Result<Vec<DetectedFace>, DomainError> {
|
||
let orig = image::load_from_memory(image_bytes)
|
||
.map_err(|e| dom(format!("decode image: {e}")))?
|
||
.to_rgb8();
|
||
let (w0, h0) = (orig.width(), orig.height());
|
||
if w0 == 0 || h0 == 0 {
|
||
return Ok(Vec::new());
|
||
}
|
||
|
||
let dets = self.detect(&orig)?;
|
||
|
||
let mut faces = Vec::new();
|
||
for det in dets.into_iter().take(MAX_FACES) {
|
||
let fw = det.bbox[2] - det.bbox[0];
|
||
let fh = det.bbox[3] - det.bbox[1];
|
||
if fw < MIN_FACE_PX || fh < MIN_FACE_PX {
|
||
continue;
|
||
}
|
||
let Some(embedding) = self.embed(&orig, &det)? else {
|
||
continue;
|
||
};
|
||
let aligned_quality = {
|
||
let inv = geom::similarity_transform_inverse(&det.kps, &geom::ARCFACE_TEMPLATE);
|
||
let aligned = geom::warp_to_aligned(&orig, &inv);
|
||
geom::laplacian_variance(&aligned)
|
||
};
|
||
let x = (det.bbox[0] / w0 as f32).clamp(0.0, 1.0);
|
||
let y = (det.bbox[1] / h0 as f32).clamp(0.0, 1.0);
|
||
let bw = (fw / w0 as f32).clamp(0.0, 1.0);
|
||
let bh = (fh / h0 as f32).clamp(0.0, 1.0);
|
||
faces.push(DetectedFace {
|
||
bbox: BoundingBox { x, y, w: bw, h: bh },
|
||
det_score: det.score,
|
||
quality: Some(aligned_quality),
|
||
embedding,
|
||
});
|
||
}
|
||
Ok(faces)
|
||
}
|
||
|
||
/// Run SCRFD and return detections in **original-image pixels**.
|
||
fn detect(&self, orig: &RgbImage) -> Result<Vec<geom::Detection>, DomainError> {
|
||
let det = self.det_size;
|
||
let (nw, nh, scale) = geom::letterbox(orig.width(), orig.height(), det);
|
||
let resized = image::imageops::resize(orig, nw, nh, image::imageops::FilterType::Triangle);
|
||
let mut canvas = RgbImage::new(det, det);
|
||
image::imageops::overlay(&mut canvas, &resized, 0, 0);
|
||
let input = geom::chw_normalized(&canvas, 127.5, 1.0 / 128.0);
|
||
let tensor =
|
||
Tensor::from_array(([1_i64, 3, det as i64, det as i64], input)).map_err(dom)?;
|
||
|
||
let layout = self.layout;
|
||
let total = layout.fmc * if layout.use_kps { 3 } else { 2 };
|
||
let raw: Vec<Vec<f32>> = {
|
||
let mut sess = self
|
||
.detector
|
||
.lock()
|
||
.map_err(|_| dom("detector mutex poisoned"))?;
|
||
let outputs = sess.run(ort::inputs![tensor]).map_err(dom)?;
|
||
(0..total)
|
||
.map(|i| {
|
||
outputs[i]
|
||
.try_extract_tensor::<f32>()
|
||
.map(|(_, data)| data.to_vec())
|
||
.map_err(dom)
|
||
})
|
||
.collect::<Result<_, _>>()?
|
||
};
|
||
|
||
let mut dets = Vec::new();
|
||
for (si, &stride) in layout.strides().iter().enumerate() {
|
||
let scores = &raw[si];
|
||
let bbox: Vec<f32> = raw[layout.fmc + si]
|
||
.iter()
|
||
.map(|v| v * stride as f32)
|
||
.collect();
|
||
let kps: Option<Vec<f32>> = if layout.use_kps {
|
||
Some(
|
||
raw[2 * layout.fmc + si]
|
||
.iter()
|
||
.map(|v| v * stride as f32)
|
||
.collect(),
|
||
)
|
||
} else {
|
||
None
|
||
};
|
||
let feat = det / stride;
|
||
geom::decode_stride(
|
||
scores,
|
||
&bbox,
|
||
kps.as_deref(),
|
||
stride,
|
||
feat,
|
||
feat,
|
||
layout.num_anchors,
|
||
self.det_threshold,
|
||
&mut dets,
|
||
);
|
||
}
|
||
|
||
// Scale detector-space coordinates back to the original image.
|
||
let inv_scale = if scale.abs() < 1e-9 { 1.0 } else { 1.0 / scale };
|
||
for d in &mut dets {
|
||
for v in &mut d.bbox {
|
||
*v *= inv_scale;
|
||
}
|
||
for k in &mut d.kps {
|
||
k[0] *= inv_scale;
|
||
k[1] *= inv_scale;
|
||
}
|
||
}
|
||
Ok(geom::nms(dets, self.nms_threshold))
|
||
}
|
||
|
||
/// Align one detection and run the ArcFace embedder. Returns `None` if the
|
||
/// embedder produces an unexpected output length.
|
||
fn embed(
|
||
&self,
|
||
orig: &RgbImage,
|
||
det: &geom::Detection,
|
||
) -> Result<Option<Vec<f32>>, DomainError> {
|
||
let inv = geom::similarity_transform_inverse(&det.kps, &geom::ARCFACE_TEMPLATE);
|
||
let aligned = geom::warp_to_aligned(orig, &inv);
|
||
let input = geom::chw_normalized(&aligned, 127.5, 1.0 / 127.5);
|
||
let size = geom::ALIGN_SIZE as i64;
|
||
let tensor = Tensor::from_array(([1_i64, 3, size, size], input)).map_err(dom)?;
|
||
|
||
let mut embedding: Vec<f32> = {
|
||
let mut sess = self
|
||
.embedder
|
||
.lock()
|
||
.map_err(|_| dom("embedder mutex poisoned"))?;
|
||
let outputs = sess.run(ort::inputs![tensor]).map_err(dom)?;
|
||
let (_, data) = outputs[0].try_extract_tensor::<f32>().map_err(dom)?;
|
||
data.to_vec()
|
||
};
|
||
if embedding.len() != EMBEDDING_DIM {
|
||
tracing::warn!(
|
||
target: "oxicloud::faces",
|
||
"embedder returned {} dims, expected {EMBEDDING_DIM}; skipping face",
|
||
embedding.len()
|
||
);
|
||
return Ok(None);
|
||
}
|
||
geom::l2_normalize(&mut embedding);
|
||
Ok(Some(embedding))
|
||
}
|
||
}
|
||
|
||
#[async_trait]
|
||
impl FaceAnalyzerPort for OnnxFaceAnalyzer {
|
||
fn is_ready(&self) -> bool {
|
||
true
|
||
}
|
||
|
||
async fn analyze(&self, image_bytes: &[u8]) -> Result<Vec<DetectedFace>, DomainError> {
|
||
let inner = self.inner.clone();
|
||
let bytes = image_bytes.to_vec();
|
||
tokio::task::spawn_blocking(move || inner.analyze_blocking(&bytes))
|
||
.await
|
||
.map_err(|e| dom(format!("inference task join: {e}")))?
|
||
}
|
||
}
|