Skip to main content

smriti/ml/
resolver.rs

1//! Context-aware resolver for face retrieval.
2//!
3//! Sits between `retrieve_candidates` and final assignment: re-scores the
4//! ranked clusters using contextual signals and returns a stronger
5//! recommendation, or flags that the decision should still go to the user.
6//!
7//! Signals fused:
8//!
9//! - Co-occurrence: if the photo already has other assigned clusters,
10//!   candidates with higher co-occurrence to those clusters are boosted.
11//!   "Same event, same cast" principle.
12//!
13//! - Temporal neighbors: if nearby-in-time photos have the candidate
14//!   cluster already assigned, boost. "Same session → same people".
15//!
16//! The embedding score is always the dominant signal; context only breaks
17//! ties and nudges ambiguous cases.
18
19use super::retrieval::RetrievalHit;
20
21/// Numeric weights controlling the relative influence of each signal.
22/// Default values roughly chosen so context nudges but doesn't dominate
23/// a strong embedding disagreement.
24#[derive(Debug, Clone, Copy)]
25pub struct ResolverWeights {
26    /// Multiplier on the cosine similarity score. Keep at 1.0.
27    pub embedding: f32,
28    /// Added per other-cluster-in-photo * cooccurrence_count / saturation.
29    pub cooccurrence: f32,
30    /// Added when the candidate is also assigned in a temporally nearby photo.
31    pub temporal: f32,
32}
33
34impl Default for ResolverWeights {
35    fn default() -> Self {
36        Self {
37            embedding: 1.0,
38            cooccurrence: 0.30,
39            temporal: 0.50,
40        }
41    }
42}
43
44/// Context provided at resolve time. All queryable from `FaceRepo`.
45#[derive(Debug, Clone)]
46pub struct ResolverContext {
47    /// Cluster IDs of other assigned faces in the same photo.
48    pub photo_other_clusters: Vec<i64>,
49    /// (candidate_cluster_id, cooccurrence_count) for each cluster
50    /// currently in the photo vs. each candidate. Caller aggregates.
51    pub cooccurrence_scores: std::collections::HashMap<i64, i64>,
52    /// Cluster IDs assigned to faces in temporally-nearby photos.
53    pub temporal_neighbor_clusters: std::collections::HashSet<i64>,
54}
55
56/// Apply contextual signals to a ranked list of retrieval hits and return
57/// the re-sorted list.
58///
59/// Output is sorted by adjusted score descending.
60pub fn rerank(
61    hits: &[RetrievalHit],
62    ctx: &ResolverContext,
63    weights: ResolverWeights,
64) -> Vec<RetrievalHit> {
65    if hits.is_empty() {
66        return Vec::new();
67    }
68
69    // Saturation normalizer: avoid letting a single prolific pair dominate.
70    // Map raw count into a bounded log-ish boost.
71    let cooccurrence_bonus = |count: i64| -> f32 {
72        if count <= 0 {
73            0.0
74        } else {
75            ((count as f32 + 1.0).ln() / 5.0).min(1.0)
76        }
77    };
78
79    let mut adjusted: Vec<RetrievalHit> = hits
80        .iter()
81        .map(|h| {
82            let mut score = h.score * weights.embedding;
83
84            if weights.cooccurrence > 0.0 && !ctx.photo_other_clusters.is_empty() {
85                let raw = ctx
86                    .cooccurrence_scores
87                    .get(&h.cluster_id)
88                    .copied()
89                    .unwrap_or(0);
90                score += weights.cooccurrence * cooccurrence_bonus(raw);
91            }
92
93            if weights.temporal > 0.0 && ctx.temporal_neighbor_clusters.contains(&h.cluster_id) {
94                score += weights.temporal * 0.1;
95            }
96
97            RetrievalHit {
98                cluster_id: h.cluster_id,
99                score,
100                best_match_face_id: h.best_match_face_id,
101                num_matches: h.num_matches,
102            }
103        })
104        .collect();
105
106    adjusted.sort_by(|a, b| {
107        b.score
108            .partial_cmp(&a.score)
109            .unwrap_or(std::cmp::Ordering::Equal)
110    });
111    adjusted
112}
113
114#[cfg(test)]
115mod tests {
116    use super::*;
117    use std::collections::{HashMap, HashSet};
118
119    fn hit(cluster_id: i64, score: f32) -> RetrievalHit {
120        RetrievalHit {
121            cluster_id,
122            score,
123            best_match_face_id: 0,
124            num_matches: 1,
125        }
126    }
127
128    #[test]
129    fn no_context_preserves_order() {
130        let hits = vec![hit(1, 0.60), hit(2, 0.55)];
131        let ctx = ResolverContext {
132            photo_other_clusters: Vec::new(),
133            cooccurrence_scores: HashMap::new(),
134            temporal_neighbor_clusters: HashSet::new(),
135        };
136        let out = rerank(&hits, &ctx, ResolverWeights::default());
137        assert_eq!(out[0].cluster_id, 1);
138        assert_eq!(out[1].cluster_id, 2);
139    }
140
141    #[test]
142    fn cooccurrence_can_flip_close_pair() {
143        // Cluster 2 has ~slightly worse embedding but co-occurs strongly with
144        // an already-present cluster. Should win after re-rank.
145        let hits = vec![hit(1, 0.52), hit(2, 0.50)];
146        let mut co = HashMap::new();
147        co.insert(2i64, 50i64); // strong co-occurrence
148        let ctx = ResolverContext {
149            photo_other_clusters: vec![99],
150            cooccurrence_scores: co,
151            temporal_neighbor_clusters: HashSet::new(),
152        };
153        let out = rerank(&hits, &ctx, ResolverWeights::default());
154        assert_eq!(out[0].cluster_id, 2);
155    }
156
157    #[test]
158    fn temporal_boost_applies() {
159        let hits = vec![hit(1, 0.50), hit(2, 0.49)];
160        let mut neighbors = HashSet::new();
161        neighbors.insert(2i64);
162        let ctx = ResolverContext {
163            photo_other_clusters: Vec::new(),
164            cooccurrence_scores: HashMap::new(),
165            temporal_neighbor_clusters: neighbors,
166        };
167        let out = rerank(&hits, &ctx, ResolverWeights::default());
168        assert_eq!(out[0].cluster_id, 2);
169    }
170}