use candle_core::Tensor;
use crate::{Result, device::default_device, distance::squared_l2};
pub fn nearest_centroid(point: &[f32], centroids: &[f32], dim: usize) -> usize {
assert_eq!(
point.len(),
dim,
"nearest_centroid: point length {} does not match dim {}",
point.len(),
dim,
);
assert!(dim > 0, "nearest_centroid: dim must be positive");
assert!(
!centroids.is_empty() && centroids.len().is_multiple_of(dim),
"nearest_centroid: centroids length {} is not a positive multiple of dim {}",
centroids.len(),
dim,
);
let mut best_idx = 0usize;
let mut best_dist = f32::INFINITY;
for (i, centroid) in centroids.chunks_exact(dim).enumerate() {
let d = squared_l2(point, centroid);
if d < best_dist {
best_dist = d;
best_idx = i;
}
}
best_idx
}
pub fn update_centroids(
points: &[f32],
assignments: &[usize],
previous: &[f32],
dim: usize,
) -> Vec<f32> {
assert!(dim > 0, "update_centroids: dim must be positive");
assert!(
points.len().is_multiple_of(dim),
"update_centroids: points length {} is not a multiple of dim {}",
points.len(),
dim,
);
assert!(
previous.len().is_multiple_of(dim) && !previous.is_empty(),
"update_centroids: previous length {} is not a positive multiple of dim {}",
previous.len(),
dim,
);
let k = previous.len() / dim;
assert_eq!(
assignments.len(),
points.len() / dim,
"update_centroids: {} assignments for {} points",
assignments.len(),
points.len() / dim,
);
let mut sums = vec![0.0f32; k * dim];
let mut counts = vec![0usize; k];
for (point, &cluster) in points.chunks_exact(dim).zip(assignments.iter()) {
assert!(
cluster < k,
"update_centroids: assignment {cluster} out of range 0..{k}"
);
let slot = &mut sums[cluster * dim..(cluster + 1) * dim];
for (s, p) in slot.iter_mut().zip(point.iter()) {
*s += *p;
}
counts[cluster] += 1;
}
let mut new_centroids = vec![0.0f32; k * dim];
for (cluster, &count) in counts.iter().enumerate() {
let start = cluster * dim;
let end = start + dim;
if count == 0 {
new_centroids[start..end].copy_from_slice(&previous[start..end]);
continue;
}
let inv = 1.0f32 / count as f32;
for (out, s) in
new_centroids[start..end].iter_mut().zip(&sums[start..end])
{
*out = s * inv;
}
}
new_centroids
}
pub fn fit_with_init(
points: &[f32],
initial: &[f32],
dim: usize,
max_iters: usize,
) -> Result<Vec<f32>> {
assert!(dim > 0, "fit_with_init: dim must be positive");
assert!(
initial.len().is_multiple_of(dim) && !initial.is_empty(),
"fit_with_init: initial centroids length {} is not a positive multiple of dim {}",
initial.len(),
dim,
);
assert!(
points.len().is_multiple_of(dim),
"fit_with_init: points length {} is not a multiple of dim {}",
points.len(),
dim,
);
let mut centroids = initial.to_vec();
if points.is_empty() || max_iters == 0 {
return Ok(centroids);
}
let n_points = points.len() / dim;
let k = initial.len() / dim;
let device = default_device();
let p_tensor = Tensor::from_slice(points, (n_points, dim), device)?;
let mut previous_assignments: Option<Vec<usize>> = None;
for _ in 0..max_iters {
let c_tensor = Tensor::from_slice(¢roids, (k, dim), device)?;
let assignments = assign_tensor(&p_tensor, &c_tensor)?;
if previous_assignments.as_deref() == Some(assignments.as_slice()) {
break;
}
centroids = update_centroids(points, &assignments, ¢roids, dim);
previous_assignments = Some(assignments);
}
Ok(centroids)
}
pub fn fit(
points: &[f32],
k: usize,
dim: usize,
max_iters: usize,
) -> Result<Vec<f32>> {
assert!(k > 0, "fit: k must be positive");
assert!(dim > 0, "fit: dim must be positive");
assert!(
points.len().is_multiple_of(dim),
"fit: points length {} is not a multiple of dim {}",
points.len(),
dim,
);
let n_points = points.len() / dim;
assert!(
n_points >= k,
"fit: need at least k={k} points, got {n_points}",
);
let initial = farthest_first_init(points, k, dim);
fit_with_init(points, &initial, dim, max_iters)
}
pub fn farthest_first_init(points: &[f32], k: usize, dim: usize) -> Vec<f32> {
assert!(k > 0, "farthest_first_init: k must be positive");
assert!(dim > 0, "farthest_first_init: dim must be positive");
assert!(
points.len().is_multiple_of(dim) && !points.is_empty(),
"farthest_first_init: points length {} is not a positive multiple of dim {}",
points.len(),
dim,
);
let n = points.len() / dim;
assert!(
n >= k,
"farthest_first_init: need at least k={k} points, got {n}",
);
let mut centroids = Vec::with_capacity(k * dim);
centroids.extend_from_slice(&points[..dim]);
let mut min_dists = vec![0.0f32; n];
for (i, p) in points.chunks_exact(dim).enumerate() {
min_dists[i] = squared_l2(p, &points[..dim]);
}
for _ in 1..k {
let mut best_idx = 0usize;
let mut best_dist = f32::NEG_INFINITY;
for (i, &d) in min_dists.iter().enumerate() {
if d > best_dist {
best_dist = d;
best_idx = i;
}
}
let new_c = &points[best_idx * dim..(best_idx + 1) * dim];
centroids.extend_from_slice(new_c);
for (i, p) in points.chunks_exact(dim).enumerate() {
let d = squared_l2(p, new_c);
if d < min_dists[i] {
min_dists[i] = d;
}
}
}
centroids
}
pub fn assign_points(
points: &[f32],
centroids: &[f32],
dim: usize,
) -> Result<Vec<usize>> {
assert!(dim > 0, "assign_points: dim must be positive");
assert!(
points.len().is_multiple_of(dim),
"assign_points: points length {} is not a multiple of dim {}",
points.len(),
dim,
);
let n_points = points.len() / dim;
if n_points == 0 {
return Ok(Vec::new());
}
assert!(
!centroids.is_empty() && centroids.len().is_multiple_of(dim),
"assign_points: centroids length {} is not a positive multiple of dim {}",
centroids.len(),
dim,
);
let k = centroids.len() / dim;
let device = default_device();
let p = Tensor::from_slice(points, (n_points, dim), device)?;
let c = Tensor::from_slice(centroids, (k, dim), device)?;
assign_tensor(&p, &c)
}
const ASSIGN_CHUNK_BYTES: usize = 128 * 1024 * 1024;
fn assign_tensor(p: &Tensor, c: &Tensor) -> Result<Vec<usize>> {
let n = p.dim(0)?;
if n == 0 {
return Ok(Vec::new());
}
let k = c.dim(0)?.max(1);
let bytes_per_row = k * std::mem::size_of::<f32>();
let chunk_rows = (ASSIGN_CHUNK_BYTES / bytes_per_row).max(1).min(n);
assign_tensor_chunked(p, c, chunk_rows)
}
fn assign_tensor_chunked(
p: &Tensor,
c: &Tensor,
chunk_rows: usize,
) -> Result<Vec<usize>> {
assert!(
chunk_rows > 0,
"assign_tensor_chunked: chunk_rows must be positive",
);
let n = p.dim(0)?;
if n == 0 {
return Ok(Vec::new());
}
let c_t = c.t()?;
let c_sq = c.sqr()?.sum_keepdim(1)?.t()?;
let mut out = Vec::with_capacity(n);
let mut start = 0usize;
while start < n {
let len = chunk_rows.min(n - start);
let p_chunk = p.narrow(0, start, len)?;
let dot = p_chunk.matmul(&c_t)?; let scores = dot.affine(-2.0, 0.0)?.broadcast_add(&c_sq)?;
let argmin_u32: Vec<u32> = scores.argmin(1)?.to_vec1::<u32>()?;
out.extend(argmin_u32.into_iter().map(|x| x as usize));
start += len;
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn nearest_centroid_returns_zero_for_single_centroid() {
let centroids = [1.0, 2.0, 3.0];
assert_eq!(nearest_centroid(&[0.0, 0.0, 0.0], ¢roids, 3), 0);
}
#[test]
fn nearest_centroid_picks_the_closest_of_many() {
let centroids = [0.0, 0.0, 5.0, 5.0, 10.0, 10.0];
assert_eq!(nearest_centroid(&[4.0, 4.0], ¢roids, 2), 1);
}
#[test]
fn nearest_centroid_breaks_ties_toward_earlier_index() {
let centroids = [0.0, 0.0, 10.0, 0.0];
assert_eq!(nearest_centroid(&[5.0, 0.0], ¢roids, 2), 0);
}
#[test]
#[should_panic(expected = "point length")]
fn nearest_centroid_panics_on_dim_mismatch() {
let centroids = [0.0, 0.0];
let _ = nearest_centroid(&[1.0], ¢roids, 2);
}
#[test]
#[should_panic(expected = "centroids length")]
fn nearest_centroid_panics_on_ragged_centroid_block() {
let centroids = [0.0, 0.0, 1.0];
let _ = nearest_centroid(&[0.0, 0.0], ¢roids, 2);
}
#[test]
fn assign_points_maps_each_point_to_its_cluster() {
let centroids = [0.0, 0.0, 10.0, 10.0];
let points = [0.1, 0.2, 9.0, 9.5, 0.0, -0.1, 10.1, 10.1];
assert_eq!(
assign_points(&points, ¢roids, 2).unwrap(),
vec![0, 1, 0, 1],
);
}
#[test]
fn assign_points_handles_empty_input() {
let centroids = [0.0, 1.0];
assert_eq!(
assign_points(&[], ¢roids, 2).unwrap(),
Vec::<usize>::new(),
);
}
#[test]
#[should_panic(expected = "points length")]
fn assign_points_panics_on_ragged_points() {
let centroids = [0.0, 0.0];
let _ = assign_points(&[1.0, 2.0, 3.0], ¢roids, 2);
}
#[test]
fn assign_tensor_chunked_agrees_with_unchunked() {
let dim = 3;
let k = 4;
let centroids: Vec<f32> = vec![
0.0, 0.0, 0.0, 10.0, 0.0, 0.0, 0.0, 10.0, 0.0, 10.0, 10.0, 0.0, ];
let n = 173;
let mut points: Vec<f32> = Vec::with_capacity(n * dim);
for i in 0..n {
let a = ((i * 37) % 11) as f32;
let b = ((i * 53) % 13) as f32;
points.extend_from_slice(&[a, b, 0.0]);
}
let device = default_device();
let p = Tensor::from_slice(&points, (n, dim), device).unwrap();
let c = Tensor::from_slice(¢roids, (k, dim), device).unwrap();
let baseline = assign_tensor_chunked(&p, &c, n).unwrap();
assert_eq!(baseline.len(), n);
for chunk in [1usize, 7, 64, n - 1, n, n + 5] {
let chunked = assign_tensor_chunked(&p, &c, chunk).unwrap();
assert_eq!(
chunked, baseline,
"chunk_rows={chunk} must agree with the unchunked result",
);
}
}
#[test]
fn update_centroids_averages_assigned_points() {
let points = [0.0, 0.0, 2.0, 0.0, 10.0, 0.0];
let assignments = [0, 0, 1];
let previous = [0.0, 0.0, 0.0, 0.0];
let updated = update_centroids(&points, &assignments, &previous, 2);
assert_eq!(updated, vec![1.0, 0.0, 10.0, 0.0]);
}
#[test]
fn update_centroids_keeps_previous_for_empty_clusters() {
let points = [0.0, 0.0, 2.0, 0.0];
let assignments = [0, 0];
let previous = [5.0, 5.0, 99.0, -99.0];
let updated = update_centroids(&points, &assignments, &previous, 2);
assert_eq!(updated[..2], [1.0, 0.0]);
assert_eq!(updated[2..], [99.0, -99.0]);
}
#[test]
#[should_panic(expected = "out of range")]
fn update_centroids_panics_on_out_of_range_assignment() {
let points = [0.0, 0.0];
let assignments = [5];
let previous = [0.0, 0.0];
let _ = update_centroids(&points, &assignments, &previous, 2);
}
#[test]
fn fit_with_init_converges_on_well_separated_clusters() {
let points = vec![
0.0, 0.0, 0.2, -0.1, -0.1, 0.1, 10.0, 10.0, 10.1, 9.9, 9.9, 10.1, ];
let initial = [0.5, 0.5, 9.5, 9.5];
let fitted = fit_with_init(&points, &initial, 2, 20).unwrap();
assert_eq!(fitted.len(), 4);
assert!((fitted[0] - 0.0333).abs() < 1e-3);
assert!((fitted[1] - 0.0).abs() < 1e-3);
assert!((fitted[2] - 10.0).abs() < 1e-3);
assert!((fitted[3] - 10.0).abs() < 1e-3);
}
#[test]
fn fit_with_init_is_idempotent_after_convergence() {
let points = vec![0.0, 0.0, 1.0, 0.0, 10.0, 0.0, 11.0, 0.0];
let initial = [0.5, 0.0, 10.5, 0.0];
let once = fit_with_init(&points, &initial, 2, 50).unwrap();
let twice = fit_with_init(&points, &once, 2, 50).unwrap();
assert_eq!(once, twice);
}
#[test]
fn fit_with_init_returns_initial_when_no_iterations_allowed() {
let points = vec![0.0, 0.0, 10.0, 10.0];
let initial = [1.0, 1.0, 9.0, 9.0];
assert_eq!(
fit_with_init(&points, &initial, 2, 0).unwrap(),
initial.to_vec(),
);
}
#[test]
fn fit_uses_farthest_first_seeding_to_spread_initial_centroids() {
let points = vec![
0.0, 0.0, 0.1, 0.1, -0.1, 0.1, 10.0, 10.0, 10.1, 9.9, 9.9, 10.1, ];
let seeds = fit(&points, 2, 2, 0).unwrap();
let c0 = &seeds[0..2];
let c1 = &seeds[2..4];
let dist = squared_l2(c0, c1);
assert!(
dist > 100.0,
"expected well-separated seeds, got {c0:?} and {c1:?} (dist²={dist})",
);
}
#[test]
fn fit_is_deterministic_across_identical_calls() {
let points = vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8];
let a = fit(&points, 3, 2, 10).unwrap();
let b = fit(&points, 3, 2, 10).unwrap();
assert_eq!(a, b);
}
#[test]
#[should_panic(expected = "need at least")]
fn fit_panics_when_asked_for_more_clusters_than_points() {
let points = vec![0.0, 0.0];
let _ = fit(&points, 5, 2, 1);
}
}