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 | |
| 3 | use crate::face::align; |
| 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 | /// Length of the vector the embedder answers. |
| 12 | pub 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)] |
| 19 | pub struct Embedding { |
| 20 | /// The unit vector. |
| 21 | pub v: [f32; DIM], |
| 22 | } |
| 23 | |
| 24 | impl 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. |
| 50 | pub 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)] |
| 60 | pub struct Embedder { |
| 61 | /// The prepared graph. |
| 62 | graph: Graph, |
| 63 | } |
| 64 | |
| 65 | impl 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)] |
| 120 | mod 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 | } |