Oregami
Repositories/oxedyne/fe2o3

oxedyne/fe2o3/fe2o3_infer/tests/models.rs

7.3 KiB, 1 run

created by r1870400018:19771, 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//! Correctness against a reference implementation, and against tract's recorded
2//! answer on a fixed input.
3//!
4//! Two of these tests want the model files. They are not in the repository --
5//! weights never are -- so point `FE2O3_INFER_MODELS` at a directory holding
6//! `face_detection_yunet_2023mar.onnx` and
7//! `face_recognition_sface_2021dec.onnx`, and the tests will run. Without it
8//! they report that they were skipped and pass, which is the only sane
9//! behaviour for a test whose fixture is thirty-nine megabytes.
10//!
11//! The recorded vectors below came from tract 0.23.4 running the same ONNX file
12//! on the same input, taken once and frozen. A test that only compared this
13//! crate against itself would prove nothing; these numbers are what an
14//! independent implementation answered.
15
16use std::env;
17use std::fs;
18
19use oxedyne_fe2o3_core::prelude::*;
20use oxedyne_fe2o3_infer::face::{detect::Detector, embed::Embedder, Image};
21use oxedyne_fe2o3_infer::kern::Cpu;
22
23/// A reproducible byte generator, so the fixed input needs no fixture file.
24fn pixels(n: usize, seed: u64) -> Vec<u8> {
25 let mut s = seed;
26 let mut v = Vec::with_capacity(n);
27 for _ in 0..n {
28 s = s.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
29 v.push((s >> 33) as u8);
30 }
31 v
32}
33
34/// Finds the model directory, or answers `None` so the test can skip.
35fn models() -> Option<String> {
36 match env::var("FE2O3_INFER_MODELS") {
37 Ok(d) if !d.is_empty() => Some(d),
38 _ => None,
39 }
40}
41
42// SFace, recorded from tract 0.23.4 on the fixed input of `pixels(112*112*3, 999)`.
43const SFACE_EXPECTED: [f32; 128] = [
44 -2.269915e-1, -1.415737e-2, -2.509517e-1, 3.997844e-1,
45 -2.476652e-1, 1.763874e-1, 1.363716e-1, -2.750306e-1,
46 -8.319645e-1, -9.862685e-2, -1.967272e-1, 2.877449e-1,
47 4.880038e-1, -4.574382e-1, 9.873680e-1, -1.919422e-1,
48 -6.378446e-1, -5.909452e-1, -1.601325e-1, 3.697520e-1,
49 8.097876e-2, -6.924052e-1, -4.088828e-1, -1.624882e-3,
50 2.270993e-1, 1.198874e-1, 4.283409e-1, 2.593386e-1,
51 -3.672851e-1, 9.240717e-3, -8.834651e-2, -4.717094e-1,
52 3.875837e-1, 7.771777e-1, 2.067817e-1, -1.768308e-1,
53 -3.115587e-1, -4.025498e-2, 3.658660e-1, 1.752793e-2,
54 1.606840e-1, -5.139388e-1, 9.671890e-1, 3.963410e-1,
55 -3.993112e-1, -6.236808e-1, 2.314502e-1, 2.565281e-1,
56 3.500144e-1, -8.477663e-2, 2.539647e-2, -2.362562e-1,
57 4.709096e-1, -3.142870e-1, 6.237337e-1, 3.424331e-1,
58 -4.658711e-1, -5.021946e-1, -1.177402e-1, -5.383276e-1,
59 -4.104891e-2, 3.315228e-1, -5.957409e-1, -7.728708e-1,
60 3.792651e-1, -5.855272e-2, -2.673874e-1, 5.746964e-1,
61 3.235634e-1, 2.786989e-1, -6.076752e-1, 3.722823e-1,
62 -6.298876e-1, -2.459088e-1, -3.353854e-1, -2.661195e-1,
63 2.093293e-1, 2.202750e-2, -7.259557e-2, 2.664874e-1,
64 -2.680961e-1, 1.110036e-1, 3.179148e-1, -7.449403e-2,
65 -4.279102e-1, 1.783608e-1, -1.274272e-2, -1.293271e-1,
66 5.955364e-1, 1.073973e-1, 3.554508e-1, -2.166391e-1,
67 -5.645101e-2, 6.639143e-1, 1.299650e-1, 6.170660e-1,
68 7.098893e-1, 6.794992e-2, 6.472573e-2, 2.854365e-1,
69 7.188817e-1, -3.434052e-1, -3.648337e-1, -1.419082e-1,
70 -3.782669e-1, -3.707471e-1, -1.573869e-1, 1.157296e-1,
71 -7.234098e-2, 5.041344e-1, 6.299373e-1, 9.830029e-1,
72 9.201707e-1, -6.767877e-2, 1.296642e-1, 1.464763e-1,
73 -2.185782e-1, 1.255877e-2, -8.435621e-1, 6.842389e-1,
74 -1.327648e-2, 3.854927e-1, -4.899058e-1, 1.003857e-1,
75 1.415506e-2, 3.207604e-1, 1.415001e-1, 9.006548e-2,
76];
77
78// YuNet, recorded from tract 0.23.4 on the fixed input of `pixels(640*640*3, 12345)`.
79// Each row is one declared output: mean, maximum, and four sampled values.
80const YUNET_EXPECTED: [[f32; 6]; 12] = [
81 [3.225795e-1, 5.002567e-1, 4.502176e-1, 4.432113e-1, 4.305283e-1, 4.268720e-1],
82 [4.158589e-1, 5.052462e-1, 4.737059e-1, 4.389420e-1, 4.639358e-1, 4.423938e-1],
83 [6.329900e-1, 7.860411e-1, 5.499988e-1, 5.024071e-1, 4.832734e-1, 5.247698e-1],
84 [1.384022e-3, 3.692466e-2, 2.051294e-4, 9.417534e-6, 6.592274e-4, 5.739927e-5],
85 [3.710113e-4, 1.439953e-2, 4.290044e-4, 2.458692e-5, 2.228022e-4, 7.522106e-5],
86 [7.854924e-4, 6.979972e-3, 1.356632e-3, 7.806122e-4, 1.905799e-3, 4.732013e-4],
87 [2.135477e-1, 1.080065e0, 5.861937e-1, 2.982946e-1, 5.664475e-1, 4.919586e-1],
88 [-1.063826e-1, 8.579210e-1, 5.123827e-1, 4.617671e-1, 5.297288e-1, 5.952517e-1],
89 [9.833903e-1, 2.517147e0, 1.240810e0, 1.038598e0, 1.207002e0, 1.187358e0],
90 [6.133286e-1, 1.641716e0, -2.000998e-1, -3.537228e-1, -2.693225e-1, -2.384279e-1],
91 [4.718952e-1, 9.449253e-1, 2.159463e-1, 1.881350e-1, 2.207145e-1, 2.293447e-1],
92 [-5.286866e-1, 5.250635e0, -6.274422e-1, -3.532023e-1, -5.473853e-1, -4.717106e-1],
93];
94
95#[test]
96fn the_embedder_answers_what_tract_answered() -> Outcome<()> {
97 let dir = match models() {
98 Some(d) => d,
99 None => {
100 println!("skipped: set FE2O3_INFER_MODELS to the directory holding the models");
101 return Ok(());
102 },
103 };
104 let bytes = res!(fs::read(format!("{}/face_recognition_sface_2021dec.onnx", dir)));
105 let emb = res!(Embedder::load(&bytes));
106 let crop = pixels(112 * 112 * 3, 999);
107 let input = res!(Embedder::input_tensor(&crop));
108 let outs = res!(emb.graph().run(Cpu::detect(), input));
109 let got = &outs[0].data;
110 req!(got.len(), 128);
111 let mut worst = 0.0f32;
112 for (g, w) in got.iter().zip(SFACE_EXPECTED.iter()) {
113 worst = worst.max((g - w).abs());
114 }
115 if worst > 2e-4 {
116 return Err(err!(
117 "The embedding differs from tract's by {:.3e}, which is more than rounding.", worst;
118 Invalid, Mismatch));
119 }
120 println!("largest difference from tract: {:.3e}", worst);
121 Ok(())
122}
123
124#[test]
125fn the_detector_answers_what_tract_answered() -> Outcome<()> {
126 let dir = match models() {
127 Some(d) => d,
128 None => {
129 println!("skipped: set FE2O3_INFER_MODELS to the directory holding the models");
130 return Ok(());
131 },
132 };
133 let bytes = res!(fs::read(format!("{}/face_detection_yunet_2023mar.onnx", dir)));
134 let det = res!(Detector::load(&bytes));
135 let px = pixels(640 * 640 * 3, 12345);
136 let img = res!(Image::new(&px, 640, 640, 3));
137 let input = res!(Detector::input_tensor(&img));
138 let outs = res!(det.graph().run(Cpu::detect(), input));
139 req!(outs.len(), 12);
140 let mut worst = 0.0f32;
141 for (i, o) in outs.iter().enumerate() {
142 let v = &o.data;
143 let mean = (v.iter().map(|x| *x as f64).sum::<f64>() / v.len() as f64) as f32;
144 let max = v.iter().fold(f32::MIN, |a, b| a.max(*b));
145 let s: Vec<f32> = (1..5).map(|j| v[j * v.len() / 5]).collect();
146 let got = [mean, max, s[0], s[1], s[2], s[3]];
147 for (g, w) in got.iter().zip(YUNET_EXPECTED[i].iter()) {
148 worst = worst.max((g - w).abs());
149 }
150 }
151 if worst > 1e-4 {
152 return Err(err!(
153 "A head differs from tract's by {:.3e}, which is more than rounding.", worst;
154 Invalid, Mismatch));
155 }
156 println!("largest difference from tract: {:.3e}", worst);
157 Ok(())
158}
159
160#[test]
161fn both_code_paths_reach_the_same_embedding() -> Outcome<()> {
162 let dir = match models() {
163 Some(d) => d,
164 None => {
165 println!("skipped: set FE2O3_INFER_MODELS to the directory holding the models");
166 return Ok(());
167 },
168 };
169 let bytes = res!(fs::read(format!("{}/face_recognition_sface_2021dec.onnx", dir)));
170 let emb = res!(Embedder::load(&bytes));
171 let crop = pixels(112 * 112 * 3, 4242);
172 let fast = res!(emb.embed_aligned(Cpu::detect(), &crop));
173 let slow = res!(emb.embed_aligned(Cpu::Baseline, &crop));
174 let mut worst = 0.0f32;
175 for i in 0..128 {
176 worst = worst.max((fast.v[i] - slow.v[i]).abs());
177 }
178 if worst > 1e-5 {
179 return Err(err!(
180 "The dispatched path and the baseline path differ by {:.3e}.", worst;
181 Invalid, Mismatch));
182 }
183 println!("largest difference between the two paths: {:.3e}", worst);
184 Ok(())
185}