use candle_core::{DType, Tensor};
use rand::{SeedableRng, rngs::StdRng, seq::SliceRandom};
use crate::{Result, device::default_device, distance::squared_l2};
const MAX_POINTS_PER_CENTROID: usize = 256;
const TRAINING_SHUFFLE_SEED: u64 = 0x7BE0_0CBE_7BE0_0CBE;
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(
p_tensor: &Tensor,
assignments: &Tensor,
previous: &Tensor,
) -> Result<Tensor> {
let (n, _) = p_tensor.dims2()?;
let (k, _) = previous.dims2()?;
let device = p_tensor.device();
let dtype = p_tensor.dtype();
if n == 0 {
return Ok(previous.clone());
}
let cluster_ids =
Tensor::arange(0u32, k as u32, device)?.reshape((1, k))?;
let one_hot = assignments
.reshape((n, 1))?
.broadcast_eq(&cluster_ids)?
.to_dtype(dtype)?;
let one_hot_t = one_hot.t()?.contiguous()?;
let sums = one_hot_t.matmul(p_tensor)?;
let counts = one_hot_t.sum_keepdim(1)?;
let safe_counts = counts.clamp(1.0f32, f32::MAX)?;
let means = sums.broadcast_div(&safe_counts)?;
let empty_mask = counts.eq(0.0f32)?.broadcast_as(means.shape())?;
Ok(empty_mask.where_cond(previous, &means)?)
}
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,
);
if points.is_empty() || max_iters == 0 {
return Ok(initial.to_vec());
}
let n_points = points.len() / dim;
let device = default_device();
let p_tensor = Tensor::from_slice(points, (n_points, dim), device)?;
fit_with_init_on_tensor(&p_tensor, initial, dim, max_iters)
}
fn fit_with_init_on_tensor(
p_tensor: &Tensor,
initial: &[f32],
dim: usize,
max_iters: usize,
) -> Result<Vec<f32>> {
let k = initial.len() / dim;
let device = p_tensor.device();
let mut centroids_t = Tensor::from_slice(initial, (k, dim), device)?;
let mut previous_assignments: Option<Tensor> = None;
for _ in 0..max_iters {
let assignments = assign_as_tensor(p_tensor, ¢roids_t)?;
if let Some(prev) = &previous_assignments {
let mismatches = assignments
.ne(prev)?
.to_dtype(DType::F32)?
.sum_all()?
.to_scalar::<f32>()?;
if mismatches == 0.0 {
break;
}
}
centroids_t = update_centroids(p_tensor, &assignments, ¢roids_t)?;
previous_assignments = Some(assignments);
}
Ok(centroids_t.flatten_all()?.to_vec1::<f32>()?)
}
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 (training_points, _) = sample_training_points(points, n_points, k, dim);
let device = default_device();
let n_train = training_points.len() / dim;
let training_tensor =
Tensor::from_slice(&training_points, (n_train, dim), device)?;
let initial = training_points[..k * dim].to_vec();
fit_with_init_on_tensor(&training_tensor, &initial, dim, max_iters)
}
pub fn fit_on_tensor(
p_tensor: &Tensor,
points: &[f32],
k: usize,
dim: usize,
max_iters: usize,
) -> Result<Vec<f32>> {
let n = p_tensor.dim(0)?;
let (training_points, same_as_pool) =
sample_training_points(points, n, k, dim);
let initial = training_points[..k * dim].to_vec();
if same_as_pool {
return fit_with_init_on_tensor(p_tensor, &initial, dim, max_iters);
}
let device = p_tensor.device();
let n_train = training_points.len() / dim;
let training_tensor =
Tensor::from_slice(&training_points, (n_train, dim), device)?;
fit_with_init_on_tensor(&training_tensor, &initial, dim, max_iters)
}
fn sample_training_points(
points: &[f32],
n: usize,
k: usize,
dim: usize,
) -> (Vec<f32>, bool) {
let target = (k * MAX_POINTS_PER_CENTROID).min(n);
if target == n {
return (points.to_vec(), true);
}
let mut rng = StdRng::seed_from_u64(TRAINING_SHUFFLE_SEED);
let mut indices: Vec<u32> = (0..n as u32).collect();
let (head, _) = indices.partial_shuffle(&mut rng, target);
let head: Vec<u32> = head.to_vec();
let mut sample = Vec::with_capacity(target * dim);
for &i in &head {
let start = i as usize * dim;
sample.extend_from_slice(&points[start..start + dim]);
}
(sample, false)
}
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());
}
let device = default_device();
let p = Tensor::from_slice(points, (n_points, dim), device)?;
assign_on_tensor(&p, centroids, dim)
}
pub fn assign_on_tensor(
p_tensor: &Tensor,
centroids: &[f32],
dim: usize,
) -> Result<Vec<usize>> {
assert!(dim > 0, "assign_on_tensor: dim must be positive");
assert!(
!centroids.is_empty() && centroids.len().is_multiple_of(dim),
"assign_on_tensor: centroids length {} is not a positive multiple of dim {}",
centroids.len(),
dim,
);
let k = centroids.len() / dim;
let device = p_tensor.device();
let c = Tensor::from_slice(centroids, (k, dim), device)?;
assign_tensor(p_tensor, &c)
}
pub fn assign_as_tensor(
p_tensor: &Tensor,
centroids_tensor: &Tensor,
) -> Result<Tensor> {
let n = p_tensor.dim(0)?;
if n == 0 {
return Ok(Tensor::zeros(
0,
candle_core::DType::U32,
p_tensor.device(),
)?);
}
let k = centroids_tensor.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);
let c_t = centroids_tensor.t()?;
let c_sq = centroids_tensor.sqr()?.sum_keepdim(1)?.t()?;
let mut chunks: Vec<Tensor> = Vec::new();
let mut start = 0usize;
while start < n {
let len = chunk_rows.min(n - start);
let p_chunk = p_tensor.narrow(0, start, len)?;
let dot = p_chunk.matmul(&c_t)?;
let scores = dot.affine(-2.0, 0.0)?.broadcast_add(&c_sq)?;
chunks.push(scores.argmin(1)?); start += len;
}
if chunks.len() == 1 {
Ok(chunks.into_iter().next().unwrap())
} else {
Ok(Tensor::cat(&chunks, 0)?)
}
}
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",
);
}
}
fn centroids_to_vec(t: &Tensor) -> Vec<f32> {
t.flatten_all().unwrap().to_vec1::<f32>().unwrap()
}
#[test]
fn update_centroids_averages_assigned_points() {
let device = default_device();
let points = Tensor::from_vec(
vec![0.0_f32, 0.0, 2.0, 0.0, 10.0, 0.0],
(3, 2),
device,
)
.unwrap();
let assignments =
Tensor::from_vec(vec![0u32, 0, 1], (3,), device).unwrap();
let previous =
Tensor::from_vec(vec![0.0_f32; 4], (2, 2), device).unwrap();
let updated =
update_centroids(&points, &assignments, &previous).unwrap();
assert_eq!(centroids_to_vec(&updated), vec![1.0, 0.0, 10.0, 0.0]);
}
#[test]
fn update_centroids_keeps_previous_for_empty_clusters() {
let device = default_device();
let points =
Tensor::from_vec(vec![0.0_f32, 0.0, 2.0, 0.0], (2, 2), device)
.unwrap();
let assignments =
Tensor::from_vec(vec![0u32, 0], (2,), device).unwrap();
let previous =
Tensor::from_vec(vec![5.0_f32, 5.0, 99.0, -99.0], (2, 2), device)
.unwrap();
let updated =
update_centroids(&points, &assignments, &previous).unwrap();
let updated = centroids_to_vec(&updated);
assert_eq!(updated[..2], [1.0, 0.0]);
assert_eq!(updated[2..], [99.0, -99.0]);
}
#[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_converges_to_well_separated_centroids_on_bimodal_corpus() {
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 fitted = fit(&points, 2, 2, 20).unwrap();
let c0 = &fitted[0..2];
let c1 = &fitted[2..4];
let dist = squared_l2(c0, c1);
assert!(
dist > 100.0,
"expected well-separated centroids after Lloyd, 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);
}
}