Skip to main content

smriti/ml/
face_embedder.rs

1//! ArcFace Face Embedding
2//!
3//! Generates 512-dimensional embeddings for face recognition.
4//! Uses the ArcFace-R100 model via ONNX Runtime, with optional
5//! offloading to a remote GPU bridge.
6
7use std::path::PathBuf;
8
9use image::RgbImage;
10use ndarray::Array1;
11use ort::session::Session;
12use ort::value::TensorRef;
13
14use super::remote_embedder::RemoteEmbedder;
15use super::OnnxRuntime;
16
17/// 512-dimensional face embedding
18#[derive(Debug, Clone)]
19pub struct FaceEmbedding {
20    /// The embedding vector
21    pub vector: Array1<f32>,
22}
23
24impl FaceEmbedding {
25    /// Create from a vector
26    pub fn new(vector: Array1<f32>) -> Self {
27        Self { vector }
28    }
29
30    /// Calculate cosine similarity with another embedding
31    pub fn cosine_similarity(&self, other: &FaceEmbedding) -> f32 {
32        let dot: f32 = self
33            .vector
34            .iter()
35            .zip(other.vector.iter())
36            .map(|(a, b)| a * b)
37            .sum();
38
39        let norm1: f32 = self.vector.iter().map(|x| x * x).sum::<f32>().sqrt();
40        let norm2: f32 = other.vector.iter().map(|x| x * x).sum::<f32>().sqrt();
41
42        if norm1 > 0.0 && norm2 > 0.0 {
43            dot / (norm1 * norm2)
44        } else {
45            0.0
46        }
47    }
48
49    /// Convert to bytes for database storage (little-endian f32 values)
50    pub fn to_bytes(&self) -> Vec<u8> {
51        self.vector.iter().flat_map(|f| f.to_le_bytes()).collect()
52    }
53
54    /// Create from bytes (expects 512 * 4 = 2048 bytes, little-endian f32)
55    pub fn from_bytes(bytes: &[u8]) -> Option<Self> {
56        if bytes.len() != 512 * 4 {
57            return None;
58        }
59
60        let vector: Array1<f32> = Array1::from_iter(
61            bytes
62                .chunks_exact(4)
63                .map(|chunk| f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]])),
64        );
65
66        Some(Self { vector })
67    }
68}
69
70/// Configuration for creating a FaceEmbedder.
71#[derive(Debug, Clone)]
72pub struct EmbedderConfig {
73    pub model_path: PathBuf,
74    pub gpu_bridge_url: Option<String>,
75    /// Name of the model the desktop is configured to use (e.g.
76    /// `adaface_ir101_webface12m.onnx`). Passed to the bridge so it
77    /// can refuse if its loaded model is different — mismatched models
78    /// produce embeddings in different metric spaces and would corrupt
79    /// clustering.
80    pub expected_model: String,
81    pub intra_threads: usize,
82}
83
84/// Dispatch enum — either local ONNX or remote GPU bridge.
85///
86/// The `Remote` variant always carries a `LocalEmbedder` alongside it.
87/// When the remote bridge fails for a batch, we fall back to local for
88/// just that batch instead of dropping the faces silently. The remote's
89/// own 3-strikes / unhealthy check still applies; this is the per-batch
90/// safety net while it's nominally healthy.
91pub enum FaceEmbedder {
92    Local(LocalEmbedder),
93    Remote {
94        remote: RemoteEmbedder,
95        local_fallback: Option<LocalEmbedder>,
96    },
97}
98
99impl FaceEmbedder {
100    /// Build the configured embedder. Always constructs a local
101    /// instance — it's required as a fallback even when the bridge is
102    /// healthy. If a remote bridge URL is configured AND the bridge
103    /// passes its initial health check (including model match), the
104    /// returned embedder routes batches through the remote first.
105    pub fn from_config(rt: &OnnxRuntime, cfg: &EmbedderConfig) -> Result<Self, ort::Error> {
106        if let Some(ref url) = cfg.gpu_bridge_url {
107            let remote = RemoteEmbedder::new(url.clone(), cfg.expected_model.clone());
108            if remote.is_healthy() {
109                let local_fallback = if cfg.model_path.exists() {
110                    Some(LocalEmbedder::new_with_threads(
111                        rt,
112                        &cfg.model_path,
113                        cfg.intra_threads,
114                    )?)
115                } else {
116                    tracing::warn!(
117                        "Local face embedder model {} missing; using remote bridge without local fallback",
118                        cfg.model_path.display()
119                    );
120                    None
121                };
122                tracing::info!("Face embedding routed to remote GPU bridge at {}", url);
123                return Ok(FaceEmbedder::Remote {
124                    remote,
125                    local_fallback,
126                });
127            }
128            tracing::info!("Remote GPU bridge unavailable; using local embedder");
129        }
130        let local = LocalEmbedder::new_with_threads(rt, &cfg.model_path, cfg.intra_threads)?;
131        Ok(FaceEmbedder::Local(local))
132    }
133
134    /// Generate embedding for an aligned face image (112x112).
135    pub fn embed(&mut self, aligned_face: &RgbImage) -> Option<FaceEmbedding> {
136        self.embed_batch(std::slice::from_ref(aligned_face))
137            .into_iter()
138            .next()
139            .flatten()
140    }
141
142    /// Generate embeddings for a batch of aligned face images.
143    /// Returns one Option per input; None only if both remote and
144    /// local-fallback produced no embedding for that input.
145    pub fn embed_batch(&mut self, faces: &[RgbImage]) -> Vec<Option<FaceEmbedding>> {
146        match self {
147            FaceEmbedder::Local(e) => e.embed_batch(faces),
148            FaceEmbedder::Remote {
149                remote,
150                local_fallback,
151            } => {
152                let mut out = remote.embed_batch(faces);
153                if out.len() != faces.len() {
154                    // Defensive: should never happen, but if it does,
155                    // recompute fully on local when available.
156                    return local_fallback
157                        .as_mut()
158                        .map_or_else(|| vec![None; faces.len()], |local| local.embed_batch(faces));
159                }
160                let missing: Vec<usize> = out
161                    .iter()
162                    .enumerate()
163                    .filter_map(|(i, r)| if r.is_none() { Some(i) } else { None })
164                    .collect();
165                if missing.is_empty() {
166                    return out;
167                }
168                if let Some(local) = local_fallback.as_mut() {
169                    let crops: Vec<RgbImage> = missing.iter().map(|&i| faces[i].clone()).collect();
170                    let fb = local.embed_batch(&crops);
171                    for (k, &orig_i) in missing.iter().enumerate() {
172                        if let Some(emb) = fb.get(k).cloned().flatten() {
173                            out[orig_i] = Some(emb);
174                        }
175                    }
176                }
177                out
178            }
179        }
180    }
181}
182
183/// Local ArcFace Face Embedder using ONNX Runtime.
184pub struct LocalEmbedder {
185    session: Session,
186}
187
188impl LocalEmbedder {
189    /// Load the ArcFace model with a specific thread count per session.
190    pub fn new_with_threads<P: AsRef<std::path::Path>>(
191        runtime: &OnnxRuntime,
192        model_path: P,
193        intra_threads: usize,
194    ) -> ort::Result<Self> {
195        let session = runtime.load_model_with_threads(model_path, intra_threads)?;
196
197        Ok(Self { session })
198    }
199
200    /// Generate embedding for an aligned face image (112x112)
201    pub fn embed(&mut self, aligned_face: &RgbImage) -> Option<FaceEmbedding> {
202        self.embed_batch(std::slice::from_ref(aligned_face))
203            .into_iter()
204            .next()
205            .flatten()
206    }
207
208    /// Generate embeddings for a batch of aligned face images.
209    /// Returns one Option per input; None if that input was invalid.
210    pub fn embed_batch(&mut self, faces: &[RgbImage]) -> Vec<Option<FaceEmbedding>> {
211        let n = faces.len();
212        if n == 0 {
213            return Vec::new();
214        }
215
216        // Validate dimensions and preprocess each face into a flat [3*112*112] row.
217        let mut valid_indices: Vec<usize> = Vec::with_capacity(n);
218        let mut batch_data: Vec<f32> = Vec::with_capacity(n * 3 * 112 * 112);
219        for (i, face) in faces.iter().enumerate() {
220            if face.width() != 112 || face.height() != 112 {
221                continue;
222            }
223            valid_indices.push(i);
224            let row = self.preprocess(face);
225            batch_data.extend_from_slice(&row);
226        }
227
228        if valid_indices.is_empty() {
229            return vec![None; n];
230        }
231
232        let batch_n = valid_indices.len();
233        let raw_outputs = match self.run_inference_batch(&batch_data, batch_n as i64) {
234            Ok(v) => v,
235            Err(e) => {
236                tracing::warn!("Batch embedding inference failed: {}", e);
237                return vec![None; n];
238            }
239        };
240
241        // Each output row is 512 floats. L2-normalize each.
242        let mut embeddings: Vec<Option<FaceEmbedding>> = vec![None; n];
243        for (out_i, orig_i) in valid_indices.iter().enumerate() {
244            let start = out_i * 512;
245            let end = start + 512;
246            if end > raw_outputs.len() {
247                break;
248            }
249            let row = Array1::from_vec(raw_outputs[start..end].to_vec());
250            let normalized = self.normalize(&row);
251            embeddings[*orig_i] = Some(FaceEmbedding::new(normalized));
252        }
253        embeddings
254    }
255
256    /// Preprocess face image for ArcFace: normalize to [-1, 1], produce NCHW vec
257    fn preprocess(&self, face: &RgbImage) -> Vec<f32> {
258        let mut input = vec![0.0f32; 3 * 112 * 112];
259        let hw = 112 * 112;
260
261        for y in 0..112u32 {
262            for x in 0..112u32 {
263                let pixel = face.get_pixel(x, y);
264                let idx = (y * 112 + x) as usize;
265                input[idx] = (pixel[0] as f32 - 127.5) / 127.5;
266                input[hw + idx] = (pixel[1] as f32 - 127.5) / 127.5;
267                input[2 * hw + idx] = (pixel[2] as f32 - 127.5) / 127.5;
268            }
269        }
270
271        input
272    }
273
274    /// Run ONNX inference with a batch of faces [N, 3, 112, 112].
275    /// Returns a flat Vec<f32> of N * 512 floats.
276    fn run_inference_batch(
277        &mut self,
278        input_data: &[f32],
279        batch_size: i64,
280    ) -> ort::Result<Vec<f32>> {
281        let input_tensor =
282            TensorRef::<f32>::from_array_view((vec![batch_size, 3, 112, 112], input_data))?;
283
284        let outputs = self.session.run(ort::inputs![input_tensor])?;
285
286        let (_name, output) = outputs
287            .iter()
288            .next()
289            .ok_or_else(|| ort::Error::new("No output tensor from ArcFace model".to_string()))?;
290
291        let (_shape, data) = output.try_extract_tensor::<f32>()?;
292        Ok(data.to_vec())
293    }
294
295    /// L2 normalize the embedding vector
296    fn normalize(&self, embedding: &Array1<f32>) -> Array1<f32> {
297        let norm: f32 = embedding.iter().map(|x| x * x).sum::<f32>().sqrt();
298
299        if norm > 0.0 {
300            embedding / norm
301        } else {
302            embedding.clone()
303        }
304    }
305}
306
307#[cfg(test)]
308mod tests {
309    use super::*;
310
311    #[test]
312    fn test_cosine_similarity_identical() {
313        let emb1 = FaceEmbedding::new(Array1::from_vec(vec![1.0, 0.0, 0.0]));
314        let emb2 = FaceEmbedding::new(Array1::from_vec(vec![1.0, 0.0, 0.0]));
315
316        assert!((emb1.cosine_similarity(&emb2) - 1.0).abs() < 0.001);
317    }
318
319    #[test]
320    fn test_cosine_similarity_orthogonal() {
321        let emb1 = FaceEmbedding::new(Array1::from_vec(vec![1.0, 0.0, 0.0]));
322        let emb2 = FaceEmbedding::new(Array1::from_vec(vec![0.0, 1.0, 0.0]));
323
324        assert!((emb1.cosine_similarity(&emb2) - 0.0).abs() < 0.001);
325    }
326
327    #[test]
328    fn test_embedding_serialization() {
329        let original = FaceEmbedding::new(Array1::from_vec(vec![1.0; 512]));
330        let bytes = original.to_bytes();
331        assert_eq!(bytes.len(), 512 * 4);
332
333        let restored = FaceEmbedding::from_bytes(&bytes).unwrap();
334        assert!((original.cosine_similarity(&restored) - 1.0).abs() < 0.001);
335    }
336
337    #[test]
338    fn test_embedding_from_bytes_wrong_size() {
339        let bytes = vec![0u8; 100];
340        assert!(FaceEmbedding::from_bytes(&bytes).is_none());
341    }
342}