use std::collections::BinaryHeap;
pub const DEFAULT_MERGE_SIM: f32 = 0.35;
pub const DEFAULT_MIN_FACE_PX: f32 = 80.0;
pub const DEFAULT_MAX_GENERIC_SIM: f32 = 0.40;
pub fn average_linkage_cosine(
points: &[(i64, Vec<f32>)],
eps: f32,
min_samples: usize,
silent: bool,
) -> Vec<(i64, Option<i64>)> {
let n = points.len();
if n == 0 { return Vec::new(); }
let clusters = agglomerate_average(points, eps, silent);
label_clusters(points, &clusters, min_samples)
}
fn agglomerate_average(points: &[(i64, Vec<f32>)], eps: f32, silent: bool) -> Vec<Vec<usize>> {
let n = points.len();
if n == 0 {
return Vec::new();
}
let mut sums: Vec<Vec<f32>> = points.iter().map(|(_, v)| v.clone()).collect();
let mut members: Vec<Vec<usize>> = (0..n).map(|i| vec![i]).collect();
let mut alive = vec![true; n];
let mut heap: BinaryHeap<HeapEntry> = BinaryHeap::new();
let progress = crate::progress::Progress::new(n as u64, silent);
seed_eps_eligible_pairs_via_gemm(points, eps, &mut heap, &progress);
progress.finish();
while let Some(HeapEntry { dist: d, i, j }) = heap.pop() {
if !alive[i] || !alive[j] {
continue;
}
let size_i = members[i].len() as f32;
let size_j = members[j].len() as f32;
let current_d = cluster_dist_from_sums(&sums[i], &sums[j], size_i, size_j);
if current_d != d {
continue; }
let moved = std::mem::take(&mut members[j]);
members[i].extend(moved);
let moved_sum = std::mem::take(&mut sums[j]);
for (s, v) in sums[i].iter_mut().zip(&moved_sum) {
*s += v;
}
alive[j] = false;
let new_size_i = members[i].len() as f32;
for k in 0..n {
if k == i || k == j || !alive[k] {
continue;
}
let size_k = members[k].len() as f32;
let new_d = cluster_dist_from_sums(&sums[i], &sums[k], new_size_i, size_k);
if new_d <= eps {
heap.push(HeapEntry { dist: new_d, i: i.min(k), j: i.max(k) });
}
}
}
(0..n).filter(|&r| alive[r]).map(|r| std::mem::take(&mut members[r])).collect()
}
const GEMM_FILTER_SLACK: f32 = 1e-4;
fn seed_eps_eligible_pairs_via_gemm(
points: &[(i64, Vec<f32>)],
eps: f32,
heap: &mut BinaryHeap<HeapEntry>,
progress: &crate::progress::Progress,
) {
seed_eps_eligible_pairs_via_gemm_blocked(points, eps, heap, progress, 1024)
}
fn seed_eps_eligible_pairs_via_gemm_blocked(
points: &[(i64, Vec<f32>)],
eps: f32,
heap: &mut BinaryHeap<HeapEntry>,
progress: &crate::progress::Progress,
block: usize,
) {
let n = points.len();
if n == 0 {
return;
}
let dim = points[0].1.len();
debug_assert!(
points.iter().all(|(_, v)| v.len() == dim),
"all embeddings must share the same dimensionality"
);
let mut flat: Vec<f32> = Vec::with_capacity(n * dim);
for (_, v) in points {
flat.extend_from_slice(v);
}
let mut block_start = 0;
while block_start < n {
let block_len = block.min(n - block_start);
let mut out = vec![0.0f32; block_len * n];
unsafe {
matrixmultiply::sgemm(
block_len, dim, n,
1.0,
flat.as_ptr().add(block_start * dim), dim as isize, 1,
flat.as_ptr(), 1, dim as isize,
0.0,
out.as_mut_ptr(), n as isize, 1,
);
}
for bi in 0..block_len {
let i = block_start + bi;
for j in (i + 1)..n {
let approx_sim = out[bi * n + j];
let approx_d = 1.0 - approx_sim;
if approx_d > eps + GEMM_FILTER_SLACK {
continue;
}
let d = cosine_dist(&points[i].1, &points[j].1);
if d <= eps {
heap.push(HeapEntry { dist: d, i, j });
}
}
progress.tick();
}
block_start += block_len;
}
}
#[cfg(test)]
mod gemm_seeding_tests {
use super::*;
#[test]
fn gemm_seeding_matches_naive_scan_exactly_including_distance_values() {
let points: Vec<(i64, Vec<f32>)> = vec![
(1, vec![1.0, 0.0, 0.0]),
(2, vec![0.99, 0.01, 0.0].iter().map(|x| x / (0.99f32 * 0.99 + 0.01 * 0.01).sqrt()).collect()),
(3, vec![0.0, 1.0, 0.0]),
(4, vec![0.0, 0.99, 0.01].iter().map(|x| x / (0.99f32 * 0.99 + 0.01 * 0.01).sqrt()).collect()),
];
let eps = 0.1;
let mut gemm_heap: BinaryHeap<HeapEntry> = BinaryHeap::new();
let progress = crate::progress::Progress::new(points.len() as u64, true);
seed_eps_eligible_pairs_via_gemm(&points, eps, &mut gemm_heap, &progress);
progress.finish();
let mut naive_pairs: Vec<(usize, usize, f32)> = Vec::new();
for i in 0..points.len() {
for j in (i + 1)..points.len() {
let d = cosine_dist(&points[i].1, &points[j].1);
if d <= eps {
naive_pairs.push((i, j, d));
}
}
}
let mut gemm_pairs: Vec<(usize, usize, f32)> =
gemm_heap.into_iter().map(|e| (e.i, e.j, e.dist)).collect();
gemm_pairs.sort_by(|a, b| (a.0, a.1).cmp(&(b.0, b.1)));
naive_pairs.sort_by(|a, b| (a.0, a.1).cmp(&(b.0, b.1)));
assert_eq!(gemm_pairs.len(), naive_pairs.len(), "pair count must match");
for (gemm, naive) in gemm_pairs.iter().zip(&naive_pairs) {
assert_eq!((gemm.0, gemm.1), (naive.0, naive.1), "pair indices must match");
assert_eq!(
gemm.2.to_bits(),
naive.2.to_bits(),
"GEMM-filtered distance must be BIT-IDENTICAL to the naive scalar distance, not just approximately equal - \
this is what the merge loop's exact-equality staleness check depends on"
);
}
let pairs_only: Vec<(usize, usize)> = naive_pairs.iter().map(|(i, j, _)| (*i, *j)).collect();
assert_eq!(pairs_only, vec![(0, 1), (2, 3)], "sanity: only the two near-identical pairs should qualify");
}
#[test]
fn gemm_seeding_handles_multiple_blocks_including_a_partial_trailing_block() {
let points: Vec<(i64, Vec<f32>)> = vec![
(1, vec![1.0, 0.0, 0.0]), (2, vec![0.95, 0.312, 0.0]), (3, vec![0.0, 1.0, 0.0]), (4, vec![0.0, 0.95, 0.312]), (5, vec![0.0, 0.0, 1.0]), ];
let eps = 0.1;
let mut naive_pairs: Vec<(usize, usize)> = Vec::new();
for i in 0..points.len() {
for j in (i + 1)..points.len() {
if cosine_dist(&points[i].1, &points[j].1) <= eps {
naive_pairs.push((i, j));
}
}
}
naive_pairs.sort();
assert!(!naive_pairs.is_empty(), "fixture must have at least one eps-eligible pair to be a meaningful test");
let mut heap: BinaryHeap<HeapEntry> = BinaryHeap::new();
let progress = crate::progress::Progress::new(points.len() as u64, true);
seed_eps_eligible_pairs_via_gemm_blocked(&points, eps, &mut heap, &progress, 2);
progress.finish();
let mut gemm_pairs: Vec<(usize, usize)> = heap.into_iter().map(|e| (e.i, e.j)).collect();
gemm_pairs.sort();
assert_eq!(gemm_pairs, naive_pairs, "blocked GEMM seeding (block=2, partial trailing block) must match the naive scan exactly");
}
#[test]
fn gemm_seeding_handles_n_equal_one_and_n_equal_block() {
let single: Vec<(i64, Vec<f32>)> = vec![(1, vec![1.0, 0.0, 0.0])];
let mut heap: BinaryHeap<HeapEntry> = BinaryHeap::new();
let progress = crate::progress::Progress::new(1, true);
seed_eps_eligible_pairs_via_gemm_blocked(&single, 0.1, &mut heap, &progress, 2);
progress.finish();
assert_eq!(heap.len(), 0, "a single point has no pairs");
let v = vec![1.0f32, 0.0, 0.0];
let exact_block: Vec<(i64, Vec<f32>)> = (0..2).map(|i| (i, v.clone())).collect();
let mut heap2: BinaryHeap<HeapEntry> = BinaryHeap::new();
let progress2 = crate::progress::Progress::new(2, true);
seed_eps_eligible_pairs_via_gemm_blocked(&exact_block, 0.01, &mut heap2, &progress2, 2);
progress2.finish();
assert_eq!(heap2.len(), 1, "n == block must not skip the (only) pair");
}
}
fn centroid(points: &[(i64, Vec<f32>)], member_idxs: &[usize]) -> Vec<f32> {
let dim = points[member_idxs[0]].1.len();
let mut sum = vec![0.0f32; dim];
for &idx in member_idxs {
for (s, v) in sum.iter_mut().zip(&points[idx].1) { *s += v; }
}
let norm = sum.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 1e-12 { for s in &mut sum { *s /= norm; } }
sum
}
fn merge_by_centroid(
points: &[(i64, Vec<f32>)],
mut clusters: Vec<Vec<usize>>,
merge_sim: f32,
) -> Vec<Vec<usize>> {
let mut centroids: Vec<Vec<f32>> = clusters.iter().map(|c| centroid(points, c)).collect();
let c = clusters.len();
let mut sim: Vec<Vec<f32>> = vec![vec![0.0f32; c]; c];
for i in 0..c {
for j in (i + 1)..c {
let s = centroids[i].iter().zip(¢roids[j]).map(|(a, b)| a * b).sum::<f32>();
sim[i][j] = s;
sim[j][i] = s;
}
}
loop {
let mut best: Option<(f32, usize, usize)> = None;
for i in 0..clusters.len() {
for j in (i + 1)..clusters.len() {
let s = sim[i][j];
if best.is_none_or(|(bs, _, _)| s > bs) {
best = Some((s, i, j));
}
}
}
let Some((s, i, j)) = best else { break };
if s < merge_sim { break; }
let moved = std::mem::take(&mut clusters[j]);
clusters[i].extend(moved);
clusters.swap_remove(j);
centroids.swap_remove(j);
centroids[i] = centroid(points, &clusters[i]);
let last = sim.len() - 1;
if j != last {
sim.swap(j, last);
for row in sim.iter_mut() {
row.swap(j, last);
}
}
sim.pop();
for row in sim.iter_mut() {
row.pop();
}
for k in 0..clusters.len() {
if k == i {
continue;
}
let s = centroids[i].iter().zip(¢roids[k]).map(|(a, b)| a * b).sum::<f32>();
sim[i][k] = s;
sim[k][i] = s;
}
}
clusters
}
fn label_clusters(
points: &[(i64, Vec<f32>)],
clusters: &[Vec<usize>],
min_samples: usize,
) -> Vec<(i64, Option<i64>)> {
let mut labels: Vec<Option<i64>> = vec![None; points.len()];
let mut cluster_id: i64 = 0;
for members in clusters {
if members.len() < min_samples { continue; }
for &idx in members {
labels[idx] = Some(cluster_id);
}
cluster_id += 1;
}
points.iter().zip(labels).map(|((id, _), lbl)| (*id, lbl)).collect()
}
#[derive(Copy, Clone, PartialEq)]
struct HeapEntry {
dist: f32,
i: usize,
j: usize,
}
impl Eq for HeapEntry {}
impl Ord for HeapEntry {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other.dist.total_cmp(&self.dist)
}
}
impl PartialOrd for HeapEntry {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> { Some(self.cmp(other)) }
}
fn cosine_dist(a: &[f32], b: &[f32]) -> f32 {
1.0 - a.iter().zip(b).map(|(x, y)| x * y).sum::<f32>()
}
fn cluster_dist_from_sums(sum_a: &[f32], sum_b: &[f32], size_a: f32, size_b: f32) -> f32 {
let dot: f32 = sum_a.iter().zip(sum_b).map(|(x, y)| x * y).sum();
1.0 - dot / (size_a * size_b)
}
#[cfg(test)]
mod sums_distance_tests {
use super::*;
#[test]
fn singleton_clusters_match_plain_cosine_distance() {
let a = vec![1.0f32, 0.0, 0.0];
let b = vec![0.6f32, 0.8, 0.0]; let expected = cosine_dist(&a, &b);
let actual = cluster_dist_from_sums(&a, &b, 1.0, 1.0);
assert!((actual - expected).abs() < 1e-6, "expected {expected}, got {actual}");
}
#[test]
fn merged_cluster_distance_matches_direct_average_over_all_cross_pairs() {
let a1 = vec![1.0f32, 0.0, 0.0];
let a2 = vec![0.6f32, 0.8, 0.0]; let b1 = vec![0.0f32, 1.0, 0.0];
let sum_a: Vec<f32> = a1.iter().zip(&a2).map(|(x, y)| x + y).collect();
let direct_avg = (cosine_dist(&a1, &b1) + cosine_dist(&a2, &b1)) / 2.0;
let via_sums = cluster_dist_from_sums(&sum_a, &b1, 2.0, 1.0);
assert!(
(via_sums - direct_avg).abs() < 1e-6,
"sums-based distance {via_sums} must match direct average {direct_avg}"
);
}
}
pub fn cluster_faces(
points: &[(i64, Vec<f32>)],
eps: f32,
min_samples: usize,
merge_sim: f32,
silent: bool,
) -> Vec<(i64, Option<i64>)> {
let n = points.len();
if n == 0 { return Vec::new(); }
let clusters = agglomerate_average(points, eps, silent);
let (mergeable, rest): (Vec<Vec<usize>>, Vec<Vec<usize>>) =
clusters.into_iter().partition(|c| c.len() >= min_samples);
let mut merged = merge_by_centroid(points, mergeable, merge_sim);
merged.extend(rest);
label_clusters(points, &merged, min_samples)
}
#[cfg(test)]
mod tests {
use super::*;
fn l2(v: Vec<f32>) -> Vec<f32> {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
v.into_iter().map(|x| x / norm).collect()
}
#[test]
fn two_close_vectors_form_cluster() {
let v1 = l2(vec![1.0f32, 0.01, 0.0]);
let v2 = l2(vec![1.0f32, 0.02, 0.0]);
let v3 = l2(vec![0.0f32, 1.0, 0.0]);
let result = average_linkage_cosine(&[(1, v1), (2, v2), (3, v3)], 0.1, 2, true);
let map: std::collections::HashMap<_, _> = result.into_iter().collect();
assert_eq!(map[&1], map[&2], "close vectors must share cluster");
assert_eq!(map[&3], None, "distant vector must be outlier");
}
#[test]
fn identical_vectors_cluster_together() {
let v = l2(vec![1.0f32, 0.0, 0.0]);
let result = average_linkage_cosine(&[(1, v.clone()), (2, v.clone()), (3, v)], 0.05, 2, true);
let ids: Vec<_> = result.iter().map(|(_, c)| *c).collect();
assert!(ids.iter().all(|c| c.is_some()), "all must be clustered");
assert_eq!(ids[0], ids[1]);
assert_eq!(ids[1], ids[2]);
}
#[test]
fn all_noise_when_min_samples_too_high() {
let v = l2(vec![1.0f32, 0.0]);
let result = average_linkage_cosine(&[(1, v.clone()), (2, v)], 0.05, 10, true);
assert!(result.iter().all(|(_, c)| c.is_none()));
}
#[test]
fn empty_input_returns_empty() {
let result = average_linkage_cosine(&[], 0.4, 2, true);
assert!(result.is_empty());
}
#[test]
fn chain_of_similar_pairs_does_not_merge_into_one_cluster() {
let angles = [0.0f32, 60.0, 120.0, 180.0, 240.0];
let points: Vec<(i64, Vec<f32>)> = angles
.iter()
.enumerate()
.map(|(i, deg)| {
let rad = deg.to_radians();
(i as i64, vec![rad.cos(), rad.sin()])
})
.collect();
let result = average_linkage_cosine(&points, 0.6, 2, true);
let map: std::collections::HashMap<_, _> = result.into_iter().collect();
let cluster_ids: std::collections::HashSet<_> =
map.values().filter_map(|c| *c).collect();
assert!(
cluster_ids.len() > 1,
"chain must not collapse into a single cluster, got cluster ids {cluster_ids:?}"
);
assert_ne!(map[&0], map[&4], "chain endpoints must not share a cluster");
}
#[test]
fn one_bad_pair_does_not_block_an_otherwise_strong_merge() {
let deg = |d: f32| { let r = d.to_radians(); vec![r.cos(), r.sin()] };
let points: Vec<(i64, Vec<f32>)> = vec![
(1, deg(0.0)), (2, deg(0.0)), (3, deg(0.0)),
(4, deg(20.0)), (5, deg(20.0)),
(6, deg(70.0)),
];
let result = average_linkage_cosine(&points, 0.6, 2, true);
let map: std::collections::HashMap<_, _> = result.into_iter().collect();
assert!(map[&1].is_some(), "the core group must still cluster");
assert_eq!(map[&1], map[&6], "the odd-angle photo must join the same person's cluster");
}
#[test]
fn two_distinct_clusters() {
let a1 = l2(vec![1.0f32, 0.0, 0.0]);
let a2 = l2(vec![0.99f32, 0.01, 0.0]);
let b1 = l2(vec![0.0f32, 1.0, 0.0]);
let b2 = l2(vec![0.0f32, 0.99, 0.01]);
let result = average_linkage_cosine(&[(1, a1), (2, a2), (3, b1), (4, b2)], 0.1, 2, true);
let map: std::collections::HashMap<_, _> = result.into_iter().collect();
assert_ne!(map[&1], map[&3]);
assert_eq!(map[&1], map[&2]);
assert_eq!(map[&3], map[&4]);
}
fn same_identity_two_subclusters() -> Vec<(i64, Vec<f32>)> {
vec![
(1, l2(vec![1.0, 1.0, 0.0, 0.15, 0.0, 0.0])),
(2, l2(vec![1.0, 1.0, 0.0, 0.0, 0.15, 0.0])),
(3, l2(vec![1.0, 1.0, 0.0, 0.0, 0.0, 0.15])),
(4, l2(vec![1.0, 0.0, 1.0, 0.15, 0.0, 0.0])),
(5, l2(vec![1.0, 0.0, 1.0, 0.0, 0.15, 0.0])),
(6, l2(vec![1.0, 0.0, 1.0, 0.0, 0.0, 0.15])),
]
}
#[test]
fn average_linkage_alone_splits_the_two_subclusters() {
let result = average_linkage_cosine(&same_identity_two_subclusters(), 0.3, 1, true);
let clusters: std::collections::HashSet<_> =
result.iter().filter_map(|(_, c)| *c).collect();
assert!(clusters.len() >= 2, "premise: average-linkage should split them, got {clusters:?}");
}
#[test]
fn centroid_merge_reunites_one_persons_fragmented_subclusters() {
let result = cluster_faces(&same_identity_two_subclusters(), 0.3, 1, 0.4, true);
let map: std::collections::HashMap<_, _> = result.into_iter().collect();
let c1 = map[&1];
assert!(c1.is_some(), "faces must be clustered, not left as noise");
for id in 2..=6 {
assert_eq!(map[&id], c1, "all six same-identity faces must share one cluster");
}
}
#[test]
fn incremental_centroid_merge_matches_full_rescan_reference() {
fn reference_merge_by_centroid(
points: &[(i64, Vec<f32>)],
mut clusters: Vec<Vec<usize>>,
merge_sim: f32,
) -> Vec<Vec<usize>> {
let mut centroids: Vec<Vec<f32>> = clusters.iter().map(|c| centroid(points, c)).collect();
loop {
let mut best: Option<(f32, usize, usize)> = None;
for i in 0..clusters.len() {
for j in (i + 1)..clusters.len() {
let s = centroids[i].iter().zip(¢roids[j]).map(|(a, b)| a * b).sum::<f32>();
if best.is_none_or(|(bs, _, _)| s > bs) {
best = Some((s, i, j));
}
}
}
let Some((s, i, j)) = best else { break };
if s < merge_sim { break; }
let moved = std::mem::take(&mut clusters[j]);
clusters[i].extend(moved);
clusters.swap_remove(j);
centroids.swap_remove(j);
centroids[i] = centroid(points, &clusters[i]);
}
clusters
}
let points = same_identity_two_subclusters();
let initial_clusters: Vec<Vec<usize>> = vec![vec![0, 1, 2], vec![3, 4, 5]];
let via_incremental = merge_by_centroid(&points, initial_clusters.clone(), 0.4);
let via_reference = reference_merge_by_centroid(&points, initial_clusters, 0.4);
let mut incremental_sorted: Vec<Vec<usize>> =
via_incremental.into_iter().map(|mut c| { c.sort(); c }).collect();
let mut reference_sorted: Vec<Vec<usize>> =
via_reference.into_iter().map(|mut c| { c.sort(); c }).collect();
incremental_sorted.sort();
reference_sorted.sort();
assert_eq!(incremental_sorted, reference_sorted);
}
#[test]
fn centroid_merge_keeps_distinct_identities_apart() {
let mut pts = same_identity_two_subclusters();
pts.truncate(3); pts.push((7, l2(vec![-1.0, 0.0, 0.0, 0.15, 0.0, 0.0])));
pts.push((8, l2(vec![-1.0, 0.0, 0.0, 0.0, 0.15, 0.0])));
pts.push((9, l2(vec![-1.0, 0.0, 0.0, 0.0, 0.0, 0.15])));
let result = cluster_faces(&pts, 0.3, 1, 0.4, true);
let map: std::collections::HashMap<_, _> = result.into_iter().collect();
assert_eq!(map[&1], map[&2], "cluster A stays together");
assert_eq!(map[&7], map[&8], "cluster C stays together");
assert_ne!(map[&1], map[&7], "different identities must not merge");
}
#[test]
fn centroid_merge_still_drops_small_clusters_below_min_samples() {
let mut pts = same_identity_two_subclusters();
pts.push((99, l2(vec![0.0, 0.0, 0.0, 0.0, 1.0, 0.0])));
let result = cluster_faces(&pts, 0.3, 2, 0.4, true);
let map: std::collections::HashMap<_, _> = result.into_iter().collect();
assert_eq!(map[&99], None, "isolated singleton must remain noise");
}
#[test]
fn agglomerate_average_matches_dense_matrix_reference_on_synthetic_data() {
fn reference_agglomerate_average(
points: &[(i64, Vec<f32>)],
eps: f32,
) -> Vec<Vec<usize>> {
let n = points.len();
let mut dist: Vec<Vec<f32>> = vec![vec![0.0f32; n]; n];
let mut heap: BinaryHeap<HeapEntry> = BinaryHeap::new();
for i in 0..n {
for j in (i + 1)..n {
let d = cosine_dist(&points[i].1, &points[j].1);
dist[i][j] = d;
dist[j][i] = d;
if d <= eps {
heap.push(HeapEntry { dist: d, i, j });
}
}
}
let mut members: Vec<Vec<usize>> = (0..n).map(|i| vec![i]).collect();
let mut alive = vec![true; n];
while let Some(HeapEntry { dist: d, i, j }) = heap.pop() {
if !alive[i] || !alive[j] {
continue;
}
if dist[i][j] != d {
continue;
}
if d > eps {
break;
}
let size_i = members[i].len() as f32;
let size_j = members[j].len() as f32;
let moved = std::mem::take(&mut members[j]);
members[i].extend(moved);
alive[j] = false;
for k in 0..n {
if k == i || k == j || !alive[k] {
continue;
}
let new_d = (size_i * dist[i][k] + size_j * dist[j][k]) / (size_i + size_j);
if new_d != dist[i][k] {
dist[i][k] = new_d;
dist[k][i] = new_d;
heap.push(HeapEntry { dist: new_d, i: i.min(k), j: i.max(k) });
}
}
}
(0..n).filter(|&r| alive[r]).map(|r| std::mem::take(&mut members[r])).collect()
}
fn lcg_next(state: &mut u64) -> f32 {
*state = state.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
((*state >> 33) as f32 / (1u64 << 31) as f32) - 1.0 }
fn normalize(v: &mut [f32]) {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 1e-12 {
for x in v.iter_mut() {
*x /= norm;
}
}
}
let dim = 64;
let num_identities = 8;
let points_per_identity = 38; let mut state = 0xC0FFEEu64;
let mut centers: Vec<Vec<f32>> = Vec::new();
for _ in 0..num_identities {
let mut c: Vec<f32> = (0..dim).map(|_| lcg_next(&mut state)).collect();
normalize(&mut c);
centers.push(c);
}
let mut points: Vec<(i64, Vec<f32>)> = Vec::new();
let mut next_id = 1i64;
for center in ¢ers {
for _ in 0..points_per_identity {
let mut v: Vec<f32> = center
.iter()
.map(|c| c + 0.15 * lcg_next(&mut state))
.collect();
normalize(&mut v);
points.push((next_id, v));
next_id += 1;
}
}
for &eps in &[0.3f32, 0.6f32, 0.9f32] {
let via_new = agglomerate_average(&points, eps, true);
let via_reference = reference_agglomerate_average(&points, eps);
let mut new_sorted: Vec<Vec<usize>> =
via_new.into_iter().map(|mut c| { c.sort(); c }).collect();
let mut reference_sorted: Vec<Vec<usize>> =
via_reference.into_iter().map(|mut c| { c.sort(); c }).collect();
new_sorted.sort();
reference_sorted.sort();
assert_eq!(
new_sorted, reference_sorted,
"partition mismatch at eps={eps} between new and reference agglomerate_average"
);
}
}
}