Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/src/face/detect.rs

9.4 KiB, 1 run

created by r1870400018:19749, which is this file's identity for as long as the history lasts, whatever it is later renamed to

download · who wrote it · its history

1//! Face detection: the anchor-free decode over three strides, and the
2//! suppression that follows it.
3
4use crate::face::Image;
5use crate::graph::Graph;
6use crate::kern::Cpu;
7use crate::tensor::Tensor;
8
9use oxedyne_fe2o3_core::prelude::*;
10
11/// The three strides the detector's heads sit on.
12pub const STRIDES: [usize; 3] = [8, 16, 32];
13
14/// The canvas extent must be a multiple of this, because the deepest head is
15/// reached by dividing by thirty-two.
16pub const DIVISOR: usize = 32;
17
18/// One detected face.
19#[derive(Clone, Copy, Debug, PartialEq)]
20pub struct Detection {
21 /// Left edge of the box, in canvas pixels.
22 pub x: f32,
23 /// Top edge of the box, in canvas pixels.
24 pub y: f32,
25 /// Box width, in canvas pixels.
26 pub w: f32,
27 /// Box height, in canvas pixels.
28 pub h: f32,
29 /// Right eye, left eye, nose tip, right mouth corner, left mouth corner.
30 pub landmarks: [(f32, f32); 5],
31 /// Confidence, the geometric mean of the classification and objectness
32 /// heads, in `[0, 1]`.
33 pub score: f32,
34}
35
36impl Detection {
37 /// Area of the box after truncation to whole pixels, which is what the
38 /// suppression works on.
39 fn int_box(&self) -> (i64, i64, i64, i64) {
40 (self.x as i64, self.y as i64, self.w as i64, self.h as i64)
41 }
42
43 /// Maps the box and the landmarks back through a letterbox.
44 pub fn unletterbox(&self, lb: &crate::face::Letterbox) -> Self {
45 let s = lb.scale as f32;
46 let mut d = *self;
47 d.x /= s;
48 d.y /= s;
49 d.w /= s;
50 d.h /= s;
51 for p in d.landmarks.iter_mut() {
52 p.0 /= s;
53 p.1 /= s;
54 }
55 d
56 }
57}
58
59/// What the decode and the suppression are allowed to keep.
60#[derive(Clone, Copy, Debug)]
61pub struct DetectorOptions {
62 /// Lowest confidence worth reporting.
63 pub score_threshold: f32,
64 /// Overlap above which the weaker of two boxes is dropped.
65 pub nms_threshold: f32,
66 /// Most candidates carried into the suppression.
67 pub top_k: usize,
68}
69
70impl Default for DetectorOptions {
71 /// The thresholds the reference implementation ships with.
72 fn default() -> Self {
73 Self { score_threshold: 0.9, nms_threshold: 0.3, top_k: 5000 }
74 }
75}
76
77/// A loaded face detector.
78#[derive(Clone, Debug)]
79pub struct Detector {
80 /// The prepared graph.
81 graph: Graph,
82 /// Which head each declared output is, as `(stride index, kind)`.
83 heads: Vec<(usize, Head)>,
84}
85
86/// Which of the four quantities a head predicts.
87#[derive(Clone, Copy, Debug, Eq, PartialEq)]
88enum Head {
89 /// Classification score.
90 Cls,
91 /// Objectness score.
92 Obj,
93 /// Box offsets and log extents.
94 Box,
95 /// Five landmark offsets.
96 Kps,
97}
98
99impl Detector {
100 /// Loads a detector from the bytes of an `.onnx` file.
101 pub fn load(onnx: &[u8]) -> Outcome<Self> {
102 let graph = res!(Graph::load(onnx));
103 let mut heads = Vec::with_capacity(graph.outputs.len());
104 for v in &graph.outputs {
105 let name = &graph.names[*v];
106 let (kind, tail) = if let Some(t) = name.strip_prefix("cls_") {
107 (Head::Cls, t)
108 } else if let Some(t) = name.strip_prefix("obj_") {
109 (Head::Obj, t)
110 } else if let Some(t) = name.strip_prefix("bbox_") {
111 (Head::Box, t)
112 } else if let Some(t) = name.strip_prefix("kps_") {
113 (Head::Kps, t)
114 } else {
115 return Err(err!(
116 "The graph output {} is not a head this detector knows.", name;
117 Invalid, Input, Mismatch));
118 };
119 let stride = res!(tail.parse::<usize>().map_err(|e| err!(e,
120 "The graph output {} does not name a stride.", name; Invalid, Input)));
121 let si = some!(STRIDES.iter().position(|s| *s == stride),
122 "The graph names a stride this detector does not carry.");
123 heads.push((si, kind));
124 }
125 if heads.len() != STRIDES.len() * 4 {
126 return Err(err!(
127 "A detector wants {} heads, the graph declares {}.",
128 STRIDES.len() * 4, heads.len();
129 Invalid, Input, Mismatch));
130 }
131 Ok(Self { graph, heads })
132 }
133
134 /// The prepared graph, for callers that want to time or inspect it.
135 pub fn graph(&self) -> &Graph {
136 &self.graph
137 }
138
139 /// Turns a red-green-blue canvas into the input tensor the detector was
140 /// exported against, which is blue-green-red and unnormalised.
141 pub fn input_tensor(img: &Image<'_>) -> Outcome<Tensor> {
142 if img.channels != 3 {
143 return Err(err!(
144 "The detector wants three channels, the image has {}.", img.channels;
145 Invalid, Input, Mismatch));
146 }
147 if img.width % DIVISOR != 0 || img.height % DIVISOR != 0 {
148 return Err(err!(
149 "The detector wants extents that are multiples of {}, found {} by {}.",
150 DIVISOR, img.width, img.height;
151 Invalid, Input, Range));
152 }
153 let mut data = vec![0.0f32; img.width * img.height * 3];
154 for p in 0..img.width * img.height {
155 data[p * 3] = img.pixels[p * 3 + 2] as f32;
156 data[p * 3 + 1] = img.pixels[p * 3 + 1] as f32;
157 data[p * 3 + 2] = img.pixels[p * 3] as f32;
158 }
159 Tensor::new(vec![1, img.height, img.width, 3], data)
160 }
161
162 /// Detects faces in a canvas whose extents are multiples of thirty-two.
163 pub fn detect(&self, cpu: Cpu, img: &Image<'_>, opts: &DetectorOptions)
164 -> Outcome<Vec<Detection>>
165 {
166 let input = res!(Self::input_tensor(img));
167 let outs = res!(self.graph.run(cpu, input));
168 self.decode(&outs, img.width, img.height, opts)
169 }
170
171 /// Decodes the twelve head outputs into detections and suppresses the
172 /// overlapping ones.
173 pub fn decode(
174 &self,
175 outs: &[Tensor],
176 width: usize,
177 height: usize,
178 opts: &DetectorOptions,
179 )
180 -> Outcome<Vec<Detection>>
181 {
182 if outs.len() != self.heads.len() {
183 return Err(err!(
184 "The graph answered {} outputs, {} were expected.", outs.len(), self.heads.len();
185 Invalid, Mismatch));
186 }
187 let mut cand: Vec<Detection> = Vec::new();
188 for (si, stride) in STRIDES.iter().enumerate() {
189 let cols = width / stride;
190 let rows = height / stride;
191 let cls = res!(self.head(outs, si, Head::Cls));
192 let obj = res!(self.head(outs, si, Head::Obj));
193 let bbox = res!(self.head(outs, si, Head::Box));
194 let kps = res!(self.head(outs, si, Head::Kps));
195 let n = rows * cols;
196 if cls.len() < n || obj.len() < n || bbox.len() < n * 4 || kps.len() < n * 10 {
197 return Err(err!(
198 "The heads at stride {} hold too few values for a {} by {} grid.",
199 stride, rows, cols;
200 Invalid, Mismatch));
201 }
202 for r in 0..rows {
203 for c in 0..cols {
204 let idx = r * cols + c;
205 let cs = cls[idx].clamp(0.0, 1.0);
206 let os = obj[idx].clamp(0.0, 1.0);
207 let score = (cs * os).sqrt();
208 if score < opts.score_threshold {
209 continue;
210 }
211 let sf = *stride as f32;
212 let cx = (c as f32 + bbox[idx * 4]) * sf;
213 let cy = (r as f32 + bbox[idx * 4 + 1]) * sf;
214 let w = bbox[idx * 4 + 2].exp() * sf;
215 let h = bbox[idx * 4 + 3].exp() * sf;
216 let mut landmarks = [(0.0f32, 0.0f32); 5];
217 for (l, lm) in landmarks.iter_mut().enumerate() {
218 lm.0 = (kps[idx * 10 + 2 * l] + c as f32) * sf;
219 lm.1 = (kps[idx * 10 + 2 * l + 1] + r as f32) * sf;
220 }
221 cand.push(Detection {
222 x: cx - w / 2.0,
223 y: cy - h / 2.0,
224 w,
225 h,
226 landmarks,
227 score,
228 });
229 }
230 }
231 }
232 Ok(suppress(cand, opts))
233 }
234
235 /// Finds the values of one head at one stride.
236 fn head<'t>(&self, outs: &'t [Tensor], si: usize, kind: Head) -> Outcome<&'t [f32]> {
237 for (i, (s, k)) in self.heads.iter().enumerate() {
238 if *s == si && *k == kind {
239 return Ok(&outs[i].data);
240 }
241 }
242 Err(err!("The graph declares no {:?} head at stride {}.", kind, STRIDES[si];
243 Invalid, Missing))
244 }
245}
246
247/// Overlap of two whole-pixel boxes, as intersection over union.
248fn iou(a: (i64, i64, i64, i64), b: (i64, i64, i64, i64)) -> f32 {
249 let x0 = a.0.max(b.0);
250 let y0 = a.1.max(b.1);
251 let x1 = (a.0 + a.2).min(b.0 + b.2);
252 let y1 = (a.1 + a.3).min(b.1 + b.3);
253 if x1 <= x0 || y1 <= y0 {
254 return 0.0;
255 }
256 let inter = ((x1 - x0) * (y1 - y0)) as f64;
257 let union = (a.2 * a.3) as f64 + (b.2 * b.3) as f64 - inter;
258 if union <= 0.0 {
259 return 0.0;
260 }
261 (inter / union) as f32
262}
263
264/// Greedy non-maximum suppression, strongest box first.
265///
266/// The boxes are truncated to whole pixels before the overlap is measured,
267/// which is what the reference implementation does and is worth matching,
268/// because on small faces the truncation changes which box survives.
269fn suppress(mut cand: Vec<Detection>, opts: &DetectorOptions) -> Vec<Detection> {
270 if cand.len() <= 1 {
271 return cand;
272 }
273 cand.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(core::cmp::Ordering::Equal));
274 if opts.top_k > 0 && cand.len() > opts.top_k {
275 cand.truncate(opts.top_k);
276 }
277 let mut kept: Vec<Detection> = Vec::new();
278 for d in cand {
279 let db = d.int_box();
280 let mut keep = true;
281 for k in &kept {
282 if iou(db, k.int_box()) > opts.nms_threshold {
283 keep = false;
284 break;
285 }
286 }
287 if keep {
288 kept.push(d);
289 }
290 }
291 kept
292}
293
294#[cfg(test)]
295mod tests {
296 use super::*;
297
298 fn det(x: f32, y: f32, w: f32, h: f32, score: f32) -> Detection {
299 Detection { x, y, w, h, landmarks: [(0.0, 0.0); 5], score }
300 }
301
302 #[test]
303 fn identical_boxes_collapse_to_the_stronger() -> Outcome<()> {
304 let c = vec![det(10.0, 10.0, 20.0, 20.0, 0.9), det(10.0, 10.0, 20.0, 20.0, 0.95)];
305 let k = suppress(c, &DetectorOptions::default());
306 req!(k.len(), 1);
307 req!(k[0].score, 0.95f32);
308 Ok(())
309 }
310
311 #[test]
312 fn separate_boxes_both_survive() -> Outcome<()> {
313 let c = vec![det(0.0, 0.0, 10.0, 10.0, 0.9), det(100.0, 100.0, 10.0, 10.0, 0.91)];
314 let k = suppress(c, &DetectorOptions::default());
315 req!(k.len(), 2);
316 Ok(())
317 }
318
319 #[test]
320 fn overlap_is_measured_on_whole_pixels() -> Outcome<()> {
321 req!(iou((0, 0, 10, 10), (0, 0, 10, 10)), 1.0f32);
322 req!(iou((0, 0, 10, 10), (20, 20, 10, 10)), 0.0f32);
323 let half = iou((0, 0, 10, 10), (5, 0, 10, 10));
324 req!(((half - 1.0 / 3.0).abs() < 1e-6), true);
325 Ok(())
326 }
327}