Skip to main content

smriti/ml/
retrieval.rs

1//! Gallery-based face retrieval.
2//!
3//! Given a query embedding and a set of per-cluster galleries (N diverse
4//! exemplars per cluster), find the best-matching clusters via k-NN.
5//! Each cluster's score is the mean of the top-k cosine similarities among its
6//! gallery members. This is strictly stronger than centroid matching: a
7//! cluster is "close" if *multiple* of its real members are close, not just
8//! its mean.
9//!
10//! The result drives confidence-banded assignment: HIGH -> auto-assign,
11//! LOW -> leave unassigned, AMBIGUOUS -> queue for user review.
12
13use super::FaceEmbedding;
14
15/// A cluster candidate ranked against the query face.
16#[derive(Debug, Clone)]
17pub struct RetrievalHit {
18    pub cluster_id: i64,
19    /// Mean of the top-k cosine similarities against this cluster's gallery.
20    /// Higher = better. Range [-1, 1], but in practice [0, 1].
21    pub score: f32,
22    /// Face id within the gallery that had the single highest similarity.
23    pub best_match_face_id: i64,
24    /// How many gallery members exceeded `min_similarity`.
25    pub num_matches: usize,
26}
27
28/// Assignment recommendation for a query face.
29#[derive(Debug, Clone)]
30pub enum ConfidenceBand {
31    /// Top candidate exceeds `high_threshold` and leads the runner-up by
32    /// at least `margin`. Safe to auto-assign.
33    High { hit: RetrievalHit },
34    /// Between thresholds, or top-2 candidates too close. Queue for review.
35    Ambiguous {
36        top: RetrievalHit,
37        runner_up: Option<RetrievalHit>,
38    },
39    /// Top score below `low_threshold`. Face does not match any existing
40    /// cluster; leave for agglomerative pass or mark as a new singleton.
41    Low,
42}
43
44/// Thresholds for banding. Tuned for ArcFace / GLinTR L2-normalized embeddings.
45#[derive(Debug, Clone, Copy)]
46pub struct BandingConfig {
47    /// Minimum score to consider any match at all.
48    pub low_threshold: f32,
49    /// Minimum top score to auto-assign without review.
50    pub high_threshold: f32,
51    /// Minimum score gap between top and runner-up to auto-assign.
52    pub margin: f32,
53}
54
55impl Default for BandingConfig {
56    fn default() -> Self {
57        // Tuned to err toward the review queue. A face gets auto-assigned
58        // only when its mean top-K cosine to a cluster's gallery is ≥ 0.65
59        // AND beats the runner-up cluster by ≥ 0.12. Anything weaker goes
60        // to AMBIGUOUS — the user confirms via PersonReview, and that
61        // confirmed face becomes a sticky `user_confirmed` gallery exemplar
62        // that drives all future retrievals. Better to ask twice than to
63        // silently misclassify.
64        Self {
65            low_threshold: 0.40,
66            high_threshold: 0.65,
67            margin: 0.12,
68        }
69    }
70}
71
72/// Retrieve ranked cluster candidates for a query embedding.
73///
74/// `galleries` is a slice of (cluster_id, Vec<(face_id, embedding)>). All
75/// embeddings are assumed L2-normalized (as produced by the embedder and
76/// stored in `person_gallery_embeddings`).
77///
78/// `top_k` is the number of best matches per cluster to average into that
79/// cluster's score. 5 is a good default.
80///
81/// `min_similarity` filters individual gallery members below this cosine
82/// similarity before they contribute to the score.
83///
84/// Clusters in `exclude` are skipped entirely -- used to honor
85/// cannot-merge constraints and same-photo conflict prevention.
86///
87/// Returns clusters sorted by descending score. Empty if no cluster scored
88/// above `min_similarity`.
89pub 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        // Score each member, keep those above min_similarity.
108        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
138/// Classify the retrieval result into a confidence band.
139pub 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); // cluster 2 is below threshold
218        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}