Skip to main content

smriti/ml/
face_detector.rs

1//! SCRFD Face Detection
2//!
3//! Detects faces in images and returns bounding boxes with landmarks.
4//! Uses the SCRFD-10GF model via ONNX Runtime.
5//!
6//! SCRFD outputs 9 tensors (3 strides x 3 types):
7//! ```text
8//!   Stride 8:  scores[12800,1], bboxes[12800,4], landmarks[12800,10]
9//!   Stride 16: scores[3200,1],  bboxes[3200,4],  landmarks[3200,10]
10//!   Stride 32: scores[800,1],   bboxes[800,4],   landmarks[800,10]
11//! ```
12//!
13//! Bboxes are anchor-distance format (left, top, right, bottom distances from anchor center).
14//! Landmarks are (dx, dy) offsets from anchor center for 5 keypoints.
15
16use std::path::Path;
17
18#[allow(unused_imports)]
19use image::GenericImageView; // needed by `DynamicImage::dimensions()` in some configurations
20use image::{DynamicImage, RgbImage};
21use ort::session::Session;
22use ort::value::TensorRef;
23
24use super::OnnxRuntime;
25
26/// A detected face with bounding box and landmarks
27#[derive(Debug, Clone)]
28pub struct DetectedFace {
29    /// Bounding box (x, y, width, height) in pixel coordinates
30    pub bbox: (f32, f32, f32, f32),
31
32    /// Normalized bounding box (0-1 range)
33    pub bbox_normalized: (f32, f32, f32, f32),
34
35    /// Detection confidence (0-1)
36    pub confidence: f32,
37
38    /// 5-point landmarks: left_eye, right_eye, nose, left_mouth, right_mouth
39    /// Each point is (x, y) in pixel coordinates
40    pub landmarks: [(f32, f32); 5],
41
42    /// Cropped and aligned face image (112x112)
43    pub aligned_face: Option<RgbImage>,
44}
45
46/// SCRFD Face Detector
47pub struct FaceDetector {
48    session: Session,
49    input_size: (u32, u32),
50    confidence_threshold: f32,
51    nms_threshold: f32,
52}
53
54/// Per-stride output from the model
55struct StrideOutput {
56    scores: Vec<f32>,
57    bboxes: Vec<f32>,
58    landmarks: Vec<f32>,
59    num_anchors: usize,
60    stride: u32,
61}
62
63impl FaceDetector {
64    /// Load the SCRFD model with a specific thread count per session.
65    pub fn new_with_threads<P: AsRef<Path>>(
66        runtime: &OnnxRuntime,
67        model_path: P,
68        intra_threads: usize,
69    ) -> ort::Result<Self> {
70        let session = runtime.load_model_with_threads(model_path, intra_threads)?;
71
72        Ok(Self {
73            session,
74            input_size: (640, 640),
75            confidence_threshold: 0.5,
76            nms_threshold: 0.4,
77        })
78    }
79
80    /// Set confidence threshold for detection
81    pub fn with_confidence_threshold(mut self, threshold: f32) -> Self {
82        self.confidence_threshold = threshold;
83        self
84    }
85
86    /// Detect faces in an image
87    pub fn detect(&mut self, image: &DynamicImage) -> Vec<DetectedFace> {
88        let (orig_width, orig_height) = image.dimensions();
89
90        // Preprocess image to NCHW float tensor
91        let input_data = self.preprocess(image);
92
93        // Run inference
94        let outputs = match self.run_inference(&input_data) {
95            Ok(o) => o,
96            Err(e) => {
97                tracing::error!("Face detection inference failed: {}", e);
98                return Vec::new();
99            }
100        };
101
102        // Post-process outputs
103        let mut faces = self.postprocess(&outputs, orig_width, orig_height);
104
105        // Apply NMS
106        faces = self.non_max_suppression(faces);
107
108        // Align faces for embedding
109        for face in &mut faces {
110            face.aligned_face = Some(self.align_face(image, &face.landmarks));
111        }
112
113        faces
114    }
115
116    /// Detect faces with adaptive multi-scale fallback for high-res images.
117    ///
118    /// Pass 1: normal full-frame detection.
119    /// Pass 2 (fallback): overlapping high-res tiles if pass 1 found no faces
120    /// and image is large, to recover small/profile faces lost by global resizing.
121    pub fn detect_adaptive(&mut self, image: &DynamicImage) -> Vec<DetectedFace> {
122        let first_pass = self.detect(image);
123        if !first_pass.is_empty() {
124            return first_pass;
125        }
126
127        let (orig_w, orig_h) = image.dimensions();
128        if orig_w.max(orig_h) <= 2048 {
129            return first_pass;
130        }
131
132        let tile = 1600u32;
133        let step = 1200u32; // overlap for boundary faces
134        let mut all_faces = Vec::new();
135
136        let max_x = orig_w.saturating_sub(1);
137        let max_y = orig_h.saturating_sub(1);
138
139        let mut y = 0u32;
140        while y <= max_y {
141            let mut x = 0u32;
142            while x <= max_x {
143                let crop_w = tile.min(orig_w.saturating_sub(x)).max(1);
144                let crop_h = tile.min(orig_h.saturating_sub(y)).max(1);
145                let crop = image.crop_imm(x, y, crop_w, crop_h);
146
147                let mut local = self.detect(&crop);
148                if !local.is_empty() {
149                    for face in &mut local {
150                        let (bx, by, bw, bh) = face.bbox;
151                        let gx = bx + x as f32;
152                        let gy = by + y as f32;
153                        face.bbox = (gx, gy, bw, bh);
154                        face.bbox_normalized = (
155                            gx / orig_w as f32,
156                            gy / orig_h as f32,
157                            bw / orig_w as f32,
158                            bh / orig_h as f32,
159                        );
160
161                        for lm in &mut face.landmarks {
162                            lm.0 += x as f32;
163                            lm.1 += y as f32;
164                        }
165                    }
166                    all_faces.extend(local);
167                }
168
169                if x + step >= orig_w {
170                    break;
171                }
172                x += step;
173            }
174
175            if y + step >= orig_h {
176                break;
177            }
178            y += step;
179        }
180
181        if all_faces.is_empty() {
182            return all_faces;
183        }
184
185        self.non_max_suppression(all_faces)
186    }
187
188    /// Preprocess image for SCRFD: resize to 640x640, normalize, produce NCHW vec
189    fn preprocess(&self, image: &DynamicImage) -> Vec<f32> {
190        let (target_w, target_h) = self.input_size;
191
192        // Resize to target dimensions
193        let resized = image.resize_exact(target_w, target_h, image::imageops::FilterType::Triangle);
194
195        let rgb = resized.to_rgb8();
196
197        // Convert to NCHW float and normalize (mean=127.5, std=128)
198        let mut input = vec![0.0f32; (3 * target_h * target_w) as usize];
199        let hw = (target_h * target_w) as usize;
200
201        for y in 0..target_h {
202            for x in 0..target_w {
203                let pixel = rgb.get_pixel(x, y);
204                let idx = (y * target_w + x) as usize;
205                input[idx] = (pixel[0] as f32 - 127.5) / 128.0; // R channel
206                input[hw + idx] = (pixel[1] as f32 - 127.5) / 128.0; // G channel
207                input[2 * hw + idx] = (pixel[2] as f32 - 127.5) / 128.0; // B channel
208            }
209        }
210
211        input
212    }
213
214    /// Run ONNX inference and extract output tensors with their shapes
215    fn run_inference(&mut self, input_data: &[f32]) -> ort::Result<Vec<(Vec<i64>, Vec<f32>)>> {
216        let (target_w, target_h) = self.input_size;
217
218        // Create input tensor (shape [1, 3, 640, 640])
219        let input_tensor = TensorRef::<f32>::from_array_view((
220            vec![1i64, 3, target_h as i64, target_w as i64],
221            input_data,
222        ))?;
223
224        let outputs = self.session.run(ort::inputs![input_tensor])?;
225
226        // Extract all output tensors with shapes
227        let mut results = Vec::new();
228        for (_name, value) in outputs.iter() {
229            let (shape, data) = value.try_extract_tensor::<f32>()?;
230            results.push((shape.to_vec(), data.to_vec()));
231        }
232
233        Ok(results)
234    }
235
236    /// Classify and group the 9 output tensors into 3 stride groups.
237    ///
238    /// SCRFD outputs 9 tensors. We classify them by their last dimension:
239    ///   - dim=1 → scores
240    ///   - dim=4 → bounding boxes (anchor-distance format)
241    ///   - dim=10 → landmarks (5 keypoints * 2 coords)
242    ///
243    /// Within each type, they are ordered by descending anchor count (stride 8, 16, 32).
244    fn group_outputs(&self, outputs: &[(Vec<i64>, Vec<f32>)]) -> Vec<StrideOutput> {
245        let strides = [8u32, 16, 32];
246
247        let mut score_tensors: Vec<(usize, &Vec<f32>)> = Vec::new();
248        let mut bbox_tensors: Vec<(usize, &Vec<f32>)> = Vec::new();
249        let mut lmk_tensors: Vec<(usize, &Vec<f32>)> = Vec::new();
250
251        for (shape, data) in outputs {
252            let last_dim = *shape.last().unwrap_or(&0);
253            let num_anchors = if shape.len() >= 2 {
254                shape[shape.len() - 2] as usize
255            } else {
256                data.len() / last_dim as usize
257            };
258
259            match last_dim {
260                1 => score_tensors.push((num_anchors, data)),
261                4 => bbox_tensors.push((num_anchors, data)),
262                10 => lmk_tensors.push((num_anchors, data)),
263                _ => {
264                    tracing::debug!("Unknown output tensor shape: {:?}", shape);
265                }
266            }
267        }
268
269        // Sort each group by descending anchor count (stride 8 has most anchors)
270        score_tensors.sort_by_key(|x| std::cmp::Reverse(x.0));
271        bbox_tensors.sort_by_key(|x| std::cmp::Reverse(x.0));
272        lmk_tensors.sort_by_key(|x| std::cmp::Reverse(x.0));
273
274        let mut stride_outputs = Vec::new();
275
276        for i in 0..strides
277            .len()
278            .min(score_tensors.len())
279            .min(bbox_tensors.len())
280        {
281            let num_anchors = score_tensors[i].0;
282            let has_landmarks = i < lmk_tensors.len();
283
284            stride_outputs.push(StrideOutput {
285                scores: score_tensors[i].1.clone(),
286                bboxes: bbox_tensors[i].1.clone(),
287                landmarks: if has_landmarks {
288                    lmk_tensors[i].1.clone()
289                } else {
290                    vec![0.0; num_anchors * 10]
291                },
292                num_anchors,
293                stride: strides[i],
294            });
295        }
296
297        stride_outputs
298    }
299
300    /// Post-process SCRFD outputs to DetectedFace structs
301    ///
302    /// SCRFD uses anchor-based detection across 3 stride levels.
303    /// Each anchor center is at (col * stride + stride/2, row * stride + stride/2).
304    /// BBox outputs are distances from anchor center: (left, top, right, bottom).
305    /// Landmark outputs are (dx, dy) offsets from anchor center for each keypoint.
306    fn postprocess(
307        &self,
308        outputs: &[(Vec<i64>, Vec<f32>)],
309        orig_width: u32,
310        orig_height: u32,
311    ) -> Vec<DetectedFace> {
312        let mut faces = Vec::new();
313
314        if outputs.len() < 6 {
315            tracing::warn!(
316                "SCRFD: Expected at least 6 output tensors, got {}",
317                outputs.len()
318            );
319            return faces;
320        }
321
322        let stride_outputs = self.group_outputs(outputs);
323
324        let (input_w, input_h) = self.input_size;
325        let scale_x = orig_width as f32 / input_w as f32;
326        let scale_y = orig_height as f32 / input_h as f32;
327
328        for so in &stride_outputs {
329            let stride = so.stride as f32;
330            let grid_w = input_w as f32 / stride;
331            let grid_h = input_h as f32 / stride;
332            let grid_cols = grid_w as usize;
333            let grid_rows = grid_h as usize;
334            let anchors_per_cell = so.num_anchors / (grid_cols * grid_rows).max(1);
335
336            for row in 0..grid_rows {
337                for col in 0..grid_cols {
338                    for a in 0..anchors_per_cell {
339                        let idx = (row * grid_cols + col) * anchors_per_cell + a;
340                        if idx >= so.num_anchors {
341                            continue;
342                        }
343
344                        // Score — SCRFD already applies sigmoid internally,
345                        // so raw output is in [0, 1] range.
346                        let score = so.scores[idx];
347
348                        if score < self.confidence_threshold {
349                            continue;
350                        }
351
352                        // Anchor center in input image coords
353                        let cx = (col as f32 + 0.5) * stride;
354                        let cy = (row as f32 + 0.5) * stride;
355
356                        // Decode bounding box (distances from anchor center)
357                        let bbox_base = idx * 4;
358                        if bbox_base + 3 >= so.bboxes.len() {
359                            continue;
360                        }
361                        let dl = so.bboxes[bbox_base] * stride;
362                        let dt = so.bboxes[bbox_base + 1] * stride;
363                        let dr = so.bboxes[bbox_base + 2] * stride;
364                        let db = so.bboxes[bbox_base + 3] * stride;
365
366                        let x1 = (cx - dl) * scale_x;
367                        let y1 = (cy - dt) * scale_y;
368                        let x2 = (cx + dr) * scale_x;
369                        let y2 = (cy + db) * scale_y;
370
371                        // Clamp to image bounds
372                        let x1 = x1.max(0.0).min(orig_width as f32);
373                        let y1 = y1.max(0.0).min(orig_height as f32);
374                        let x2 = x2.max(0.0).min(orig_width as f32);
375                        let y2 = y2.max(0.0).min(orig_height as f32);
376
377                        let width = x2 - x1;
378                        let height = y2 - y1;
379
380                        if width < 10.0 || height < 10.0 {
381                            continue;
382                        }
383
384                        // Decode landmarks
385                        let lm_base = idx * 10;
386                        let mut lmks = [(0.0f32, 0.0f32); 5];
387                        if lm_base + 9 < so.landmarks.len() {
388                            for (j, lmk) in lmks.iter_mut().enumerate() {
389                                *lmk = (
390                                    (cx + so.landmarks[lm_base + j * 2] * stride) * scale_x,
391                                    (cy + so.landmarks[lm_base + j * 2 + 1] * stride) * scale_y,
392                                );
393                            }
394                        }
395
396                        faces.push(DetectedFace {
397                            bbox: (x1, y1, width, height),
398                            bbox_normalized: (
399                                x1 / orig_width as f32,
400                                y1 / orig_height as f32,
401                                width / orig_width as f32,
402                                height / orig_height as f32,
403                            ),
404                            confidence: score,
405                            landmarks: lmks,
406                            aligned_face: None,
407                        });
408                    }
409                }
410            }
411        }
412
413        tracing::debug!(
414            "SCRFD raw detections: {} (before NMS, threshold={})",
415            faces.len(),
416            self.confidence_threshold
417        );
418
419        faces
420    }
421
422    /// Non-maximum suppression to remove overlapping detections
423    fn non_max_suppression(&self, mut faces: Vec<DetectedFace>) -> Vec<DetectedFace> {
424        if faces.is_empty() {
425            return faces;
426        }
427
428        // Sort by confidence (descending)
429        faces.sort_by(|a, b| {
430            b.confidence
431                .partial_cmp(&a.confidence)
432                .unwrap_or(std::cmp::Ordering::Equal)
433        });
434
435        let mut keep = vec![true; faces.len()];
436
437        for i in 0..faces.len() {
438            if !keep[i] {
439                continue;
440            }
441
442            for j in (i + 1)..faces.len() {
443                if !keep[j] {
444                    continue;
445                }
446
447                let iou = self.calculate_iou(&faces[i].bbox, &faces[j].bbox);
448                if iou > self.nms_threshold {
449                    keep[j] = false;
450                }
451            }
452        }
453
454        faces
455            .into_iter()
456            .enumerate()
457            .filter(|(i, _)| keep[*i])
458            .map(|(_, f)| f)
459            .collect()
460    }
461
462    /// Calculate Intersection over Union for two bounding boxes
463    /// Each box is (x, y, width, height)
464    fn calculate_iou(&self, box1: &(f32, f32, f32, f32), box2: &(f32, f32, f32, f32)) -> f32 {
465        let (x1, y1, w1, h1) = *box1;
466        let (x2, y2, w2, h2) = *box2;
467
468        let xi1 = x1.max(x2);
469        let yi1 = y1.max(y2);
470        let xi2 = (x1 + w1).min(x2 + w2);
471        let yi2 = (y1 + h1).min(y2 + h2);
472
473        let inter_width = (xi2 - xi1).max(0.0);
474        let inter_height = (yi2 - yi1).max(0.0);
475        let inter_area = inter_width * inter_height;
476
477        let area1 = w1 * h1;
478        let area2 = w2 * h2;
479        let union_area = area1 + area2 - inter_area;
480
481        if union_area > 0.0 {
482            inter_area / union_area
483        } else {
484            0.0
485        }
486    }
487
488    /// Align face using all 5 landmarks via similarity transform onto the
489    /// canonical InsightFace 112x112 template. Falls back to eye-center crop
490    /// if the landmarks are degenerate (rare).
491    fn align_face(&self, image: &DynamicImage, landmarks: &[(f32, f32); 5]) -> RgbImage {
492        if let Some(aligned) = super::alignment::align_face_112(image, landmarks) {
493            return aligned;
494        }
495
496        // Degenerate landmarks: fall back to the old eye-center crop.
497        let left_eye = landmarks[0];
498        let right_eye = landmarks[1];
499        let eye_center = (
500            (left_eye.0 + right_eye.0) / 2.0,
501            (left_eye.1 + right_eye.1) / 2.0,
502        );
503        let eye_dist =
504            ((right_eye.0 - left_eye.0).powi(2) + (right_eye.1 - left_eye.1).powi(2)).sqrt();
505        let face_size = (eye_dist * 2.5).max(10.0);
506
507        let img_w = image.width();
508        let img_h = image.height();
509        let x = (eye_center.0 - face_size / 2.0).max(0.0) as u32;
510        let y = (eye_center.1 - face_size / 2.0).max(0.0) as u32;
511        let x = x.min(img_w.saturating_sub(1));
512        let y = y.min(img_h.saturating_sub(1));
513        let crop_w = (face_size as u32).min(img_w.saturating_sub(x)).max(1);
514        let crop_h = (face_size as u32).min(img_h.saturating_sub(y)).max(1);
515
516        let cropped = image.crop_imm(x, y, crop_w, crop_h);
517        let resized = cropped.resize_exact(112, 112, image::imageops::FilterType::Lanczos3);
518        resized.to_rgb8()
519    }
520}
521
522/// Estimate yaw angle (degrees) from 5-point SCRFD landmarks.
523///
524/// Uses the nose offset from the eyes-midpoint, normalised by
525/// inter-ocular distance. Positive = right-facing, negative = left-facing.
526/// The absolute value indicates how far the face is turned from frontal.
527pub fn estimate_yaw_from_landmarks(landmarks: &[(f32, f32); 5]) -> f32 {
528    let left_eye = landmarks[0];
529    let right_eye = landmarks[1];
530    let nose = landmarks[2];
531
532    let eye_center_x = (left_eye.0 + right_eye.0) / 2.0;
533    let eye_center_y = (left_eye.1 + right_eye.1) / 2.0;
534    let iod = ((right_eye.0 - left_eye.0).powi(2) + (right_eye.1 - left_eye.1).powi(2)).sqrt();
535
536    if iod < 1.0 {
537        return 0.0; // degenerate landmarks
538    }
539
540    let nose_offset_x = nose.0 - eye_center_x;
541    let _nose_offset_y = nose.1 - eye_center_y;
542
543    // The yaw is roughly the horizontal offset normalised by IOD.
544    // A small vertical correction accounts for head tilt.
545    let yaw_rad = (nose_offset_x / iod).atan();
546    let yaw_deg = yaw_rad.to_degrees();
547
548    // Return absolute yaw for gating purposes.
549    yaw_deg.abs()
550}
551
552#[cfg(test)]
553mod tests {
554    /// Helper to create a FaceDetector-like struct for testing IoU
555    struct IouTester;
556
557    impl IouTester {
558        fn calculate_iou(box1: &(f32, f32, f32, f32), box2: &(f32, f32, f32, f32)) -> f32 {
559            let (x1, y1, w1, h1) = *box1;
560            let (x2, y2, w2, h2) = *box2;
561
562            let xi1 = x1.max(x2);
563            let yi1 = y1.max(y2);
564            let xi2 = (x1 + w1).min(x2 + w2);
565            let yi2 = (y1 + h1).min(y2 + h2);
566
567            let inter_width = (xi2 - xi1).max(0.0);
568            let inter_height = (yi2 - yi1).max(0.0);
569            let inter_area = inter_width * inter_height;
570
571            let area1 = w1 * h1;
572            let area2 = w2 * h2;
573            let union_area = area1 + area2 - inter_area;
574
575            if union_area > 0.0 {
576                inter_area / union_area
577            } else {
578                0.0
579            }
580        }
581    }
582
583    #[test]
584    fn test_iou_same_box() {
585        let box1 = (0.0, 0.0, 100.0, 100.0);
586        assert!((IouTester::calculate_iou(&box1, &box1) - 1.0).abs() < 0.001);
587    }
588
589    #[test]
590    fn test_iou_non_overlapping() {
591        let box1 = (0.0, 0.0, 100.0, 100.0);
592        let box2 = (200.0, 200.0, 100.0, 100.0);
593        assert!((IouTester::calculate_iou(&box1, &box2) - 0.0).abs() < 0.001);
594    }
595
596    #[test]
597    fn test_iou_partial_overlap() {
598        let box1 = (0.0, 0.0, 100.0, 100.0);
599        let box2 = (50.0, 50.0, 100.0, 100.0);
600        let iou = IouTester::calculate_iou(&box1, &box2);
601        // Intersection: 50x50 = 2500, Union: 10000 + 10000 - 2500 = 17500
602        assert!((iou - 2500.0 / 17500.0).abs() < 0.001);
603    }
604
605    #[test]
606    fn test_estimate_yaw_frontal() {
607        use super::estimate_yaw_from_landmarks;
608        // Frontal face: eyes level, nose centered
609        let lmks = [
610            (100.0, 100.0), // left eye
611            (200.0, 100.0), // right eye
612            (150.0, 150.0), // nose (centered)
613            (120.0, 200.0), // left mouth
614            (180.0, 200.0), // right mouth
615        ];
616        let yaw = estimate_yaw_from_landmarks(&lmks);
617        assert!(yaw < 5.0, "frontal face yaw={} should be <5", yaw);
618    }
619
620    #[test]
621    fn test_estimate_yaw_profile() {
622        use super::estimate_yaw_from_landmarks;
623        // Right-facing face: nose offset to the right
624        let lmks = [
625            (100.0, 100.0), // left eye
626            (200.0, 100.0), // right eye
627            (220.0, 140.0), // nose (right offset)
628            (120.0, 200.0), // left mouth
629            (180.0, 200.0), // right mouth
630        ];
631        let yaw = estimate_yaw_from_landmarks(&lmks);
632        assert!(yaw > 10.0, "profile face yaw={} should be >10", yaw);
633    }
634}