stats_claw/algorithms/clustering/
mean_shift.rs1use crate::algorithms::{centroid, euclidean_sq};
11
12#[derive(Debug, Clone)]
14pub struct MeanShiftResult {
15 pub labels: Vec<usize>,
17 pub n_clusters: usize,
19 pub centers: Vec<Vec<f64>>,
21}
22
23const MAX_ITER: usize = 300;
25const SHIFT_TOL: f64 = 1e-3;
27
28#[must_use]
54pub fn mean_shift(data: &[Vec<f64>], bandwidth: f64) -> MeanShiftResult {
55 if data.is_empty() || bandwidth <= 0.0 {
56 return MeanShiftResult {
57 labels: Vec::new(),
58 n_clusters: 0,
59 centers: Vec::new(),
60 };
61 }
62 let band_sq = bandwidth * bandwidth;
63 let modes: Vec<Vec<f64>> = data
64 .iter()
65 .map(|seed| shift_to_mode(seed, data, band_sq))
66 .collect();
67 let centers = merge_modes(&modes, band_sq);
68 let labels: Vec<usize> = data
69 .iter()
70 .map(|point| nearest_center(point, ¢ers))
71 .collect();
72 let n_clusters = centers.len();
73 MeanShiftResult {
74 labels,
75 n_clusters,
76 centers,
77 }
78}
79
80fn shift_to_mode(seed: &[f64], data: &[Vec<f64>], band_sq: f64) -> Vec<f64> {
83 let dim = seed.len();
84 let mut current = seed.to_vec();
85 let tol_sq = band_sq * SHIFT_TOL * SHIFT_TOL;
86 for _ in 0..MAX_ITER {
87 let within: Vec<&[f64]> = data
88 .iter()
89 .filter(|p| euclidean_sq(¤t, p) <= band_sq)
90 .map(Vec::as_slice)
91 .collect();
92 if within.is_empty() {
93 break;
94 }
95 let next = centroid(&within, dim);
96 let moved = euclidean_sq(¤t, &next);
97 current = next;
98 if moved <= tol_sq {
99 break;
100 }
101 }
102 current
103}
104
105fn merge_modes(modes: &[Vec<f64>], band_sq: f64) -> Vec<Vec<f64>> {
108 let mut centers: Vec<Vec<f64>> = Vec::new();
109 for mode in modes {
110 let near = centers.iter().any(|c| euclidean_sq(mode, c) <= band_sq);
111 if !near {
112 centers.push(mode.clone());
113 }
114 }
115 centers
116}
117
118fn nearest_center(point: &[f64], centers: &[Vec<f64>]) -> usize {
120 let mut best = 0_usize;
121 let mut best_d = f64::INFINITY;
122 for (i, c) in centers.iter().enumerate() {
123 let d = euclidean_sq(point, c);
124 if d < best_d {
125 best_d = d;
126 best = i;
127 }
128 }
129 best
130}
131
132#[cfg(test)]
133mod tests {
134 use super::*;
135
136 #[test]
137 fn finds_two_modes_in_two_blobs() {
138 let data = vec![
139 vec![0.0],
140 vec![0.1],
141 vec![0.2],
142 vec![10.0],
143 vec![10.1],
144 vec![10.2],
145 ];
146 let r = mean_shift(&data, 1.0);
147 assert_eq!(r.n_clusters, 2, "modes were {}", r.n_clusters);
148 assert_eq!(r.labels.first(), r.labels.get(2), "first blob split");
149 assert_ne!(r.labels.first(), r.labels.get(3), "blobs merged");
150 }
151
152 #[test]
153 fn deterministic_for_fixed_inputs() {
154 let data = vec![vec![0.0], vec![0.1], vec![5.0], vec![5.1]];
155 assert_eq!(mean_shift(&data, 1.0).labels, mean_shift(&data, 1.0).labels);
156 }
157
158 #[test]
159 fn empty_input_is_empty_result() {
160 let r = mean_shift(&[], 1.0);
161 assert!(r.labels.is_empty(), "labels not empty");
162 assert_eq!(r.n_clusters, 0, "clusters not zero");
163 }
164}