use crate::algorithms::{centroid, euclidean_sq};
#[derive(Debug, Clone)]
pub struct MeanShiftResult {
pub labels: Vec<usize>,
pub n_clusters: usize,
pub centers: Vec<Vec<f64>>,
}
const MAX_ITER: usize = 300;
const SHIFT_TOL: f64 = 1e-3;
#[must_use]
pub fn mean_shift(data: &[Vec<f64>], bandwidth: f64) -> MeanShiftResult {
if data.is_empty() || bandwidth <= 0.0 {
return MeanShiftResult {
labels: Vec::new(),
n_clusters: 0,
centers: Vec::new(),
};
}
let band_sq = bandwidth * bandwidth;
let modes: Vec<Vec<f64>> = data
.iter()
.map(|seed| shift_to_mode(seed, data, band_sq))
.collect();
let centers = merge_modes(&modes, band_sq);
let labels: Vec<usize> = data
.iter()
.map(|point| nearest_center(point, ¢ers))
.collect();
let n_clusters = centers.len();
MeanShiftResult {
labels,
n_clusters,
centers,
}
}
fn shift_to_mode(seed: &[f64], data: &[Vec<f64>], band_sq: f64) -> Vec<f64> {
let dim = seed.len();
let mut current = seed.to_vec();
let tol_sq = band_sq * SHIFT_TOL * SHIFT_TOL;
for _ in 0..MAX_ITER {
let within: Vec<&[f64]> = data
.iter()
.filter(|p| euclidean_sq(¤t, p) <= band_sq)
.map(Vec::as_slice)
.collect();
if within.is_empty() {
break;
}
let next = centroid(&within, dim);
let moved = euclidean_sq(¤t, &next);
current = next;
if moved <= tol_sq {
break;
}
}
current
}
fn merge_modes(modes: &[Vec<f64>], band_sq: f64) -> Vec<Vec<f64>> {
let mut centers: Vec<Vec<f64>> = Vec::new();
for mode in modes {
let near = centers.iter().any(|c| euclidean_sq(mode, c) <= band_sq);
if !near {
centers.push(mode.clone());
}
}
centers
}
fn nearest_center(point: &[f64], centers: &[Vec<f64>]) -> usize {
let mut best = 0_usize;
let mut best_d = f64::INFINITY;
for (i, c) in centers.iter().enumerate() {
let d = euclidean_sq(point, c);
if d < best_d {
best_d = d;
best = i;
}
}
best
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn finds_two_modes_in_two_blobs() {
let data = vec![
vec![0.0],
vec![0.1],
vec![0.2],
vec![10.0],
vec![10.1],
vec![10.2],
];
let r = mean_shift(&data, 1.0);
assert_eq!(r.n_clusters, 2, "modes were {}", r.n_clusters);
assert_eq!(r.labels.first(), r.labels.get(2), "first blob split");
assert_ne!(r.labels.first(), r.labels.get(3), "blobs merged");
}
#[test]
fn deterministic_for_fixed_inputs() {
let data = vec![vec![0.0], vec![0.1], vec![5.0], vec![5.1]];
assert_eq!(mean_shift(&data, 1.0).labels, mean_shift(&data, 1.0).labels);
}
#[test]
fn empty_input_is_empty_result() {
let r = mean_shift(&[], 1.0);
assert!(r.labels.is_empty(), "labels not empty");
assert_eq!(r.n_clusters, 0, "clusters not zero");
}
}