use crate::common::assert_allclose;
use approx::assert_abs_diff_eq;
use ndarray::{Array1, Array2, array};
use rustyml::error::Error;
use rustyml::machine_learning::KMeans;
fn three_blob_data() -> Array2<f64> {
#[rustfmt::skip]
let data = array![
[-0.05, 0.03],
[ 0.04, -0.02],
[ 0.01, 0.05],
[-0.03, -0.04],
[ 0.02, 0.01],
[ 9.95, 0.03],
[10.04, -0.02],
[10.01, 0.05],
[ 9.97, -0.04],
[10.02, 0.01],
[ 4.95, 10.03],
[ 5.04, 9.98],
[ 5.01, 10.05],
[ 4.97, 9.96],
[ 5.02, 10.01],
];
data
}
fn assert_blob_structure(labels: &Array1<usize>, blob_size: usize) {
let n_blobs = labels.len() / blob_size;
assert_eq!(
labels.len(),
n_blobs * blob_size,
"labels length must be a multiple of blob_size"
);
let blob_labels: Vec<usize> = (0..n_blobs).map(|b| labels[b * blob_size]).collect();
for b in 0..n_blobs {
for i in 0..blob_size {
assert_eq!(
labels[b * blob_size + i],
blob_labels[b],
"blob {b} point {i} has label {} but expected {}",
labels[b * blob_size + i],
blob_labels[b],
);
}
}
for i in 0..n_blobs {
for j in (i + 1)..n_blobs {
assert_ne!(
blob_labels[i], blob_labels[j],
"blobs {i} and {j} share label {}",
blob_labels[i],
);
}
}
}
#[test]
fn constructor_zero_clusters_is_invalid() {
let err = KMeans::new(0, 100, 1e-4).unwrap_err();
assert!(
matches!(err, Error::InvalidParameter { .. }),
"expected InvalidParameter, got {err:?}"
);
}
#[test]
fn constructor_zero_max_iter_is_invalid() {
let err = KMeans::new(3, 0, 1e-4).unwrap_err();
assert!(
matches!(err, Error::InvalidParameter { .. }),
"expected InvalidParameter, got {err:?}"
);
}
#[test]
fn constructor_zero_tolerance_is_invalid() {
let err = KMeans::new(3, 100, 0.0).unwrap_err();
assert!(
matches!(err, Error::InvalidParameter { .. }),
"expected InvalidParameter, got {err:?}"
);
}
#[test]
fn constructor_negative_tolerance_is_invalid() {
let err = KMeans::new(3, 100, -1e-4).unwrap_err();
assert!(
matches!(err, Error::InvalidParameter { .. }),
"expected InvalidParameter, got {err:?}"
);
}
#[test]
fn constructor_nan_tolerance_is_invalid() {
let err = KMeans::new(3, 100, f64::NAN).unwrap_err();
assert!(
matches!(err, Error::InvalidParameter { .. }),
"expected InvalidParameter, got {err:?}"
);
}
#[test]
fn constructor_inf_tolerance_is_invalid() {
let err = KMeans::new(3, 100, f64::INFINITY).unwrap_err();
assert!(
matches!(err, Error::InvalidParameter { .. }),
"expected InvalidParameter, got {err:?}"
);
}
#[test]
fn constructor_valid_stores_params() {
let km = KMeans::new(3, 200, 1e-4).unwrap().with_random_state(42);
assert_eq!(km.get_n_clusters(), 3);
assert_eq!(km.get_max_iterations(), 200);
assert_abs_diff_eq!(km.get_tolerance(), 1e-4, epsilon = 1e-15);
assert_eq!(km.get_random_state(), Some(42));
assert!(km.get_centroids().is_none());
assert!(km.get_labels().is_none());
assert!(km.get_inertia().is_none());
assert!(km.get_actual_iterations().is_none());
}
#[test]
fn default_constructor_has_documented_defaults() {
let km = KMeans::default();
assert_eq!(km.get_n_clusters(), 8);
assert_eq!(km.get_max_iterations(), 300);
assert_abs_diff_eq!(km.get_tolerance(), 1e-4, epsilon = 1e-15);
assert_eq!(km.get_random_state(), None);
}
#[test]
fn fit_empty_data_is_error() {
let mut km = KMeans::new(1, 100, 1e-4).unwrap().with_random_state(0);
let data: Array2<f64> = Array2::zeros((0, 2));
let err = km.fit(&data).unwrap_err();
assert!(
matches!(err, Error::EmptyInput(_)),
"expected EmptyInput, got {err:?}"
);
}
#[test]
fn fit_nan_data_is_error() {
let mut km = KMeans::new(1, 100, 1e-4).unwrap().with_random_state(0);
let data = array![[1.0, f64::NAN], [2.0, 3.0]];
let err = km.fit(&data).unwrap_err();
assert!(
matches!(err, Error::NonFinite(_)),
"expected NonFinite, got {err:?}"
);
}
#[test]
fn fit_inf_data_is_error() {
let mut km = KMeans::new(1, 100, 1e-4).unwrap().with_random_state(0);
let data = array![[1.0, f64::INFINITY], [2.0, 3.0]];
let err = km.fit(&data).unwrap_err();
assert!(
matches!(err, Error::NonFinite(_)),
"expected NonFinite, got {err:?}"
);
}
#[test]
fn fit_fewer_samples_than_clusters_is_error() {
let mut km = KMeans::new(5, 100, 1e-4).unwrap().with_random_state(0);
let data = array![[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]];
let err = km.fit(&data).unwrap_err();
assert!(
matches!(err, Error::InvalidInput(_)),
"expected InvalidInput, got {err:?}"
);
}
#[test]
fn predict_before_fit_is_not_fitted() {
let km = KMeans::new(3, 100, 1e-4).unwrap().with_random_state(42);
let data = array![[1.0, 2.0]];
let err = km.predict(&data).unwrap_err();
assert!(
matches!(err, Error::NotFitted(_)),
"expected NotFitted, got {err:?}"
);
}
#[test]
fn predict_wrong_feature_count_is_dimension_mismatch() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
let data = three_blob_data(); km.fit(&data).unwrap();
let bad = array![[1.0, 2.0, 3.0]];
let err = km.predict(&bad).unwrap_err();
assert!(
matches!(err, Error::DimensionMismatch { .. }),
"expected DimensionMismatch, got {err:?}"
);
}
#[test]
fn predict_empty_input_after_fit_is_empty_input() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km.fit(&three_blob_data()).unwrap();
let empty: Array2<f64> = Array2::zeros((0, 2));
let err = km.predict(&empty).unwrap_err();
assert!(
matches!(err, Error::EmptyInput(_)),
"expected EmptyInput, got {err:?}"
);
}
#[test]
fn predict_nan_input_is_non_finite() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
let data = three_blob_data();
km.fit(&data).unwrap();
let bad = array![[f64::NAN, 1.0]];
let err = km.predict(&bad).unwrap_err();
assert!(
matches!(err, Error::NonFinite(_)),
"expected NonFinite, got {err:?}"
);
}
#[test]
fn fit_k3_centroids_near_true_blob_means() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
let data = three_blob_data();
km.fit(&data).unwrap();
let centroids = km.get_centroids().unwrap();
assert_eq!(centroids.nrows(), 3, "should have 3 centroids");
assert_eq!(centroids.ncols(), 2, "centroids should be 2-dimensional");
let true_means: [[f64; 2]; 3] = [[0.0, 0.0], [10.0, 0.0], [5.0, 10.0]];
let tol = 0.1_f64;
for &[tx, ty] in &true_means {
let matched = (0..3).any(|i| {
let cx = centroids[(i, 0)];
let cy = centroids[(i, 1)];
(cx - tx).abs() < tol && (cy - ty).abs() < tol
});
assert!(
matched,
"no centroid within tol={tol} of true mean ({tx}, {ty}); centroids = {centroids:?}"
);
}
}
#[test]
fn fit_k3_labels_respect_blob_structure() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
let data = three_blob_data();
km.fit(&data).unwrap();
let labels = km.get_labels().unwrap();
assert_eq!(labels.len(), 15);
assert_blob_structure(labels, 5);
}
#[test]
fn predict_on_training_data_matches_fit_labels() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
let data = three_blob_data();
km.fit(&data).unwrap();
let train_labels = km.get_labels().unwrap().clone();
let pred_labels = km.predict(&data).unwrap();
assert_eq!(train_labels.len(), pred_labels.len());
for (t, p) in train_labels.iter().zip(pred_labels.iter()) {
assert_eq!(t, p, "training label {t} != predicted label {p}");
}
}
#[test]
fn predict_new_points_near_blob_centres() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
let data = three_blob_data();
km.fit(&data).unwrap();
let train_labels = km.get_labels().unwrap();
let label_blob0 = train_labels[0];
let label_blob1 = train_labels[5];
let label_blob2 = train_labels[10];
let new_points = array![
[0.0, 0.0], [10.0, 0.0], [5.0, 10.0], ];
let pred = km.predict(&new_points).unwrap();
assert_eq!(
pred[0], label_blob0,
"point near (0,0) should map to blob 0"
);
assert_eq!(
pred[1], label_blob1,
"point near (10,0) should map to blob 1"
);
assert_eq!(
pred[2], label_blob2,
"point near (5,10) should map to blob 2"
);
}
#[test]
fn fit_predict_equals_fit_then_get_labels() {
let data = three_blob_data();
let mut km1 = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km1.fit(&data).unwrap();
let labels_via_fit = km1.get_labels().unwrap().clone();
let mut km2 = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
let labels_via_fp = km2.fit_predict(&data).unwrap();
assert_eq!(labels_via_fit.len(), labels_via_fp.len());
for (a, b) in labels_via_fit.iter().zip(labels_via_fp.iter()) {
assert_eq!(a, b);
}
}
#[test]
fn inertia_is_positive_on_blob_data() {
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km.fit(&three_blob_data()).unwrap();
let inertia = km.get_inertia().unwrap();
assert!(inertia > 0.0, "inertia must be positive, got {inertia}");
}
#[test]
fn inertia_decreases_as_k_grows() {
let data = three_blob_data();
let mut km1 = KMeans::new(1, 300, 1e-4).unwrap().with_random_state(42);
km1.fit(&data).unwrap();
let inertia_k1 = km1.get_inertia().unwrap();
let mut km3 = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km3.fit(&data).unwrap();
let inertia_k3 = km3.get_inertia().unwrap();
assert!(
inertia_k1 >= inertia_k3,
"k=1 inertia ({inertia_k1}) should be >= k=3 inertia ({inertia_k3})"
);
}
#[test]
fn k1_centroid_equals_global_mean() {
let data = three_blob_data();
let mut km = KMeans::new(1, 300, 1e-4).unwrap().with_random_state(42);
km.fit(&data).unwrap();
let centroids = km.get_centroids().unwrap();
assert_eq!(centroids.nrows(), 1);
let n = data.nrows() as f64;
let true_mean_x: f64 = data.column(0).iter().sum::<f64>() / n;
let true_mean_y: f64 = data.column(1).iter().sum::<f64>() / n;
assert_abs_diff_eq!(centroids[(0, 0)], true_mean_x, epsilon = 1e-9);
assert_abs_diff_eq!(centroids[(0, 1)], true_mean_y, epsilon = 1e-9);
}
#[test]
fn k1_all_labels_are_zero() {
let data = three_blob_data();
let mut km = KMeans::new(1, 300, 1e-4).unwrap().with_random_state(42);
km.fit(&data).unwrap();
let labels = km.get_labels().unwrap();
for &l in labels.iter() {
assert_eq!(l, 0, "k=1: all labels must be 0, got {l}");
}
}
#[test]
fn k_equals_n_samples_boundary() {
let data = three_blob_data();
let n = data.nrows();
let mut km_n = KMeans::new(n, 300, 1e-4).unwrap().with_random_state(42);
km_n.fit(&data).unwrap();
let centroids = km_n.get_centroids().unwrap();
assert_eq!(centroids.nrows(), n, "k=n_samples: should have n centroids");
assert_eq!(centroids.ncols(), 2);
let labels = km_n.get_labels().unwrap();
assert_eq!(labels.len(), n);
for &l in labels.iter() {
assert!(l < n, "label {l} out of range [0, {n})");
}
let inertia_n = km_n.get_inertia().unwrap();
assert!(
inertia_n >= 0.0,
"inertia must be non-negative, got {inertia_n}"
);
let mut km3 = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km3.fit(&data).unwrap();
let inertia_3 = km3.get_inertia().unwrap();
assert!(
inertia_n <= inertia_3 + 1e-9,
"k=n_samples inertia ({inertia_n}) should be <= k=3 inertia ({inertia_3})"
);
}
#[test]
fn same_seed_gives_identical_results() {
let data = three_blob_data();
let mut km_a = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km_a.fit(&data).unwrap();
let mut km_b = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km_b.fit(&data).unwrap();
let ca = km_a.get_centroids().unwrap();
let cb = km_b.get_centroids().unwrap();
assert_allclose(ca, cb, 0.0_f64);
let la = km_a.get_labels().unwrap();
let lb = km_b.get_labels().unwrap();
for (a, b) in la.iter().zip(lb.iter()) {
assert_eq!(a, b, "label mismatch between identical-seed runs");
}
let ia = km_a.get_inertia().unwrap();
let ib = km_b.get_inertia().unwrap();
assert_abs_diff_eq!(ia, ib, epsilon = 0.0);
let na = km_a.get_actual_iterations().unwrap();
let nb = km_b.get_actual_iterations().unwrap();
assert_eq!(na, nb, "n_iter mismatch between identical-seed runs");
}
#[test]
fn different_seeds_still_converge_to_same_partition_on_blob_data() {
let data = three_blob_data();
let mut km_a = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(0);
km_a.fit(&data).unwrap();
let mut km_b = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(99);
km_b.fit(&data).unwrap();
assert_blob_structure(km_a.get_labels().unwrap(), 5);
assert_blob_structure(km_b.get_labels().unwrap(), 5);
}
#[test]
fn actual_iterations_is_in_valid_range_after_fit() {
let max_iter = 300;
let mut km = KMeans::new(3, max_iter, 1e-4)
.unwrap()
.with_random_state(42);
km.fit(&three_blob_data()).unwrap();
let n_iter = km.get_actual_iterations().unwrap();
assert!(n_iter >= 1, "n_iter must be >= 1, got {n_iter:?}");
assert!(
n_iter <= max_iter,
"n_iter {n_iter:?} exceeds max_iter {max_iter}"
);
}
#[test]
fn save_load_round_trip_preserves_predictions() {
let data = three_blob_data();
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km.fit(&data).unwrap();
let path = "/tmp/rustyml_test_kmeans_round_trip.json";
km.save_to_path(path).expect("save_to_path should succeed");
let km_loaded = KMeans::load_from_path(path).expect("load_from_path should succeed");
let pred_original = km.predict(&data).unwrap();
let pred_loaded = km_loaded.predict(&data).unwrap();
assert_eq!(pred_original.len(), pred_loaded.len());
for (o, l) in pred_original.iter().zip(pred_loaded.iter()) {
assert_eq!(o, l, "prediction mismatch after save/load round-trip");
}
let co = km.get_centroids().unwrap();
let cl = km_loaded.get_centroids().unwrap();
assert_allclose(co, cl, 0.0_f64);
}
#[test]
fn load_from_nonexistent_path_is_io_error() {
let err =
KMeans::load_from_path("/tmp/this_path_does_not_exist_rustyml_kmeans.json").unwrap_err();
assert!(
matches!(err, Error::Io(_)),
"expected Error::Io, got {err:?}"
);
}
fn three_blobs_1200() -> Array2<f64> {
let centers = [(0.0_f64, 0.0_f64), (10.0, 0.0), (5.0, 10.0)];
let mut v = Vec::with_capacity(1200 * 2);
for (cx, cy) in centers {
for k in 0..400u32 {
v.push(cx + ((k * 7) % 11) as f64 * 0.01 - 0.05);
v.push(cy + ((k * 5) % 13) as f64 * 0.008 - 0.048);
}
}
Array2::from_shape_vec((1200, 2), v).unwrap()
}
#[test]
fn fit_parallel_branch_k3_centroids_near_true_means_1200() {
let data = three_blobs_1200();
assert_eq!(data.nrows(), 1200, "dataset must cross the 1000 threshold");
let mut km = KMeans::new(3, 300, 1e-4).unwrap().with_random_state(42);
km.fit(&data).unwrap();
let centroids = km.get_centroids().unwrap();
assert_eq!(centroids.nrows(), 3, "should have 3 centroids");
assert_eq!(centroids.ncols(), 2);
let true_means: [[f64; 2]; 3] = [[0.0, 0.0], [10.0, 0.0], [5.0, 10.0]];
let tol = 0.1_f64;
for &[tx, ty] in &true_means {
let matched = (0..3).any(|i| {
let cx = centroids[(i, 0)];
let cy = centroids[(i, 1)];
(cx - tx).abs() < tol && (cy - ty).abs() < tol
});
assert!(
matched,
"no centroid within tol={tol} of true mean ({tx}, {ty}); centroids = {centroids:?}"
);
}
let labels = km.get_labels().unwrap();
assert_eq!(labels.len(), 1200);
let block_label = [labels[0], labels[400], labels[800]];
assert_ne!(block_label[0], block_label[1]);
assert_ne!(block_label[0], block_label[2]);
assert_ne!(block_label[1], block_label[2]);
for (blk, &lab) in block_label.iter().enumerate() {
let start = blk * 400;
for i in start..start + 400 {
assert_eq!(
labels[i], lab,
"row {i} (blob {blk}) should share label {lab}, got {}",
labels[i]
);
}
}
}
#[test]
fn fit_parallel_accumulate_matches_serial_means_exactly() {
let n = 16_384_usize;
let d = 16_usize;
let k = 4_usize;
assert!(
n * d >= 262_144,
"dataset must cross the sum-gate work metric"
);
let mut v = Vec::with_capacity(n * d);
for i in 0..n {
for j in 0..d {
v.push((((i * 31 + j * 7) % 13) + (i % k) * 100) as f64);
}
}
let data = Array2::from_shape_vec((n, d), v).unwrap();
let mut km = KMeans::new(k, 1, 1e-12).unwrap().with_random_state(7);
km.fit(&data).unwrap();
let labels = km.get_labels().unwrap();
let centroids = km.get_centroids().unwrap();
let mut sums = Array2::<f64>::zeros((k, d));
let mut counts = vec![0usize; k];
for (i, &lab) in labels.iter().enumerate() {
counts[lab] += 1;
for j in 0..d {
sums[(lab, j)] += data[(i, j)];
}
}
for c in 0..k {
assert!(
counts[c] > 0,
"cluster {c} must be non-empty on this designed dataset (else the reseed \
path replaces its centroid and the mean identity no longer applies)"
);
for j in 0..d {
let expected = sums[(c, j)] / counts[c] as f64;
assert_eq!(
centroids[(c, j)],
expected,
"centroid ({c},{j}): parallel accumulate diverged from the exact serial mean"
);
}
}
}
#[test]
fn converged_centroids_equal_their_cluster_means() {
let mut km = KMeans::new(3, 300, 1e-6).unwrap().with_random_state(42);
let data = three_blob_data();
km.fit(&data).unwrap();
let centroids = km.get_centroids().unwrap();
let labels = km.get_labels().unwrap();
let n_features = data.ncols();
for k in 0..3 {
let mut sum = vec![0.0_f64; n_features];
let mut count = 0usize;
for (i, &lbl) in labels.iter().enumerate() {
if lbl == k {
for j in 0..n_features {
sum[j] += data[[i, j]];
}
count += 1;
}
}
assert!(count > 0, "cluster {k} is empty");
for j in 0..n_features {
let mean_j = sum[j] / count as f64;
assert_abs_diff_eq!(centroids[[k, j]], mean_j, epsilon = 1e-3);
}
}
}