Skip to main content

stats_claw/algorithms/clustering/
mean_shift.rs

1//! Mean-shift clustering with a flat (uniform) kernel.
2//!
3//! Each point is treated as a seed and iteratively shifted to the mean of the
4//! points inside its `bandwidth` window until the shift falls below a tolerance.
5//! Converged seeds that land within one bandwidth of each other collapse to a
6//! single mode; every original point then takes the label of its nearest mode.
7//! This mirrors `scikit-learn`'s `MeanShift` with a flat kernel, so the number of
8//! discovered modes is the identifiable scalar the equivalence suite checks.
9
10use crate::algorithms::{centroid, euclidean_sq};
11
12/// Outcome of a mean-shift run.
13#[derive(Debug, Clone)]
14pub struct MeanShiftResult {
15    /// Cluster id per input point, in input order.
16    pub labels: Vec<usize>,
17    /// Number of distinct modes (clusters) discovered.
18    pub n_clusters: usize,
19    /// The discovered cluster centers (modes), one per cluster.
20    pub centers: Vec<Vec<f64>>,
21}
22
23/// Maximum mean-shift iterations per seed before giving up on convergence.
24const MAX_ITER: usize = 300;
25/// Relative shift (as a fraction of bandwidth) below which a seed has converged.
26const SHIFT_TOL: f64 = 1e-3;
27
28/// Clusters `data` by shifting every point toward its local density mode.
29///
30/// Deterministic — no RNG is used; seeds are processed in input order and modes
31/// are merged greedily, so repeated runs match. Empty input or a non-positive
32/// bandwidth yields an empty result.
33///
34/// # Arguments
35///
36/// * `data` — observations; each inner slice is one point of equal dimension.
37/// * `bandwidth` — radius of the flat kernel window; must be `> 0`.
38///
39/// # Returns
40///
41/// A [`MeanShiftResult`] with per-point labels, the discovered mode count, and the
42/// mode centers.
43///
44/// # Examples
45///
46/// ```
47/// use stats_claw::algorithms::clustering::mean_shift;
48///
49/// let data = vec![vec![0.0], vec![0.1], vec![10.0], vec![10.1]];
50/// let r = mean_shift(&data, 1.0);
51/// assert_eq!(r.n_clusters, 2, "modes were {}", r.n_clusters);
52/// ```
53#[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, &centers))
71        .collect();
72    let n_clusters = centers.len();
73    MeanShiftResult {
74        labels,
75        n_clusters,
76        centers,
77    }
78}
79
80/// Iteratively shifts `seed` to the mean of points within the flat kernel until
81/// the move shrinks below the tolerance or the iteration cap is hit.
82fn 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(&current, 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(&current, &next);
97        current = next;
98        if moved <= tol_sq {
99            break;
100        }
101    }
102    current
103}
104
105/// Collapses converged modes into representative centers: a mode joins an existing
106/// center when it lies within one bandwidth, otherwise it starts a new center.
107fn 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
118/// Returns the index of the center nearest to `point` (0 when none exist).
119fn 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}