Skip to main content

smriti/ml/
clustering.rs

1//! Face clustering using agglomerative complete-linkage.
2//!
3//! This avoids DBSCAN chaining by only merging clusters where *all*
4//! pairwise distances remain under a strict threshold.
5
6use std::collections::HashMap;
7
8use super::FaceEmbedding;
9
10// HNSW deliberately randomizes graph levels. Exact complete-link clustering
11// is affordable for small batches and makes repeated runs stable; large
12// libraries retain the O(n log n) path.
13const EXACT_CLUSTER_LIMIT: usize = 256;
14
15#[derive(Debug, Clone)]
16pub struct ClusterInput {
17    pub face_id: i64,
18    pub photo_id: i64,
19    /// Current cluster assignment (if any) — used with face_negatives to
20    /// avoid clustering a face back into a cluster it was rejected from.
21    pub current_cluster_id: Option<i64>,
22    pub embedding: FaceEmbedding,
23}
24
25/// Agglomerative clusterer with complete linkage.
26pub struct FaceClusterer {
27    /// Maximum cosine distance (1 - cosine similarity) allowed within cluster.
28    max_distance: f32,
29}
30
31impl FaceClusterer {
32    /// Strict default to reduce false merges.
33    pub fn new() -> Self {
34        Self { max_distance: 0.28 }
35    }
36
37    pub fn with_max_distance(mut self, max_distance: f32) -> Self {
38        self.max_distance = max_distance.clamp(0.05, 1.0);
39        self
40    }
41
42    /// Cluster faces and return map of face_id -> cluster_label (-1 = singleton/noise).
43    ///
44    /// `negatives` is an optional set of (face_id, not_cluster_id) pairs from
45    /// the `face_negatives` table. When present, two faces whose current clusters
46    /// are blocked by a negative constraint will not be merged.
47    pub fn cluster(
48        &self,
49        faces: &[ClusterInput],
50        negatives: Option<&std::collections::HashSet<(i64, i64)>>,
51    ) -> HashMap<i64, i32> {
52        #[cfg(feature = "hnsw_clustering")]
53        {
54            if faces.len() <= EXACT_CLUSTER_LIMIT {
55                self.cluster_complete_link(faces, negatives)
56            } else {
57                cluster_hnsw(faces, self.max_distance, negatives)
58            }
59        }
60        #[cfg(not(feature = "hnsw_clustering"))]
61        {
62            self.cluster_complete_link(faces, negatives)
63        }
64    }
65
66    /// Legacy O(n²) complete-link path. Kept compiled when
67    /// `hnsw_clustering` is off so we can A/B against the new
68    /// implementation on a real library by toggling features.
69    fn cluster_complete_link(
70        &self,
71        faces: &[ClusterInput],
72        negatives: Option<&std::collections::HashSet<(i64, i64)>>,
73    ) -> HashMap<i64, i32> {
74        if faces.is_empty() {
75            return HashMap::new();
76        }
77
78        // Start with each face in its own cluster.
79        let mut clusters: Vec<Vec<usize>> = (0..faces.len()).map(|i| vec![i]).collect();
80
81        loop {
82            let mut best_pair: Option<(usize, usize, f32)> = None;
83
84            for i in 0..clusters.len() {
85                for j in (i + 1)..clusters.len() {
86                    let d = self.complete_link_distance(&clusters[i], &clusters[j], faces);
87                    if d > self.max_distance {
88                        continue;
89                    }
90
91                    let mut merged = clusters[i].clone();
92                    merged.extend_from_slice(&clusters[j]);
93                    if Self::has_same_photo_conflict(&merged, faces) {
94                        continue;
95                    }
96                    if Self::has_negative_conflict(&clusters[i], &clusters[j], faces, negatives) {
97                        continue;
98                    }
99
100                    match best_pair {
101                        Some((_, _, best_d)) if d >= best_d => {}
102                        _ => best_pair = Some((i, j, d)),
103                    }
104                }
105            }
106
107            let Some((a, b, _)) = best_pair else {
108                break;
109            };
110
111            let mut merged = clusters[a].clone();
112            merged.extend_from_slice(&clusters[b]);
113
114            // Merge b into a
115            clusters[a] = merged;
116            clusters.remove(b);
117        }
118
119        let mut out = HashMap::new();
120        let mut label = 0i32;
121        for members in clusters {
122            if members.len() < 2 {
123                out.insert(faces[members[0]].face_id, -1);
124            } else {
125                for idx in members {
126                    out.insert(faces[idx].face_id, label);
127                }
128                label += 1;
129            }
130        }
131        out
132    }
133
134    fn complete_link_distance(&self, a: &[usize], b: &[usize], faces: &[ClusterInput]) -> f32 {
135        let mut max_d = 0.0f32;
136        for &i in a {
137            for &j in b {
138                let sim = faces[i].embedding.cosine_similarity(&faces[j].embedding);
139                let d = 1.0 - sim;
140                if d > max_d {
141                    max_d = d;
142                }
143            }
144        }
145        max_d
146    }
147
148    fn has_same_photo_conflict(cluster: &[usize], faces: &[ClusterInput]) -> bool {
149        let mut seen = std::collections::HashSet::new();
150        for &idx in cluster {
151            let pid = faces[idx].photo_id;
152            if !seen.insert(pid) {
153                return true;
154            }
155        }
156        false
157    }
158
159    /// If any face in cluster A has a (face_id, current_cluster_id-of-some-face-in-B)
160    /// pair in `negatives` — or vice versa — the merge is forbidden.
161    fn has_negative_conflict(
162        a: &[usize],
163        b: &[usize],
164        faces: &[ClusterInput],
165        negatives: Option<&std::collections::HashSet<(i64, i64)>>,
166    ) -> bool {
167        let neg = match negatives {
168            Some(n) if !n.is_empty() => n,
169            _ => return false,
170        };
171        for &ai in a {
172            for &bi in b {
173                if let Some(bc) = faces[bi].current_cluster_id {
174                    if neg.contains(&(faces[ai].face_id, bc)) {
175                        return true;
176                    }
177                }
178                if let Some(ac) = faces[ai].current_cluster_id {
179                    if neg.contains(&(faces[bi].face_id, ac)) {
180                        return true;
181                    }
182                }
183            }
184        }
185        false
186    }
187}
188
189/// HNSW-based clustering — replaces complete-link's O(n²) with
190/// O(n log n) index build + O(n · k) k-NN queries + O(n · α) union-find.
191/// Wall time on 10k faces: <30 s vs hours for the legacy path.
192///
193/// 1. Build an HNSW index over all embeddings using cosine distance.
194/// 2. For each face, query top-K nearest neighbours (K=15).
195/// 3. Edges with similarity ≥ (1 − max_distance) and no same-photo
196///    conflict get unioned via union-find.
197/// 4. Each connected component with ≥ 2 members becomes a cluster.
198#[cfg(feature = "hnsw_clustering")]
199fn cluster_hnsw(
200    faces: &[ClusterInput],
201    max_distance: f32,
202    negatives: Option<&std::collections::HashSet<(i64, i64)>>,
203) -> HashMap<i64, i32> {
204    use hnsw_rs::prelude::*;
205    use std::collections::HashSet;
206    if faces.is_empty() {
207        return HashMap::new();
208    }
209    if faces.len() == 1 {
210        let mut out = HashMap::new();
211        out.insert(faces[0].face_id, -1);
212        return out;
213    }
214
215    let n = faces.len();
216    // M=16 / ef_c=200 are the de-facto-good defaults for cosine
217    // recall on ~512-dim vectors.
218    let max_nb_connection = 16;
219    let ef_c = 200;
220    let nb_layer = 16.min(((n as f32).ln().trunc() as usize).max(1));
221    let hns: Hnsw<f32, DistCosine> = Hnsw::new(max_nb_connection, n, nb_layer, ef_c, DistCosine {});
222
223    // Insert with HNSW's per-row index as the data id; we map back
224    // to face_id via the input slice.
225    let data: Vec<(&[f32], usize)> = faces
226        .iter()
227        .enumerate()
228        .map(|(i, f)| (f.embedding.vector.as_slice().unwrap_or(&[]), i))
229        .collect();
230    hns.parallel_insert_slice(&data);
231
232    // Union-find over face indices.
233    let mut parent: Vec<usize> = (0..n).collect();
234    fn find(parent: &mut [usize], mut x: usize) -> usize {
235        while parent[x] != x {
236            parent[x] = parent[parent[x]];
237            x = parent[x];
238        }
239        x
240    }
241    fn union(parent: &mut [usize], a: usize, b: usize) {
242        let ra = find(parent, a);
243        let rb = find(parent, b);
244        if ra != rb {
245            parent[ra] = rb;
246        }
247    }
248
249    // Build a per-component photo set so we don't union two
250    // components if doing so would put the same photo's faces
251    // into one cluster (same-photo conflict).
252    let mut component_photos: HashMap<usize, std::collections::HashSet<i64>> = HashMap::new();
253    for (i, f) in faces.iter().enumerate() {
254        component_photos.entry(i).or_default().insert(f.photo_id);
255    }
256
257    let knbn = 15;
258    let ef_search = ef_c;
259    for (i, face_i) in faces.iter().enumerate() {
260        let query = match face_i.embedding.vector.as_slice() {
261            Some(s) => s,
262            None => continue,
263        };
264        let neighbours = hns.search(query, knbn, ef_search);
265        for nb in neighbours {
266            let j = nb.d_id;
267            if j == i {
268                continue;
269            }
270            // DistCosine returns 1 - cos_sim, which is exactly our
271            // existing "max_distance" threshold semantics.
272            if nb.distance > max_distance {
273                continue;
274            }
275            let ra = find(&mut parent, i);
276            let rb = find(&mut parent, j);
277            if ra == rb {
278                continue;
279            }
280            // Check same-photo conflict across the two components.
281            let photos_a = component_photos.get(&ra).cloned().unwrap_or_default();
282            let photos_b = component_photos.get(&rb).cloned().unwrap_or_default();
283            if photos_a.intersection(&photos_b).next().is_some() {
284                continue;
285            }
286
287            // Check face_negatives constraints: if any face in component A
288            // has a negative against the current cluster of any face in
289            // component B (or vice versa), skip the merge.
290            if let Some(neg) = negatives {
291                // Collect face_ids in each component.
292                let faces_a: HashSet<i64> = faces
293                    .iter()
294                    .enumerate()
295                    .filter(|(fi, _)| find(&mut parent, *fi) == ra)
296                    .map(|(_, f)| f.face_id)
297                    .collect();
298                let faces_b: HashSet<i64> = faces
299                    .iter()
300                    .enumerate()
301                    .filter(|(fi, _)| find(&mut parent, *fi) == rb)
302                    .map(|(_, f)| f.face_id)
303                    .collect();
304
305                // Check A→B: any face in A has negative against B's cluster?
306                let conflict_a_to_b = faces_a.iter().any(|fa| {
307                    faces_b.iter().any(|fb| {
308                        let fb_cluster = faces
309                            .iter()
310                            .find(|f| f.face_id == *fb)
311                            .and_then(|f| f.current_cluster_id);
312                        fb_cluster.is_some_and(|cid| neg.contains(&(*fa, cid)))
313                    })
314                });
315                // Check B→A: any face in B has negative against A's cluster?
316                let conflict_b_to_a = faces_b.iter().any(|fb| {
317                    faces_a.iter().any(|fa| {
318                        let fa_cluster = faces
319                            .iter()
320                            .find(|f| f.face_id == *fa)
321                            .and_then(|f| f.current_cluster_id);
322                        fa_cluster.is_some_and(|cid| neg.contains(&(*fb, cid)))
323                    })
324                });
325
326                if conflict_a_to_b || conflict_b_to_a {
327                    continue;
328                }
329            }
330
331            union(&mut parent, ra, rb);
332            // Move the smaller set into the larger; track on the new root.
333            let new_root = find(&mut parent, ra);
334            let other = if new_root == ra { rb } else { ra };
335            let mut merged = photos_a;
336            merged.extend(photos_b);
337            component_photos.insert(new_root, merged);
338            component_photos.remove(&other);
339        }
340    }
341
342    // Group by root → emit cluster labels.
343    let mut groups: HashMap<usize, Vec<usize>> = HashMap::new();
344    for i in 0..n {
345        let r = find(&mut parent, i);
346        groups.entry(r).or_default().push(i);
347    }
348
349    let mut out = HashMap::new();
350    let mut label = 0i32;
351    for (_root, members) in groups {
352        if members.len() < 2 {
353            out.insert(faces[members[0]].face_id, -1);
354        } else {
355            for idx in members {
356                out.insert(faces[idx].face_id, label);
357            }
358            label += 1;
359        }
360    }
361    out
362}
363
364impl Default for FaceClusterer {
365    fn default() -> Self {
366        Self::new()
367    }
368}
369
370#[cfg(test)]
371mod tests {
372    use super::*;
373
374    fn emb(v: &[f32]) -> FaceEmbedding {
375        FaceEmbedding::new(ndarray::Array1::from_vec(v.to_vec()))
376    }
377
378    #[test]
379    fn complete_link_rejects_chain() {
380        // A-B close, B-C close, A-C far => should not all merge.
381        let mk = |id: i64, p: i64, vec: Vec<f32>| ClusterInput {
382            face_id: id,
383            photo_id: p,
384            current_cluster_id: None,
385            embedding: emb(&vec),
386        };
387        let a = mk(1, 11, vec![1.0, 0.0, 0.0]);
388        let b = mk(2, 12, vec![0.8, 0.2, 0.0]);
389        let c = mk(3, 13, vec![0.0, 1.0, 0.0]);
390
391        let clusterer = FaceClusterer::new().with_max_distance(0.35);
392        let labels = clusterer.cluster(&[a, b, c], None);
393        // Should not place all three in one non-negative cluster.
394        let vals: Vec<i32> = labels.values().copied().filter(|v| *v >= 0).collect();
395        assert!(vals.len() < 3);
396    }
397
398    #[test]
399    fn prevents_same_photo_merge() {
400        let mk = |id: i64, p: i64| ClusterInput {
401            face_id: id,
402            photo_id: p,
403            current_cluster_id: None,
404            embedding: emb(&[1.0, 0.0, 0.0]),
405        };
406
407        let clusterer = FaceClusterer::new();
408        let labels = clusterer.cluster(&[mk(1, 42), mk(2, 42)], None);
409        assert_eq!(labels[&1], -1);
410        assert_eq!(labels[&2], -1);
411    }
412}