1use std::collections::HashMap;
7
8use super::FaceEmbedding;
9
10const EXACT_CLUSTER_LIMIT: usize = 256;
14
15#[derive(Debug, Clone)]
16pub struct ClusterInput {
17 pub face_id: i64,
18 pub photo_id: i64,
19 pub current_cluster_id: Option<i64>,
22 pub embedding: FaceEmbedding,
23}
24
25pub struct FaceClusterer {
27 max_distance: f32,
29}
30
31impl FaceClusterer {
32 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 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 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 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 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 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#[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 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 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 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 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 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 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 if let Some(neg) = negatives {
291 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 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 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 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 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 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 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}