use crate::cluster::kmeans::KMeansInit::KMeansPlusPlus;
use crate::error::{DatarustError, Result};
use crate::matrix::Matrix;
use crate::model_selection::rng::Rng;
use crate::traits::{Clusterer, ParamValue, Params};
type LloydResult = (Vec<Vec<f64>>, Vec<usize>, f64, usize);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum KMeansInit {
#[default]
KMeansPlusPlus,
Random,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct KMeans {
n_clusters: usize,
init: KMeansInit,
max_iter: usize,
tol: f64,
n_init: usize,
random_state: Option<u64>,
cluster_centers_: Vec<Vec<f64>>,
labels_: Vec<usize>,
inertia_: f64,
n_iter_: usize,
n_features_in_: usize,
fitted: bool,
}
impl Default for KMeans {
fn default() -> Self {
Self::new()
}
}
impl KMeans {
fn validate_fitted_state(&self) -> Result<()> {
if self.cluster_centers_.len() != self.n_clusters
|| self.n_features_in_ == 0
|| !self.inertia_.is_finite()
|| self
.cluster_centers_
.iter()
.any(|center| center.len() != self.n_features_in_)
|| self
.cluster_centers_
.iter()
.flatten()
.any(|v| !v.is_finite())
{
return Err(DatarustError::InvalidInput(
"KMeans has inconsistent fitted state".into(),
));
}
Ok(())
}
pub fn new() -> Self {
Self {
n_clusters: 8,
init: KMeansPlusPlus,
max_iter: 300,
tol: 1e-4,
n_init: 10,
random_state: None,
cluster_centers_: Vec::new(),
labels_: Vec::new(),
inertia_: 0.0,
n_iter_: 0,
n_features_in_: 0,
fitted: false,
}
}
pub fn with_n_clusters(mut self, n_clusters: usize) -> Self {
self.n_clusters = n_clusters;
self
}
pub fn with_init(mut self, init: KMeansInit) -> Self {
self.init = init;
self
}
pub fn with_max_iter(mut self, max_iter: usize) -> Self {
self.max_iter = max_iter;
self
}
pub fn with_tol(mut self, tol: f64) -> Self {
self.tol = tol;
self
}
pub fn with_n_init(mut self, n_init: usize) -> Self {
self.n_init = n_init;
self
}
pub fn with_random_state(mut self, seed: u64) -> Self {
self.random_state = Some(seed);
self
}
pub fn cluster_centers(&self) -> &[Vec<f64>] {
&self.cluster_centers_
}
pub fn labels(&self) -> &[usize] {
&self.labels_
}
pub fn inertia(&self) -> f64 {
self.inertia_
}
pub fn n_iter(&self) -> usize {
self.n_iter_
}
pub fn n_features_in(&self) -> usize {
self.n_features_in_
}
fn validate(&self, x: &Matrix) -> Result<(usize, usize)> {
let n = x.nrows();
let p = x.ncols();
if n == 0 {
return Err(DatarustError::EmptyInput("X has no rows".into()));
}
if p == 0 {
return Err(DatarustError::EmptyInput("X has no columns".into()));
}
if self.n_clusters == 0 {
return Err(DatarustError::InvalidConfig(
"n_clusters must be >= 1".into(),
));
}
x.validate_finite()?;
if self.n_clusters > n {
return Err(DatarustError::InvalidConfig(format!(
"n_clusters ({}) cannot be greater than n_samples ({})",
self.n_clusters, n
)));
}
if self.max_iter == 0 {
return Err(DatarustError::InvalidConfig("max_iter must be > 0".into()));
}
if self.n_init == 0 {
return Err(DatarustError::InvalidConfig("n_init must be > 0".into()));
}
if !self.tol.is_finite() || self.tol < 0.0 {
return Err(DatarustError::InvalidConfig(format!(
"tol must be finite and >= 0, got {}",
self.tol
)));
}
Ok((n, p))
}
#[inline]
fn sq_dist(a: &[f64], b: &[f64]) -> f64 {
a.iter()
.zip(b.iter())
.map(|(ai, bi)| {
let d = ai - bi;
d * d
})
.sum()
}
fn init_centroids(&self, x: &Matrix, n: usize, p: usize, rng: &mut Rng) -> Vec<Vec<f64>> {
match self.init {
KMeansInit::Random => {
let mut indices: Vec<usize> = (0..n).collect();
rng.shuffle(&mut indices);
indices[..self.n_clusters]
.iter()
.map(|&i| x.row(i).to_vec())
.collect()
}
KMeansInit::KMeansPlusPlus => {
let first = rng.next_usize(n);
let mut centers: Vec<Vec<f64>> = Vec::with_capacity(self.n_clusters);
centers.push(x.row(first).to_vec());
let mut nearest_sq: Vec<f64> = (0..n)
.map(|i| Self::sq_dist(x.row(i), ¢ers[0]))
.collect();
for _ in 1..self.n_clusters {
let total: f64 = nearest_sq.iter().sum();
let mut centers_row = vec![0.0_f64; p];
if total <= 0.0 {
let next = centers.len().min(n - 1);
centers_row.copy_from_slice(x.row(next));
} else {
let r = rng.next_unit() * total;
let mut acc = 0.0;
let mut chosen = n - 1;
for (i, &d) in nearest_sq.iter().enumerate() {
acc += d;
if acc >= r {
chosen = i;
break;
}
}
centers_row.copy_from_slice(x.row(chosen));
}
for (i, nearest) in nearest_sq.iter_mut().enumerate() {
let d = Self::sq_dist(x.row(i), ¢ers_row);
if d < *nearest {
*nearest = d;
}
}
centers.push(centers_row);
}
centers
}
}
}
fn lloyds_run(
&self,
x: &Matrix,
n: usize,
p: usize,
mut centers: Vec<Vec<f64>>,
) -> LloydResult {
let mut labels = vec![0usize; n];
let mut centroid_shift_sq = f64::MAX;
let mut iter = 0;
while iter < self.max_iter && centroid_shift_sq > self.tol {
for (label, i) in labels.iter_mut().zip(0..n) {
let row = x.row(i);
let (best_idx, _) = centers
.iter()
.enumerate()
.map(|(c, ctr)| (c, Self::sq_dist(row, ctr)))
.min_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.unwrap_or((0, 0.0));
*label = best_idx;
}
let mut new_centers = vec![vec![0.0_f64; p]; self.n_clusters];
let mut counts = vec![0usize; self.n_clusters];
for (i, &label) in labels.iter().enumerate() {
counts[label] += 1;
for (j, center) in new_centers[label].iter_mut().enumerate() {
*center += x.get(i, j);
}
}
for (c, center) in new_centers.iter_mut().enumerate() {
if counts[c] > 0 {
for val in center.iter_mut() {
*val /= counts[c] as f64;
}
} else {
center.copy_from_slice(¢ers[c]);
}
}
centroid_shift_sq = new_centers
.iter()
.zip(centers.iter())
.map(|(new, old)| Self::sq_dist(new, old))
.sum::<f64>()
/ self.n_clusters as f64;
centers = new_centers;
iter += 1;
}
let inertia = (0..n)
.map(|i| Self::sq_dist(x.row(i), ¢ers[labels[i]]))
.sum::<f64>();
(centers, labels, inertia, iter)
}
}
impl Params for KMeans {
fn get_params(&self) -> Vec<(&'static str, ParamValue)> {
vec![
("n_clusters", ParamValue::Int(self.n_clusters)),
("max_iter", ParamValue::Int(self.max_iter)),
("tol", ParamValue::Float(self.tol)),
("n_init", ParamValue::Int(self.n_init)),
]
}
fn set_params(&mut self, name: &str, value: ParamValue) -> Result<()> {
match (name, value) {
("n_clusters", ParamValue::Int(0)) => {
return Err(DatarustError::InvalidConfig(
"n_clusters must be >= 1".into(),
));
}
("n_clusters", ParamValue::Int(v)) => self.n_clusters = v,
("max_iter", ParamValue::Int(0)) => {
return Err(DatarustError::InvalidConfig("max_iter must be > 0".into()));
}
("max_iter", ParamValue::Int(v)) => self.max_iter = v,
("tol", ParamValue::Float(v)) if !v.is_finite() || v < 0.0 => {
return Err(DatarustError::InvalidConfig(format!(
"tol must be finite and >= 0, got {v}"
)));
}
("tol", ParamValue::Float(v)) => self.tol = v,
("n_init", ParamValue::Int(0)) => {
return Err(DatarustError::InvalidConfig("n_init must be > 0".into()));
}
("n_init", ParamValue::Int(v)) => self.n_init = v,
(other, _) => {
return Err(DatarustError::InvalidInput(format!(
"KMeans has no tunable parameter '{other}'"
)));
}
}
self.fitted = false;
Ok(())
}
}
impl Clusterer for KMeans {
fn name(&self) -> &'static str {
"KMeans"
}
fn fit(&mut self, x: &Matrix) -> Result<()> {
let (n, p) = self.validate(x)?;
self.n_features_in_ = p;
let base_seed = self.random_state.unwrap_or(0x9E3779B97F4A7C15);
let mut best: Option<LloydResult> = None;
for run in 0..self.n_init {
let mut rng = Rng::new(base_seed.wrapping_add(run as u64));
let init_centers = self.init_centroids(x, n, p, &mut rng);
let (centers, labels, inertia, iters) = self.lloyds_run(x, n, p, init_centers);
match &best {
None => best = Some((centers, labels, inertia, iters)),
Some((_, _, best_inertia, _)) if inertia < *best_inertia => {
best = Some((centers, labels, inertia, iters));
}
_ => {}
}
}
let (centers, labels, inertia, iters) =
best.expect("n_init >= 1 checked in validate; best is always Some");
self.cluster_centers_ = centers;
self.labels_ = labels;
self.inertia_ = inertia;
self.n_iter_ = iters;
self.fitted = true;
Ok(())
}
fn predict(&self, x: &Matrix) -> Result<Vec<usize>> {
if !self.fitted {
return Err(DatarustError::NotFitted("KMeans".into()));
}
self.validate_fitted_state()?;
if x.ncols() != self.n_features_in_ {
return Err(DatarustError::ShapeMismatch {
expected: format!("{} features", self.n_features_in_),
actual: format!("{} features", x.ncols()),
});
}
x.validate_finite()?;
let n = x.nrows();
let mut out = vec![0usize; n];
for (i, slot) in out.iter_mut().enumerate() {
let row = x.row(i);
let (best_idx, _) = self
.cluster_centers_
.iter()
.enumerate()
.map(|(c, ctr)| (c, Self::sq_dist(row, ctr)))
.min_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.unwrap_or((0, 0.0));
*slot = best_idx;
}
Ok(out)
}
fn fit_predict(&mut self, x: &Matrix) -> Result<Vec<usize>> {
self.fit(x)?;
Ok(self.labels_.clone())
}
fn n_clusters(&self) -> usize {
self.n_clusters
}
fn is_fitted(&self) -> bool {
self.fitted
}
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn finds_three_blobs() {
let rows: Vec<Vec<f64>> = (0..30)
.map(|i| {
let base = (i / 10) as f64 * 10.0;
vec![base, base]
})
.collect();
let x = Matrix::new(rows).unwrap();
let mut km = KMeans::new()
.with_n_clusters(3)
.with_n_init(10)
.with_random_state(0);
let labels = km.fit_predict(&x).unwrap();
for blob in 0..3 {
let first = labels[blob * 10];
for i in 1..10 {
assert_eq!(labels[blob * 10 + i], first, "blob {blob} not homogeneous");
}
}
assert_ne!(labels[0], labels[10]);
assert_ne!(labels[10], labels[20]);
assert_ne!(labels[0], labels[20]);
}
#[test]
fn recovers_known_centroids() {
let mut rows: Vec<Vec<f64>> = Vec::new();
let centers = [[0.0, 0.0], [10.0, 10.0], [-10.0, -10.0]];
for [cx, cy] in centers {
for dx in [-0.1, 0.0, 0.1] {
for dy in [-0.1, 0.0, 0.1] {
rows.push(vec![cx + dx, cy + dy]);
}
}
}
let x = Matrix::new(rows).unwrap();
let mut km = KMeans::new().with_n_clusters(3).with_random_state(42);
km.fit(&x).unwrap();
let mut matched = 0;
for center in km.cluster_centers() {
let close = centers
.iter()
.any(|tc| approx(center[0], tc[0], 0.2) && approx(center[1], tc[1], 0.2));
if close {
matched += 1;
}
}
assert_eq!(
matched,
3,
"not all centroids recovered: {:?}",
km.cluster_centers()
);
}
#[test]
fn inertia_non_negative() {
let rows = vec![vec![1.0], vec![2.0], vec![10.0], vec![11.0]];
let x = Matrix::new(rows).unwrap();
let mut km = KMeans::new().with_n_clusters(2).with_random_state(1);
km.fit(&x).unwrap();
assert!(km.inertia() >= 0.0);
}
#[test]
fn predict_assigns_new_points() {
let rows = vec![vec![0.0], vec![1.0], vec![100.0], vec![101.0]];
let x = Matrix::new(rows).unwrap();
let mut km = KMeans::new().with_n_clusters(2).with_random_state(0);
km.fit(&x).unwrap();
let train_labels = km.labels();
let left_cluster = train_labels[0]; let right_cluster = train_labels[2]; let test = Matrix::new(vec![vec![0.5], vec![100.5]]).unwrap();
let pred = km.predict(&test).unwrap();
assert_eq!(pred[0], left_cluster);
assert_eq!(pred[1], right_cluster);
assert_ne!(pred[0], pred[1]);
}
#[test]
fn deterministic_same_seed() {
let rows: Vec<Vec<f64>> = (0..20).map(|i| vec![i as f64, i as f64 * 2.0]).collect();
let x = Matrix::new(rows).unwrap();
let mut a = KMeans::new().with_n_clusters(3).with_random_state(7);
let mut b = KMeans::new().with_n_clusters(3).with_random_state(7);
let la = a.fit_predict(&x).unwrap();
let lb = b.fit_predict(&x).unwrap();
assert_eq!(la, lb);
assert_eq!(a.cluster_centers(), b.cluster_centers());
}
#[test]
fn n_clusters_zero_errors() {
let x = Matrix::new(vec![vec![1.0], vec![2.0]]).unwrap();
let mut km = KMeans::new().with_n_clusters(0);
assert!(km.fit(&x).is_err());
}
#[test]
fn n_clusters_exceeds_samples_errors() {
let x = Matrix::new(vec![vec![1.0], vec![2.0]]).unwrap();
let mut km = KMeans::new().with_n_clusters(5);
assert!(km.fit(&x).is_err());
}
#[test]
fn invalid_iteration_and_tolerance_configuration_errors() {
let x = Matrix::new(vec![vec![1.0], vec![2.0]]).unwrap();
assert!(KMeans::new()
.with_n_clusters(2)
.with_max_iter(0)
.fit(&x)
.is_err());
assert!(KMeans::new()
.with_n_clusters(2)
.with_n_init(0)
.fit(&x)
.is_err());
for tol in [-1.0, f64::NAN, f64::INFINITY] {
assert!(KMeans::new()
.with_n_clusters(2)
.with_tol(tol)
.fit(&x)
.is_err());
}
}
#[test]
fn non_finite_features_error_during_fit_and_predict() {
let train = Matrix::new(vec![vec![0.0], vec![1.0]]).unwrap();
for invalid in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
let bad = Matrix::new(vec![vec![0.0], vec![invalid]]).unwrap();
assert!(KMeans::new().with_n_clusters(1).fit(&bad).is_err());
let mut fitted = KMeans::new().with_n_clusters(1);
fitted.fit(&train).unwrap();
assert!(fitted.predict(&bad).is_err());
}
}
#[test]
fn params_reject_invalid_values_without_mutating_state() {
let mut model = KMeans::new();
for (name, value) in [
("n_clusters", ParamValue::Int(0)),
("max_iter", ParamValue::Int(0)),
("tol", ParamValue::Float(f64::NAN)),
("n_init", ParamValue::Int(0)),
] {
assert!(matches!(
model.set_params(name, value),
Err(DatarustError::InvalidConfig(_))
));
}
assert_eq!(
model.get_params(),
vec![
("n_clusters", ParamValue::Int(8)),
("max_iter", ParamValue::Int(300)),
("tol", ParamValue::Float(1e-4)),
("n_init", ParamValue::Int(10)),
]
);
}
#[test]
fn predict_before_fit_errors() {
let km = KMeans::new().with_n_clusters(2);
let x = Matrix::new(vec![vec![1.0], vec![2.0]]).unwrap();
assert!(km.predict(&x).is_err());
}
#[test]
fn single_sample_single_cluster_works() {
let x = Matrix::new(vec![vec![1.0]]).unwrap();
let mut km = KMeans::new().with_n_clusters(1);
assert!(km.fit(&x).is_ok());
assert_eq!(km.labels(), &[0]);
}
#[test]
fn n_iter_recorded() {
let rows = vec![vec![0.0], vec![1.0], vec![10.0], vec![11.0]];
let x = Matrix::new(rows).unwrap();
let mut km = KMeans::new().with_n_clusters(2).with_random_state(0);
km.fit(&x).unwrap();
assert!(km.n_iter() >= 1);
}
#[test]
fn kmeans_plus_plus_better_than_or_equal_random() {
let mut rows: Vec<Vec<f64>> = Vec::new();
for center in &[0.0_f64, 10.0, 20.0] {
for d in &[-0.2, -0.1, 0.0, 0.1, 0.2] {
rows.push(vec![center + d]);
}
}
let x = Matrix::new(rows).unwrap();
let mut pp = KMeans::new()
.with_n_clusters(3)
.with_init(KMeansInit::KMeansPlusPlus)
.with_random_state(0);
pp.fit(&x).unwrap();
let mut rnd = KMeans::new()
.with_n_clusters(3)
.with_init(KMeansInit::Random)
.with_random_state(0);
rnd.fit(&x).unwrap();
assert_eq!(
pp.labels()
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>()
.len(),
3
);
assert!(pp.inertia() <= rnd.inertia() + 1e-9 || pp.inertia() < 1e-6);
}
#[test]
fn predict_shape_mismatch_errors() {
let x = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
let mut km = KMeans::new().with_n_clusters(2).with_random_state(0);
km.fit(&x).unwrap();
let wrong = Matrix::new(vec![vec![1.0], vec![2.0]]).unwrap();
assert!(km.predict(&wrong).is_err());
}
#[test]
#[allow(clippy::float_cmp)]
fn fit_transform_returns_one_hot() {
let rows = vec![vec![0.0], vec![1.0], vec![10.0], vec![11.0]];
let x = Matrix::new(rows).unwrap();
let mut km = KMeans::new().with_n_clusters(2).with_random_state(0);
let out = km.fit_transform(&x).unwrap();
assert_eq!(out.nrows(), 4);
assert_eq!(out.ncols(), 2);
for i in 0..4 {
let sum: f64 = (0..2).map(|j| out.get(i, j)).sum();
assert!(approx(sum, 1.0, 1e-12));
}
}
}