1use super::FaceEmbedding;
14
15#[derive(Debug, Clone)]
17pub struct RetrievalHit {
18 pub cluster_id: i64,
19 pub score: f32,
22 pub best_match_face_id: i64,
24 pub num_matches: usize,
26}
27
28#[derive(Debug, Clone)]
30pub enum ConfidenceBand {
31 High { hit: RetrievalHit },
34 Ambiguous {
36 top: RetrievalHit,
37 runner_up: Option<RetrievalHit>,
38 },
39 Low,
42}
43
44#[derive(Debug, Clone, Copy)]
46pub struct BandingConfig {
47 pub low_threshold: f32,
49 pub high_threshold: f32,
51 pub margin: f32,
53}
54
55impl Default for BandingConfig {
56 fn default() -> Self {
57 Self {
65 low_threshold: 0.40,
66 high_threshold: 0.65,
67 margin: 0.12,
68 }
69 }
70}
71
72pub fn retrieve_candidates(
90 query: &FaceEmbedding,
91 galleries: &[(i64, Vec<(i64, FaceEmbedding)>)],
92 top_k: usize,
93 min_similarity: f32,
94 exclude: &std::collections::HashSet<i64>,
95) -> Vec<RetrievalHit> {
96 if top_k == 0 {
97 return Vec::new();
98 }
99
100 let mut hits: Vec<RetrievalHit> = Vec::with_capacity(galleries.len());
101
102 for (cluster_id, members) in galleries {
103 if exclude.contains(cluster_id) || members.is_empty() {
104 continue;
105 }
106
107 let mut scored: Vec<(i64, f32)> = members
109 .iter()
110 .map(|(face_id, emb)| (*face_id, query.cosine_similarity(emb)))
111 .filter(|(_, s)| *s >= min_similarity)
112 .collect();
113
114 if scored.is_empty() {
115 continue;
116 }
117
118 scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
119 let k = top_k.min(scored.len());
120 let mean_top_k: f32 = scored.iter().take(k).map(|(_, s)| s).sum::<f32>() / k as f32;
121
122 hits.push(RetrievalHit {
123 cluster_id: *cluster_id,
124 score: mean_top_k,
125 best_match_face_id: scored[0].0,
126 num_matches: scored.len(),
127 });
128 }
129
130 hits.sort_by(|a, b| {
131 b.score
132 .partial_cmp(&a.score)
133 .unwrap_or(std::cmp::Ordering::Equal)
134 });
135 hits
136}
137
138pub fn classify(hits: &[RetrievalHit], cfg: &BandingConfig) -> ConfidenceBand {
140 match hits.len() {
141 0 => ConfidenceBand::Low,
142 1 => {
143 let top = hits[0].clone();
144 if top.score >= cfg.high_threshold {
145 ConfidenceBand::High { hit: top }
146 } else if top.score >= cfg.low_threshold {
147 ConfidenceBand::Ambiguous {
148 top,
149 runner_up: None,
150 }
151 } else {
152 ConfidenceBand::Low
153 }
154 }
155 _ => {
156 let top = hits[0].clone();
157 let runner_up = hits[1].clone();
158 if top.score < cfg.low_threshold {
159 ConfidenceBand::Low
160 } else if top.score >= cfg.high_threshold && (top.score - runner_up.score) >= cfg.margin
161 {
162 ConfidenceBand::High { hit: top }
163 } else {
164 ConfidenceBand::Ambiguous {
165 top,
166 runner_up: Some(runner_up),
167 }
168 }
169 }
170 }
171}
172
173#[cfg(test)]
174mod tests {
175 use super::*;
176 use ndarray::Array1;
177 use std::collections::HashSet;
178
179 fn emb(v: Vec<f32>) -> FaceEmbedding {
180 let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
181 let normed: Vec<f32> = v.iter().map(|x| x / norm).collect();
182 FaceEmbedding::new(Array1::from_vec(normed))
183 }
184
185 #[test]
186 fn empty_galleries_yield_no_hits() {
187 let q = emb(vec![1.0, 0.0, 0.0]);
188 let hits = retrieve_candidates(&q, &[], 3, 0.3, &HashSet::new());
189 assert!(hits.is_empty());
190 }
191
192 #[test]
193 fn zero_top_k_yields_no_hits() {
194 let q = emb(vec![1.0, 0.0, 0.0]);
195 let galleries = vec![(1, vec![(10, emb(vec![1.0, 0.0, 0.0]))])];
196
197 let hits = retrieve_candidates(&q, &galleries, 0, 0.3, &HashSet::new());
198
199 assert!(hits.is_empty());
200 }
201
202 #[test]
203 fn near_identical_gallery_scores_highest() {
204 let q = emb(vec![1.0, 0.0, 0.0]);
205 let galleries: Vec<(i64, Vec<(i64, FaceEmbedding)>)> = vec![
206 (
207 1,
208 vec![
209 (10, emb(vec![0.99, 0.01, 0.0])),
210 (11, emb(vec![0.98, 0.0, 0.1])),
211 ],
212 ),
213 (2, vec![(20, emb(vec![0.0, 1.0, 0.0]))]),
214 ];
215
216 let hits = retrieve_candidates(&q, &galleries, 3, 0.3, &HashSet::new());
217 assert_eq!(hits.len(), 1); assert_eq!(hits[0].cluster_id, 1);
219 assert!(hits[0].score > 0.98);
220 }
221
222 #[test]
223 fn exclusion_is_respected() {
224 let q = emb(vec![1.0, 0.0, 0.0]);
225 let galleries: Vec<(i64, Vec<(i64, FaceEmbedding)>)> = vec![
226 (1, vec![(10, emb(vec![0.99, 0.01, 0.0]))]),
227 (2, vec![(20, emb(vec![0.95, 0.05, 0.0]))]),
228 ];
229
230 let mut exclude = HashSet::new();
231 exclude.insert(1);
232 let hits = retrieve_candidates(&q, &galleries, 3, 0.3, &exclude);
233 assert_eq!(hits.len(), 1);
234 assert_eq!(hits[0].cluster_id, 2);
235 }
236
237 #[test]
238 fn classify_bands() {
239 let cfg = BandingConfig::default();
240
241 let low = vec![RetrievalHit {
242 cluster_id: 1,
243 score: 0.2,
244 best_match_face_id: 1,
245 num_matches: 1,
246 }];
247 assert!(matches!(classify(&low, &cfg), ConfidenceBand::Low));
248
249 let high = vec![RetrievalHit {
250 cluster_id: 1,
251 score: 0.7,
252 best_match_face_id: 1,
253 num_matches: 3,
254 }];
255 assert!(matches!(classify(&high, &cfg), ConfidenceBand::High { .. }));
256
257 let ambiguous_pair = vec![
258 RetrievalHit {
259 cluster_id: 1,
260 score: 0.60,
261 best_match_face_id: 1,
262 num_matches: 2,
263 },
264 RetrievalHit {
265 cluster_id: 2,
266 score: 0.57,
267 best_match_face_id: 2,
268 num_matches: 2,
269 },
270 ];
271 assert!(matches!(
272 classify(&ambiguous_pair, &cfg),
273 ConfidenceBand::Ambiguous { .. }
274 ));
275 }
276}