Oregami
Repositories/oxedyne/fe2o3

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

4.3 KiB, 1 run

created by r1870400018:19751, 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 embedding: an aligned crop in, a unit vector out.
2
3use crate::face::align;
4use crate::face::Image;
5use crate::graph::Graph;
6use crate::kern::Cpu;
7use crate::tensor::Tensor;
8
9use oxedyne_fe2o3_core::prelude::*;
10
11/// Length of the vector the embedder answers.
12pub const DIM: usize = 128;
13
14/// A face, as a point on the unit sphere.
15///
16/// Two of these are compared with [`cosine`], which is a dot product because
17/// the vector is already normalised.
18#[derive(Clone, Copy, Debug, PartialEq)]
19pub struct Embedding {
20 /// The unit vector.
21 pub v: [f32; DIM],
22}
23
24impl Embedding {
25 /// Normalises a raw network output onto the unit sphere.
26 pub fn from_raw(raw: &[f32]) -> Outcome<Self> {
27 if raw.len() != DIM {
28 return Err(err!(
29 "An embedding of {} values was expected, found {}.", DIM, raw.len();
30 Invalid, Input, Mismatch));
31 }
32 let norm = raw.iter().map(|v| (*v as f64) * (*v as f64)).sum::<f64>().sqrt();
33 if norm <= 0.0 {
34 return Err(err!("An embedding of zero length cannot be normalised.";
35 Invalid, Input, Range));
36 }
37 let mut v = [0.0f32; DIM];
38 for (o, r) in v.iter_mut().zip(raw.iter()) {
39 *o = (*r as f64 / norm) as f32;
40 }
41 Ok(Self { v })
42 }
43}
44
45/// Cosine similarity of two embeddings, in `[-1, 1]`.
46///
47/// The reference implementation calls two faces the same person above `0.363`.
48/// A clustering threshold should sit higher than a verification one, because a
49/// cluster that splits is easy to mend and a cluster that merges is not.
50pub fn cosine(a: &Embedding, b: &Embedding) -> f32 {
51 let mut s = 0.0f64;
52 for i in 0..DIM {
53 s += a.v[i] as f64 * b.v[i] as f64;
54 }
55 s as f32
56}
57
58/// A loaded face embedder.
59#[derive(Clone, Debug)]
60pub struct Embedder {
61 /// The prepared graph.
62 graph: Graph,
63}
64
65impl Embedder {
66 /// Loads an embedder from the bytes of an `.onnx` file.
67 pub fn load(onnx: &[u8]) -> Outcome<Self> {
68 let graph = res!(Graph::load(onnx));
69 if graph.outputs.len() != 1 {
70 return Err(err!(
71 "An embedder wants one output, the graph declares {}.", graph.outputs.len();
72 Invalid, Input, Mismatch));
73 }
74 Ok(Self { graph })
75 }
76
77 /// The prepared graph, for callers that want to time or inspect it.
78 pub fn graph(&self) -> &Graph {
79 &self.graph
80 }
81
82 /// Turns an aligned crop into the input tensor the embedder was exported
83 /// against, which is red-green-blue and unnormalised -- the subtraction and
84 /// the scaling are the first two operators of the graph itself.
85 pub fn input_tensor(crop: &[u8]) -> Outcome<Tensor> {
86 let want = align::CROP * align::CROP * 3;
87 if crop.len() != want {
88 return Err(err!(
89 "An aligned crop of {} bytes was expected, found {}.", want, crop.len();
90 Invalid, Input, Mismatch));
91 }
92 let data = crop.iter().map(|v| *v as f32).collect::<Vec<_>>();
93 Tensor::new(vec![1, align::CROP, align::CROP, 3], data)
94 }
95
96 /// Embeds a crop that has already been warped onto the template.
97 pub fn embed_aligned(&self, cpu: Cpu, crop: &[u8]) -> Outcome<Embedding> {
98 let input = res!(Self::input_tensor(crop));
99 let outs = res!(self.graph.run(cpu, input));
100 let raw = some!(outs.first(), "The embedder answered no output.");
101 Embedding::from_raw(&raw.data)
102 }
103
104 /// Embeds a face out of a photograph, given its five landmarks in that
105 /// photograph's own coordinates.
106 pub fn embed(&self, cpu: Cpu, img: &Image<'_>, landmarks: &[(f32, f32); 5])
107 -> Outcome<Embedding>
108 {
109 if img.channels != 3 {
110 return Err(err!(
111 "The embedder wants three channels, the image has {}.", img.channels;
112 Invalid, Input, Mismatch));
113 }
114 let crop = res!(align::align_crop(img, landmarks));
115 self.embed_aligned(cpu, &crop)
116 }
117}
118
119#[cfg(test)]
120mod tests {
121 use super::*;
122
123 #[test]
124 fn a_vector_normalises_and_matches_itself() -> Outcome<()> {
125 let raw = (0..DIM).map(|i| (i as f32) - 63.5).collect::<Vec<_>>();
126 let e = res!(Embedding::from_raw(&raw));
127 let n = e.v.iter().map(|v| (*v as f64) * (*v as f64)).sum::<f64>();
128 req!(((n - 1.0).abs() < 1e-6), true);
129 req!(((cosine(&e, &e) - 1.0).abs() < 1e-5), true);
130 Ok(())
131 }
132
133 #[test]
134 fn an_opposite_vector_scores_minus_one() -> Outcome<()> {
135 let raw = (0..DIM).map(|i| (i as f32) - 63.5).collect::<Vec<_>>();
136 let neg = raw.iter().map(|v| -*v).collect::<Vec<_>>();
137 let a = res!(Embedding::from_raw(&raw));
138 let b = res!(Embedding::from_raw(&neg));
139 req!(((cosine(&a, &b) + 1.0).abs() < 1e-5), true);
140 Ok(())
141 }
142}