Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/tests/detector.rs

9.4 KiB, 7 runs

created by r1870400018:21035, 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//! A category detector's graph, held sample for sample against another engine.
2//!
3//! The face models are checked against `tract` on a synthetic input. This one is
4//! checked against OpenCV's own DNN engine on a real photograph, because it
5//! exercises the operators `tract` never made this crate need: a pooled window
6//! with padding, a bilinear resample under PyTorch's half-pixel rule, a channel
7//! split, a concatenation and a channel shuffle. Each of those is written
8//! against a channels-last layout while the model names channels-first axes, so
9//! agreeing with an independent engine is the only thing that says they are
10//! right.
11//!
12//! The fixture is made by `oracle.py` in the scoping directory, which runs the
13//! same `.onnx` through `cv2.dnn` and writes each tensor as raw little-endian
14//! `f32` beside a `.shape` file. Point `FE2O3_INFER_DETECTOR` at the model and
15//! `FE2O3_INFER_DETECTOR_REF` at the directory of tensors; without both, the
16//! test says it was skipped and passes.
17
18use std::env;
19use std::fs;
20use std::path::{Path, PathBuf};
21
22use oxedyne_fe2o3_core::prelude::*;
23use oxedyne_fe2o3_infer::graph::Graph;
24use oxedyne_fe2o3_infer::kern::Cpu;
25use oxedyne_fe2o3_infer::object::{self, Options};
26use oxedyne_fe2o3_infer::tensor::Tensor;
27
28/// The model and the reference directory, or `None` so the test can skip.
29fn fixture() -> Option<(PathBuf, PathBuf)> {
30 let model = env::var("FE2O3_INFER_DETECTOR").ok()?;
31 let refs = env::var("FE2O3_INFER_DETECTOR_REF").ok()?;
32 if model.is_empty() || refs.is_empty() {
33 return None;
34 }
35 Some((PathBuf::from(model), PathBuf::from(refs)))
36}
37
38/// Reads one reference tensor: raw `f32` beside a file naming its shape.
39fn reference(dir: &Path, name: &str) -> Outcome<Tensor> {
40 let shape = res!(fs::read_to_string(dir.join(format!("{}.shape", name))));
41 let mut dims = Vec::new();
42 for word in shape.split_whitespace() {
43 dims.push(res!(word.parse::<usize>()));
44 }
45 let raw = res!(fs::read(dir.join(format!("{}.f32", name))));
46 if raw.len() % 4 != 0 {
47 return Err(err!("The reference {} is {} bytes, not a run of f32.", name, raw.len();
48 Invalid, Input, Mismatch));
49 }
50 let mut data = Vec::with_capacity(raw.len() / 4);
51 for chunk in raw.chunks_exact(4) {
52 let mut b = [0u8; 4];
53 b.copy_from_slice(chunk);
54 data.push(f32::from_le_bytes(b));
55 }
56 Tensor::new(dims, data)
57}
58
59/// The outputs the reference holds. Only used to check that every one of the
60/// graph's own outputs is covered: the two engines list them in different
61/// orders, so they are paired by name and never by position.
62fn output_names(dir: &Path) -> Outcome<Vec<String>> {
63 let meta = res!(fs::read_to_string(dir.join("meta.txt")));
64 for line in meta.lines() {
65 if let Some(rest) = line.strip_prefix("outputs ") {
66 return Ok(rest.split_whitespace().map(|s| s.to_string()).collect());
67 }
68 }
69 Err(err!("The reference directory names no outputs."; Invalid, Input, Missing))
70}
71
72#[test]
73fn the_detector_answers_what_another_engine_answered() -> Outcome<()> {
74 let (model, root) = match fixture() {
75 Some(f) => f,
76 None => {
77 println!("skipped: set FE2O3_INFER_DETECTOR and FE2O3_INFER_DETECTOR_REF");
78 return Ok(());
79 },
80 };
81
82 let bytes = res!(fs::read(&model));
83 let g = res!(Graph::load(&bytes));
84
85 // A directory of tensors, or a directory of directories of them.
86 let mut cases = Vec::new();
87 if root.join("meta.txt").exists() {
88 cases.push(root.clone());
89 } else {
90 for entry in res!(fs::read_dir(&root)) {
91 let path = res!(entry).path();
92 if path.join("meta.txt").exists() {
93 cases.push(path);
94 }
95 }
96 cases.sort();
97 }
98 if cases.is_empty() {
99 return Err(err!("No reference tensors under {:?}.", root; Invalid, Input, Missing));
100 }
101
102 let mut over = 0.0f32;
103 let mut over_at = String::new();
104 for refs in &cases {
105 let (worst, at) = res!(one_case(&g, refs));
106 let name = refs.file_name().map(|s| s.to_string_lossy().to_string())
107 .unwrap_or_else(|| fmt!("{:?}", refs));
108 println!("{:8} worst difference {:e}, at {}", name, worst, at);
109 if worst > over {
110 over = worst;
111 over_at = fmt!("{}: {}", name, at);
112 }
113 }
114
115 println!("{} photographs, worst of all {:e}, at {}", cases.len(), over, over_at);
116 // Two engines summing the same convolution in different orders differ in the
117 // last bits and no more. A wrong operator is not a rounding difference: the
118 // faults this test exists to catch move a value by whole units, and dropping
119 // the channel shuffle or reading the resize by the other coordinate rule was
120 // measured at 16.9 and 2.7 against the 5e-6 two right answers differ by.
121 let close = over < 2.0e-3;
122 req!(close, true, "The detector differs from the reference by {}, at {}.", over, over_at);
123 Ok(())
124}
125
126/// Runs one photograph and answers the worst disagreement and where it was.
127fn one_case(g: &Graph, refs: &Path) -> Outcome<(f32, String)> {
128 // The reference input is `[1, 3, 416, 416]`, channels first, as the model
129 // declares it; this crate wants it channels last.
130 let blob = res!(reference(refs, "input"));
131 if blob.dims.len() != 4 {
132 return Err(err!("The reference input has shape {:?}.", blob.dims; Invalid, Input));
133 }
134 let (n, c, h, w) = (blob.dims[0], blob.dims[1], blob.dims[2], blob.dims[3]);
135 let mut nhwc = vec![0.0f32; blob.len()];
136 for ci in 0..c {
137 for p in 0..h * w {
138 nhwc[p * c + ci] = blob.data[ci * h * w + p];
139 }
140 }
141 let input = res!(Tensor::new(vec![n, h, w, c], nhwc));
142
143 let got = res!(g.run(Cpu::detect(), input));
144 let held = res!(output_names(refs));
145 if got.len() != held.len() {
146 return Err(err!(
147 "The graph answered {} outputs and the reference holds {}.", got.len(), held.len();
148 Invalid, Mismatch));
149 }
150 // Each output is paired with the reference of the same name. The two engines
151 // declare the six in different orders, and pairing by position silently
152 // compares a box prediction with a class score.
153 let names = g.outputs.iter()
154 .map(|i| g.names[*i].clone())
155 .collect::<Vec<_>>();
156 for name in &names {
157 if !held.contains(name) {
158 return Err(err!(
159 "The graph answers an output {} that the reference does not hold.", name;
160 Invalid, Mismatch));
161 }
162 }
163
164 let mut worst = 0.0f32;
165 let mut worst_at = String::new();
166 for (name, mine) in names.iter().zip(got.iter()) {
167 let want = res!(reference(refs, name));
168 if mine.len() != want.len() {
169 return Err(err!(
170 "Output {} has {} values against the reference's {} (shapes {:?} and {:?}).",
171 name, mine.len(), want.len(), mine.dims, want.dims;
172 Invalid, Mismatch));
173 }
174 for (i, (a, b)) in mine.data.iter().zip(want.data.iter()).enumerate() {
175 let d = (a - b).abs();
176 if d > worst {
177 worst = d;
178 worst_at = fmt!("{} at {} ({} against {})", name, i, a, b);
179 }
180 }
181 }
182
183 // The decode is held to the reference's own boxes, which is a separate claim
184 // from the network agreeing: the distribution integral, the anchor positions
185 // and the suppression are all this crate's and none of them is exercised by
186 // comparing tensors.
187 res!(check_detections(refs, &got));
188
189 Ok((worst, worst_at))
190}
191
192/// One line of the reference's decoded output.
193struct Ref {
194 /// Index into the category names.
195 class: usize,
196 /// Confidence.
197 score: f32,
198 /// Box, `(x, y, w, h)` in canvas pixels.
199 rect: (f32, f32, f32, f32),
200}
201
202/// Compares this crate's decode with the reference's, box for box.
203fn check_detections(refs: &Path, out: &[Tensor]) -> Outcome<()> {
204 let path = refs.join("detections.txt");
205 if !path.exists() {
206 return Ok(());
207 }
208 let text = res!(fs::read_to_string(&path));
209 let mut want = Vec::new();
210 for line in text.lines() {
211 let f = line.split('\t').collect::<Vec<_>>();
212 if f.len() != 6 {
213 return Err(err!("A reference detection has {} fields.", f.len();
214 Invalid, Input, Mismatch));
215 }
216 want.push(Ref {
217 class: res!(f[0].parse::<usize>()),
218 score: res!(f[1].parse::<f32>()),
219 rect: (
220 res!(f[2].parse::<f32>()),
221 res!(f[3].parse::<f32>()),
222 res!(f[4].parse::<f32>()),
223 res!(f[5].parse::<f32>()),
224 ),
225 });
226 }
227
228 let got = res!(object::decode(out, &Options::default()));
229 if got.len() != want.len() {
230 let mine = got.iter()
231 .map(|o| fmt!("{} {:.2}", o.name(), o.score))
232 .collect::<Vec<_>>().join(", ");
233 let theirs = want.iter()
234 .map(|r| fmt!("{} {:.2}", object::NAMES[r.class], r.score))
235 .collect::<Vec<_>>().join(", ");
236 return Err(err!(
237 "The decode found {} objects and the reference {}. Mine: [{}]. Theirs: [{}].",
238 got.len(), want.len(), mine, theirs;
239 Invalid, Mismatch));
240 }
241
242 // Both are strongest first, so they line up.
243 for (i, (mine, theirs)) in got.iter().zip(want.iter()).enumerate() {
244 if mine.class != theirs.class {
245 return Err(err!(
246 "Detection {} is a {} here and a {} in the reference.",
247 i, mine.name(), object::NAMES[theirs.class];
248 Invalid, Mismatch));
249 }
250 let ds = (mine.score - theirs.score).abs();
251 if ds > 1.0e-5 {
252 return Err(err!(
253 "Detection {} scores {} here and {} in the reference.",
254 i, mine.score, theirs.score;
255 Invalid, Mismatch));
256 }
257 let (x, y, w, h) = theirs.rect;
258 let off = (mine.x - x).abs()
259 .max((mine.y - y).abs())
260 .max((mine.w - w).abs())
261 .max((mine.h - h).abs());
262 // A quarter of a pixel. The faults worth catching -- the anchor half a
263 // sample out, the distribution read the wrong way round -- move an edge
264 // by whole pixels at the finest head and by tens at the coarsest.
265 if off > 0.25 {
266 return Err(err!(
267 "Detection {} ({}) is at ({}, {}, {}, {}) here and ({}, {}, {}, {}) in the \
268 reference, {} pixels apart.",
269 i, mine.name(), mine.x, mine.y, mine.w, mine.h, x, y, w, h, off;
270 Invalid, Mismatch));
271 }
272 }
273 Ok(())
274}