smriti/ml/
face_embedder.rs1use 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#[derive(Debug, Clone)]
19pub struct FaceEmbedding {
20 pub vector: Array1<f32>,
22}
23
24impl FaceEmbedding {
25 pub fn new(vector: Array1<f32>) -> Self {
27 Self { vector }
28 }
29
30 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 pub fn to_bytes(&self) -> Vec<u8> {
51 self.vector.iter().flat_map(|f| f.to_le_bytes()).collect()
52 }
53
54 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#[derive(Debug, Clone)]
72pub struct EmbedderConfig {
73 pub model_path: PathBuf,
74 pub gpu_bridge_url: Option<String>,
75 pub expected_model: String,
81 pub intra_threads: usize,
82}
83
84pub enum FaceEmbedder {
92 Local(LocalEmbedder),
93 Remote {
94 remote: RemoteEmbedder,
95 local_fallback: Option<LocalEmbedder>,
96 },
97}
98
99impl FaceEmbedder {
100 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 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 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 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
183pub struct LocalEmbedder {
185 session: Session,
186}
187
188impl LocalEmbedder {
189 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 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 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 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 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 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 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 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}