1use std::path::Path;
17
18#[allow(unused_imports)]
19use image::GenericImageView; use image::{DynamicImage, RgbImage};
21use ort::session::Session;
22use ort::value::TensorRef;
23
24use super::OnnxRuntime;
25
26#[derive(Debug, Clone)]
28pub struct DetectedFace {
29 pub bbox: (f32, f32, f32, f32),
31
32 pub bbox_normalized: (f32, f32, f32, f32),
34
35 pub confidence: f32,
37
38 pub landmarks: [(f32, f32); 5],
41
42 pub aligned_face: Option<RgbImage>,
44}
45
46pub struct FaceDetector {
48 session: Session,
49 input_size: (u32, u32),
50 confidence_threshold: f32,
51 nms_threshold: f32,
52}
53
54struct 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 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 pub fn with_confidence_threshold(mut self, threshold: f32) -> Self {
82 self.confidence_threshold = threshold;
83 self
84 }
85
86 pub fn detect(&mut self, image: &DynamicImage) -> Vec<DetectedFace> {
88 let (orig_width, orig_height) = image.dimensions();
89
90 let input_data = self.preprocess(image);
92
93 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 let mut faces = self.postprocess(&outputs, orig_width, orig_height);
104
105 faces = self.non_max_suppression(faces);
107
108 for face in &mut faces {
110 face.aligned_face = Some(self.align_face(image, &face.landmarks));
111 }
112
113 faces
114 }
115
116 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; 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 fn preprocess(&self, image: &DynamicImage) -> Vec<f32> {
190 let (target_w, target_h) = self.input_size;
191
192 let resized = image.resize_exact(target_w, target_h, image::imageops::FilterType::Triangle);
194
195 let rgb = resized.to_rgb8();
196
197 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; input[hw + idx] = (pixel[1] as f32 - 127.5) / 128.0; input[2 * hw + idx] = (pixel[2] as f32 - 127.5) / 128.0; }
209 }
210
211 input
212 }
213
214 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 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 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 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 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 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 let score = so.scores[idx];
347
348 if score < self.confidence_threshold {
349 continue;
350 }
351
352 let cx = (col as f32 + 0.5) * stride;
354 let cy = (row as f32 + 0.5) * stride;
355
356 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 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 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 fn non_max_suppression(&self, mut faces: Vec<DetectedFace>) -> Vec<DetectedFace> {
424 if faces.is_empty() {
425 return faces;
426 }
427
428 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 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 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 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
522pub 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; }
539
540 let nose_offset_x = nose.0 - eye_center_x;
541 let _nose_offset_y = nose.1 - eye_center_y;
542
543 let yaw_rad = (nose_offset_x / iod).atan();
546 let yaw_deg = yaw_rad.to_degrees();
547
548 yaw_deg.abs()
550}
551
552#[cfg(test)]
553mod tests {
554 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 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 let lmks = [
610 (100.0, 100.0), (200.0, 100.0), (150.0, 150.0), (120.0, 200.0), (180.0, 200.0), ];
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 let lmks = [
625 (100.0, 100.0), (200.0, 100.0), (220.0, 140.0), (120.0, 200.0), (180.0, 200.0), ];
631 let yaw = estimate_yaw_from_landmarks(&lmks);
632 assert!(yaw > 10.0, "profile face yaw={} should be >10", yaw);
633 }
634}