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 | |
| 4 | use crate::face::Image; |
| 5 | use crate::graph::Graph; |
| 6 | use crate::kern::Cpu; |
| 7 | use crate::tensor::Tensor; |
| 8 | |
| 9 | use oxedyne_fe2o3_core::prelude::*; |
| 10 | |
| 11 | /// The three strides the detector's heads sit on. |
| 12 | pub 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. |
| 16 | pub const DIVISOR: usize = 32; |
| 17 | |
| 18 | /// One detected face. |
| 19 | #[derive(Clone, Copy, Debug, PartialEq)] |
| 20 | pub 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 | |
| 36 | impl 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)] |
| 61 | pub 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 | |
| 70 | impl 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)] |
| 79 | pub 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)] |
| 88 | enum 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 | |
| 99 | impl 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. |
| 248 | fn 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. |
| 269 | fn 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)] |
| 295 | mod 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 | } |