Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/src/object.rs

16.3 KiB, 1 run

created by r1870400018:21039, 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//! Category detection: what is in a photograph, from a fixed list of everyday
2//! things.
3//!
4//! This is the second detector the crate carries and it answers a different
5//! question from the first. The face detector finds one kind of thing and says
6//! where its eyes are; this one finds eighty kinds and says only where each is.
7//!
8//! # What the network hands back, and what has to be done with it
9//!
10//! Three heads, at strides of eight, sixteen and thirty-two, each giving a class
11//! score per position and a box per position. The box is not four numbers. It is
12//! four *distributions* -- one per side, over eight bins -- and the distance to
13//! that side is their mean, which is the arrangement a generalised focal loss
14//! trains. So each side is a softmax and then a dot product against `0..8`,
15//! multiplied by the stride, and the box is that distance out from the position
16//! rather than a corner in its own right.
17//!
18//! The positions are not at the centres of their cells. They sit at
19//! `i·stride + (stride − 1)/2`, which is half a sample short of the centre, and
20//! every reference implementation of this model does the same. Taking the centre
21//! instead moves every box by up to fifteen pixels at the coarsest head.
22//!
23//! # Channel order and normalisation
24//!
25//! The network was exported against blue-green-red input, and against the
26//! ImageNet statistics in that order. [`Detector::input_tensor`] takes ordinary
27//! red-green-blue pixels and puts them the way the network was trained, so a
28//! caller never has to know, exactly as the face detector does.
29
30use crate::face::{Image, Letterbox};
31use crate::graph::Graph;
32use crate::kern::Cpu;
33use crate::tensor::Tensor;
34
35use oxedyne_fe2o3_core::prelude::*;
36
37/// The strides the three heads predict at.
38pub const STRIDES: [usize; 3] = [8, 16, 32];
39
40/// Bins in each side's distribution.
41pub const BINS: usize = 8;
42
43/// Sides of a box, in the order left, top, right, bottom.
44pub const SIDES: usize = 4;
45
46/// Categories the network was trained on.
47pub const CATEGORIES: usize = 80;
48
49/// The canvas the network wants, in pixels each way.
50pub const SIDE: usize = 416;
51
52/// The eighty categories, in the order the network scores them.
53pub const NAMES: [&str; CATEGORIES] = [
54 "person", "bicycle", "car", "motorcycle", "airplane", "bus", "train",
55 "truck", "boat", "traffic light", "fire hydrant", "stop sign",
56 "parking meter", "bench", "bird", "cat", "dog", "horse", "sheep", "cow",
57 "elephant", "bear", "zebra", "giraffe", "backpack", "umbrella", "handbag",
58 "tie", "suitcase", "frisbee", "skis", "snowboard", "sports ball", "kite",
59 "baseball bat", "baseball glove", "skateboard", "surfboard",
60 "tennis racket", "bottle", "wine glass", "cup", "fork", "knife", "spoon",
61 "bowl", "banana", "apple", "sandwich", "orange", "broccoli", "carrot",
62 "hot dog", "pizza", "donut", "cake", "chair", "couch", "potted plant",
63 "bed", "dining table", "toilet", "tv", "laptop", "mouse", "remote",
64 "keyboard", "cell phone", "microwave", "oven", "toaster", "sink",
65 "refrigerator", "book", "clock", "vase", "scissors", "teddy bear",
66 "hair drier", "toothbrush",
67];
68
69/// The categories that are animals, by index into [`NAMES`].
70///
71/// Every four-legged and winged thing the list carries, wild ones included: a
72/// caller after somebody's pets wants the whole set, because a network that has
73/// to choose between `dog` and `bear` for a large dark animal is answering a
74/// question nobody asked.
75pub const ANIMALS: [usize; 10] = [14, 15, 16, 17, 18, 19, 20, 21, 22, 23];
76
77/// Whether a category is one of the animals.
78pub fn is_animal(class: usize) -> bool {
79 ANIMALS.contains(&class)
80}
81
82/// One detected thing.
83#[derive(Clone, Copy, Debug, PartialEq)]
84pub struct Object {
85 /// Index into [`NAMES`].
86 pub class: usize,
87 /// Confidence in `[0, 1]`.
88 pub score: f32,
89 /// Left edge of the box, in canvas pixels.
90 pub x: f32,
91 /// Top edge, in canvas pixels.
92 pub y: f32,
93 /// Width, in canvas pixels.
94 pub w: f32,
95 /// Height, in canvas pixels.
96 pub h: f32,
97}
98
99impl Object {
100 /// The name of the category.
101 pub fn name(&self) -> &'static str {
102 NAMES.get(self.class).copied().unwrap_or("?")
103 }
104
105 /// Whether this is one of the animals.
106 pub fn is_animal(&self) -> bool {
107 is_animal(self.class)
108 }
109
110 /// Area of the box after truncation to whole pixels, which is what the
111 /// suppression works on.
112 fn int_box(&self) -> (i64, i64, i64, i64) {
113 (self.x as i64, self.y as i64, self.w as i64, self.h as i64)
114 }
115
116 /// Maps the box back through a letterbox, into the original frame.
117 pub fn unletterbox(&self, lb: &Letterbox) -> Self {
118 let s = lb.scale as f32;
119 let mut o = *self;
120 o.x /= s;
121 o.y /= s;
122 o.w /= s;
123 o.h /= s;
124 o
125 }
126}
127
128/// What the decode and the suppression are allowed to keep.
129#[derive(Clone, Copy, Debug)]
130pub struct Options {
131 /// Lowest confidence worth reporting.
132 pub score_threshold: f32,
133 /// Overlap above which the weaker of two boxes is dropped.
134 pub nms_threshold: f32,
135 /// Most candidates carried out of one head, before the threshold.
136 pub pre_k: usize,
137 /// Most candidates carried into the suppression.
138 pub top_k: usize,
139}
140
141impl Default for Options {
142 /// The thresholds the reference implementation ships with.
143 ///
144 /// They are the reference's and not a recommendation: on a real photograph
145 /// library a score of `0.35` admits far more than it should, and a caller
146 /// after a usable answer should raise it.
147 fn default() -> Self {
148 Self {
149 score_threshold: 0.35,
150 nms_threshold: 0.6,
151 pre_k: 1000,
152 top_k: 0,
153 }
154 }
155}
156
157/// A loaded category detector.
158#[derive(Clone, Debug)]
159pub struct Detector {
160 /// The prepared graph.
161 graph: Graph,
162}
163
164impl Detector {
165 /// Reads a model from the bytes of an `.onnx` file.
166 pub fn load(onnx: &[u8]) -> Outcome<Self> {
167 let graph = res!(Graph::load(onnx));
168 Ok(Self { graph })
169 }
170
171 /// The prepared graph, for a caller that wants to run it itself.
172 pub fn graph(&self) -> &Graph {
173 &self.graph
174 }
175
176 /// Turns a letterboxed canvas into the tensor the network wants.
177 ///
178 /// The image must already be the canvas: square, [`SIDE`] each way, with the
179 /// photograph fitted into it. Channels are reversed and the ImageNet
180 /// statistics applied in the network's own order.
181 pub fn input_tensor(img: &Image<'_>) -> Outcome<Tensor> {
182 if img.width != SIDE || img.height != SIDE {
183 return Err(err!(
184 "The detector wants a canvas of {} by {}, and was given {} by {}.",
185 SIDE, SIDE, img.width, img.height;
186 Invalid, Input, Mismatch));
187 }
188 if img.channels < 3 {
189 return Err(err!(
190 "The detector wants three channels, and was given {}.", img.channels;
191 Invalid, Input, Mismatch));
192 }
193 // Blue, green, red -- the order the network was exported against.
194 const MEAN: [f32; 3] = [103.53, 116.28, 123.675];
195 const STD: [f32; 3] = [57.375, 57.12, 58.395];
196 let n = SIDE * SIDE;
197 let mut data = vec![0.0f32; n * 3];
198 for p in 0..n {
199 let src = p * img.channels;
200 for c in 0..3 {
201 // Channel `c` of the network is channel `2 - c` of the image.
202 let v = img.pixels[src + (2 - c)] as f32;
203 data[p * 3 + c] = (v - MEAN[c]) / STD[c];
204 }
205 }
206 Tensor::new(vec![1, SIDE, SIDE, 3], data)
207 }
208
209 /// Runs the network over a canvas and answers what it found.
210 pub fn detect(&self, cpu: Cpu, img: &Image<'_>, opts: &Options) -> Outcome<Vec<Object>> {
211 let input = res!(Self::input_tensor(img));
212 let out = res!(self.graph.run(cpu, input));
213 decode(&out, opts)
214 }
215}
216
217/// Turns the network's outputs into boxes.
218///
219/// Takes the tensors rather than the model, because nothing here needs the
220/// weights: a caller holding the outputs from anywhere can decode them, which is
221/// what lets the decode be checked against another implementation on its own.
222///
223/// The heads arrive in whatever order the model declared them, so they are
224/// paired by what they are rather than by position: a head with [`CATEGORIES`]
225/// values a position is the scores, one with `SIDES · BINS` is the boxes, and
226/// the two belonging together have the same number of positions.
227pub fn decode(out: &[Tensor], opts: &Options) -> Outcome<Vec<Object>> {
228 let mut scores: Vec<(usize, &Tensor)> = Vec::new();
229 let mut boxes: Vec<(usize, &Tensor)> = Vec::new();
230 for t in out {
231 if t.dims.len() != 3 {
232 return Err(err!(
233 "A head of shape {:?} is not a run of positions.", t.dims;
234 Invalid, Input, Mismatch));
235 }
236 let (points, width) = (t.dims[1], t.dims[2]);
237 if width == CATEGORIES {
238 scores.push((points, t));
239 } else if width == SIDES * BINS {
240 boxes.push((points, t));
241 } else {
242 return Err(err!(
243 "A head of {} values a position is neither scores nor boxes.", width;
244 Invalid, Input, Mismatch));
245 }
246 }
247 if scores.len() != boxes.len() {
248 return Err(err!(
249 "The network gave {} score heads and {} box heads.", scores.len(), boxes.len();
250 Invalid, Input, Mismatch));
251 }
252
253 let mut cand: Vec<Object> = Vec::new();
254 for (points, cls) in &scores {
255 let bx = match boxes.iter().find(|(p, _)| p == points) {
256 Some((_, t)) => *t,
257 None => return Err(err!(
258 "No box head has the {} positions the scores do.", points;
259 Invalid, Input, Mismatch)),
260 };
261 res!(level(*points, cls, bx, opts, &mut cand));
262 }
263
264 Ok(suppress(cand, opts))
265}
266
267/// Decodes one head.
268fn level(
269 points: usize,
270 cls: &Tensor,
271 bx: &Tensor,
272 opts: &Options,
273 out: &mut Vec<Object>,
274 )
275 -> Outcome<()>
276 {
277 // The head is a square grid over the canvas, so its side gives its stride.
278 let side = (points as f64).sqrt().round() as usize;
279 if side * side != points || side == 0 {
280 return Err(err!(
281 "A head of {} positions is not a square grid.", points;
282 Invalid, Input, Mismatch));
283 }
284 if SIDE % side != 0 {
285 return Err(err!(
286 "A grid of {} does not divide a canvas of {}.", side, SIDE;
287 Invalid, Input, Mismatch));
288 }
289 let stride = SIDE / side;
290 if !STRIDES.contains(&stride) {
291 return Err(err!(
292 "A head at stride {} is not one this detector predicts at.", stride;
293 Invalid, Input, Unimplemented));
294 }
295
296 // The strongest class at each position, and the order to consider them.
297 let mut best: Vec<(f32, usize)> = Vec::with_capacity(points);
298 for p in 0..points {
299 let row = &cls.data[p * CATEGORIES..(p + 1) * CATEGORIES];
300 let mut top = (0.0f32, 0usize);
301 for (c, v) in row.iter().enumerate() {
302 if *v > top.0 {
303 top = (*v, c);
304 }
305 }
306 best.push(top);
307 }
308 let mut order: Vec<usize> = (0..points).collect();
309 if opts.pre_k > 0 && points > opts.pre_k {
310 // Only the strongest positions are decoded at all, which is what the
311 // reference does before it thresholds.
312 order.sort_by(|a, b| best[*b].0.partial_cmp(&best[*a].0)
313 .unwrap_or(core::cmp::Ordering::Equal));
314 order.truncate(opts.pre_k);
315 }
316
317 let limit = SIDE as f32;
318 for p in order {
319 let (score, class) = best[p];
320 if score < opts.score_threshold {
321 continue;
322 }
323 // The four sides, each the mean of its own distribution.
324 let row = &bx.data[p * SIDES * BINS..(p + 1) * SIDES * BINS];
325 let mut d = [0.0f32; SIDES];
326 for (s, dist) in d.iter_mut().enumerate() {
327 let bins = &row[s * BINS..(s + 1) * BINS];
328 let top = bins.iter().copied().fold(f32::NEG_INFINITY, f32::max);
329 let mut sum = 0.0f32;
330 let mut acc = 0.0f32;
331 for (i, v) in bins.iter().enumerate() {
332 let e = (v - top).exp();
333 sum += e;
334 acc += e * i as f32;
335 }
336 *dist = if sum > 0.0 { acc / sum * stride as f32 } else { 0.0 };
337 }
338
339 // The position, which is half a sample short of the cell's centre.
340 let (gx, gy) = (p % side, p / side);
341 let cx = (gx * stride) as f32 + 0.5 * (stride as f32 - 1.0);
342 let cy = (gy * stride) as f32 + 0.5 * (stride as f32 - 1.0);
343 let x1 = (cx - d[0]).clamp(0.0, limit);
344 let y1 = (cy - d[1]).clamp(0.0, limit);
345 let x2 = (cx + d[2]).clamp(0.0, limit);
346 let y2 = (cy + d[3]).clamp(0.0, limit);
347 out.push(Object {
348 class,
349 score,
350 x: x1,
351 y: y1,
352 w: x2 - x1,
353 h: y2 - y1,
354 });
355 }
356 Ok(())
357}
358
359/// Overlap of two boxes, each `(x, y, w, h)` in whole pixels.
360fn iou(a: (i64, i64, i64, i64), b: (i64, i64, i64, i64)) -> f32 {
361 let x0 = a.0.max(b.0);
362 let y0 = a.1.max(b.1);
363 let x1 = (a.0 + a.2).min(b.0 + b.2);
364 let y1 = (a.1 + a.3).min(b.1 + b.3);
365 if x1 <= x0 || y1 <= y0 {
366 return 0.0;
367 }
368 let inter = ((x1 - x0) * (y1 - y0)) as f64;
369 let union = (a.2 * a.3) as f64 + (b.2 * b.3) as f64 - inter;
370 if union <= 0.0 {
371 return 0.0;
372 }
373 (inter / union) as f32
374}
375
376/// Greedy non-maximum suppression, strongest box first.
377///
378/// The suppression does not know about categories, which is the reference's
379/// behaviour and is right here: two boxes on the same animal, one calling it a
380/// dog and the other a cat, are one animal and not two, and keeping both would
381/// report the disagreement as a pair of findings.
382fn suppress(mut cand: Vec<Object>, opts: &Options) -> Vec<Object> {
383 if cand.len() <= 1 {
384 return cand;
385 }
386 cand.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(core::cmp::Ordering::Equal));
387 if opts.top_k > 0 && cand.len() > opts.top_k {
388 cand.truncate(opts.top_k);
389 }
390 let mut kept: Vec<Object> = Vec::new();
391 for o in cand {
392 let b = o.int_box();
393 if kept.iter().any(|k| iou(k.int_box(), b) > opts.nms_threshold) {
394 continue;
395 }
396 kept.push(o);
397 }
398 kept
399}
400
401#[cfg(test)]
402mod tests {
403 use super::*;
404
405 #[test]
406 fn the_animals_are_the_ones_the_names_say_they_are() -> Outcome<()> {
407 let want = ["bird", "cat", "dog", "horse", "sheep", "cow", "elephant",
408 "bear", "zebra", "giraffe"];
409 for (i, name) in ANIMALS.iter().zip(want.iter()) {
410 req!(NAMES[*i], *name);
411 }
412 req!(is_animal(15), true, "A cat is an animal.");
413 req!(is_animal(0), false, "A person is not one of the animals here.");
414 Ok(())
415 }
416
417 #[test]
418 fn a_distribution_decodes_to_its_mean() -> Outcome<()> {
419 // The coarsest head: a 13 by 13 grid at stride 32. One position carries a
420 // dog, and each of its four sides puts all its weight on bin 2, so every
421 // distance is 2 x 32 = 64 out from the position.
422 let side = 13;
423 let points = side * side;
424 let stride = SIDE / side;
425 let at = 5 * side + 7; // grid column 7, row 5
426
427 let mut c = vec![0.0f32; points * CATEGORIES];
428 c[at * CATEGORIES + 16] = 0.9; // dog
429 let cls = res!(Tensor::new(vec![1, points, CATEGORIES], c));
430
431 let mut b = vec![-30.0f32; points * SIDES * BINS];
432 for s in 0..SIDES {
433 b[at * SIDES * BINS + s * BINS + 2] = 30.0;
434 }
435 let bx = res!(Tensor::new(vec![1, points, SIDES * BINS], b));
436
437 let opts = Options { score_threshold: 0.5, ..Options::default() };
438 let mut out = Vec::new();
439 res!(level(points, &cls, &bx, &opts, &mut out));
440
441 req!(out.len(), 1, "One position was above the threshold.");
442 let o = out[0];
443 req!(o.name(), "dog");
444 req!(o.is_animal(), true);
445
446 // The position sits half a sample short of the cell's centre, and the box
447 // reaches 64 out from it on every side.
448 let cx = (7 * stride) as f32 + 0.5 * (stride as f32 - 1.0);
449 let cy = (5 * stride) as f32 + 0.5 * (stride as f32 - 1.0);
450 let want = 2.0 * stride as f32;
451 let left = (o.x - (cx - want)).abs() < 1e-3;
452 let top = (o.y - (cy - want)).abs() < 1e-3;
453 let wide = (o.w - 2.0 * want).abs() < 1e-3;
454 let tall = (o.h - 2.0 * want).abs() < 1e-3;
455 req!(left, true, "The left edge is at {}, wanted {}.", o.x, cx - want);
456 req!(top, true, "The top edge is at {}, wanted {}.", o.y, cy - want);
457 req!(wide, true, "The box is {} wide, wanted {}.", o.w, 2.0 * want);
458 req!(tall, true, "The box is {} tall, wanted {}.", o.h, 2.0 * want);
459
460 // Taking the cell's centre instead of the position would move it by half
461 // a sample, which is the fault this arithmetic is easiest to get wrong in.
462 let centred = (7 * stride) as f32 + 0.5 * stride as f32;
463 let apart = (centred - cx).abs() > 1e-3;
464 req!(apart, true, "The position and the cell centre are not distinguishable.");
465 Ok(())
466 }
467
468 #[test]
469 fn overlap_is_measured_on_whole_pixels() -> Outcome<()> {
470 req!(iou((0, 0, 10, 10), (0, 0, 10, 10)), 1.0f32);
471 req!(iou((0, 0, 10, 10), (20, 20, 10, 10)), 0.0f32);
472 let half = iou((0, 0, 10, 10), (5, 0, 10, 10));
473 let third = (half - 1.0 / 3.0).abs() < 1e-6;
474 req!(third, true);
475 Ok(())
476 }
477
478 #[test]
479 fn the_suppression_does_not_care_what_a_box_is_called() -> Outcome<()> {
480 // The same animal, called two things. One finding, not two.
481 let a = Object { class: 16, score: 0.6, x: 10.0, y: 10.0, w: 50.0, h: 50.0 };
482 let b = Object { class: 15, score: 0.5, x: 11.0, y: 11.0, w: 50.0, h: 50.0 };
483 let kept = suppress(vec![a, b], &Options::default());
484 req!(kept.len(), 1, "Two names for one animal came back as two animals.");
485 req!(kept[0].name(), "dog");
486 Ok(())
487 }
488}