1use crate::SeededSampler;
4use std::{error::Error, fmt};
5
6#[derive(Clone, Debug, PartialEq)]
8pub enum ClusteringError {
9 EmptyInput,
11 ZeroDimension,
13 DimensionMismatch {
15 expected: usize,
17 actual: usize,
19 point: usize,
21 },
22 NonFiniteInput {
24 point: usize,
26 coordinate: usize,
28 value: f64,
30 },
31 InvalidComponentCount {
33 components: usize,
35 points: usize,
37 },
38 InvalidControl {
40 field: &'static str,
42 reason: &'static str,
44 },
45 WorkLimit {
47 limit: u64,
49 used: u64,
51 },
52 ArithmeticOverflow {
54 operation: &'static str,
56 },
57 SingularComponent {
59 component: usize,
61 },
62 NumericalFailure {
64 operation: &'static str,
66 },
67}
68
69impl fmt::Display for ClusteringError {
70 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
71 match self {
72 Self::EmptyInput => write!(f, "clustering requires at least one point"),
73 Self::ZeroDimension => write!(f, "clustering points must have at least one coordinate"),
74 Self::DimensionMismatch {
75 expected,
76 actual,
77 point,
78 } => write!(
79 f,
80 "clustering point {point} has dimension {actual}, expected {expected}"
81 ),
82 Self::NonFiniteInput {
83 point,
84 coordinate,
85 value,
86 } => write!(
87 f,
88 "clustering point {point} coordinate {coordinate} is not finite: {value}"
89 ),
90 Self::InvalidComponentCount { components, points } => write!(
91 f,
92 "clustering component count must be in 1..={points}, got {components}"
93 ),
94 Self::InvalidControl { field, reason } => {
95 write!(f, "invalid clustering control {field}: {reason}")
96 }
97 Self::WorkLimit { limit, used } => write!(
98 f,
99 "clustering work limit {limit} cannot admit another complete step after {used} units"
100 ),
101 Self::ArithmeticOverflow { operation } => {
102 write!(f, "clustering accounting overflowed during {operation}")
103 }
104 Self::SingularComponent { component } => {
105 write!(
106 f,
107 "mixture component {component} is singular under the selected policy"
108 )
109 }
110 Self::NumericalFailure { operation } => {
111 write!(
112 f,
113 "clustering produced a non-finite result during {operation}"
114 )
115 }
116 }
117 }
118}
119
120impl Error for ClusteringError {}
121
122#[derive(Clone, Copy, Debug, PartialEq)]
124pub struct KMeansControl {
125 pub seed: u64,
127 pub max_iterations: usize,
129 pub tolerance: f64,
131 pub max_work: u64,
133 pub restarts: usize,
135}
136
137impl KMeansControl {
138 pub fn new(
140 seed: u64,
141 max_iterations: usize,
142 tolerance: f64,
143 max_work: u64,
144 restarts: usize,
145 ) -> Result<Self, ClusteringError> {
146 let control = Self {
147 seed,
148 max_iterations,
149 tolerance,
150 max_work,
151 restarts,
152 };
153 control.validate()?;
154 Ok(control)
155 }
156
157 fn validate(self) -> Result<(), ClusteringError> {
158 for (field, valid, reason) in [
159 (
160 "max_iterations",
161 self.max_iterations > 0,
162 "must be greater than zero",
163 ),
164 ("max_work", self.max_work > 0, "must be greater than zero"),
165 ("restarts", self.restarts > 0, "must be greater than zero"),
166 (
167 "tolerance",
168 self.tolerance.is_finite() && self.tolerance >= 0.0,
169 "must be finite and nonnegative",
170 ),
171 ] {
172 if !valid {
173 return Err(ClusteringError::InvalidControl { field, reason });
174 }
175 }
176 Ok(())
177 }
178}
179
180impl Default for KMeansControl {
181 fn default() -> Self {
182 Self {
183 seed: 0,
184 max_iterations: 100,
185 tolerance: 1.0e-8,
186 max_work: 100_000,
187 restarts: 1,
188 }
189 }
190}
191
192#[derive(Clone, Copy, Debug, PartialEq, Eq)]
194pub enum KMeansTermination {
195 Converged,
197 IterationLimit,
199 WorkLimit,
201}
202
203#[derive(Clone, Copy, Debug, PartialEq, Eq)]
205pub enum KMeansSearchTermination {
206 Completed,
208 WorkLimit,
210}
211
212#[derive(Clone, Debug, PartialEq)]
214pub struct KMeansModel {
215 pub centroids: Vec<Vec<f64>>,
217 pub assignments: Vec<usize>,
219}
220
221#[derive(Clone, Debug, PartialEq)]
223pub struct KMeansRestartEvidence {
224 pub restart: usize,
226 pub seed: u64,
228 pub inertia: f64,
230 pub iterations: usize,
232 pub converged: bool,
234 pub empty_cluster_repairs: u64,
236 pub work: u64,
238 pub termination: KMeansTermination,
240}
241
242#[derive(Clone, Debug, PartialEq)]
244pub struct KMeansReport {
245 pub model: KMeansModel,
247 pub selected_restart: usize,
249 pub restarts: Vec<KMeansRestartEvidence>,
251 pub requested_restarts: usize,
253 pub work: u64,
255 pub termination: KMeansSearchTermination,
257}
258
259pub fn fit_kmeans(
265 points: &[Vec<f64>],
266 clusters: usize,
267 control: KMeansControl,
268) -> Result<KMeansReport, ClusteringError> {
269 validate_points(points)?;
270 validate_components(points.len(), clusters)?;
271 control.validate()?;
272
273 let mut meter = WorkMeter::new(control.max_work);
274 let mut candidates = Vec::with_capacity(control.restarts);
275 let mut models = Vec::with_capacity(control.restarts);
276 let mut search_termination = KMeansSearchTermination::Completed;
277
278 for restart in 0..control.restarts {
279 let seed = derived_seed(control.seed, restart);
280 match run_kmeans(points, clusters, control, restart, seed, &mut meter) {
281 Ok((model, evidence)) => {
282 let stopped = evidence.termination == KMeansTermination::WorkLimit;
283 models.push(model);
284 candidates.push(evidence);
285 if stopped {
286 search_termination = KMeansSearchTermination::WorkLimit;
287 break;
288 }
289 }
290 Err(ClusteringError::WorkLimit { .. }) if !models.is_empty() => {
291 search_termination = KMeansSearchTermination::WorkLimit;
292 break;
293 }
294 Err(error) => return Err(error),
295 }
296 }
297
298 let selected_restart = candidates
299 .iter()
300 .enumerate()
301 .min_by(|(left_index, left), (right_index, right)| {
302 left.inertia
303 .total_cmp(&right.inertia)
304 .then_with(|| left_index.cmp(right_index))
305 })
306 .map(|(index, _)| index)
307 .ok_or(ClusteringError::WorkLimit {
308 limit: control.max_work,
309 used: meter.used,
310 })?;
311
312 Ok(KMeansReport {
313 model: models.swap_remove(selected_restart),
314 selected_restart,
315 restarts: candidates,
316 requested_restarts: control.restarts,
317 work: meter.used,
318 termination: search_termination,
319 })
320}
321
322fn run_kmeans(
323 points: &[Vec<f64>],
324 clusters: usize,
325 control: KMeansControl,
326 restart: usize,
327 seed: u64,
328 meter: &mut WorkMeter,
329) -> Result<(KMeansModel, KMeansRestartEvidence), ClusteringError> {
330 let start_work = meter.used;
331 let mut random = SplitMix64::new(seed);
332 let centroids = kmeans_plus_plus(points, clusters, &mut random, meter)?;
333 let (assignments, residuals, inertia) = assign_points(points, ¢roids, meter)?;
334 let mut model = KMeansModel {
335 centroids,
336 assignments,
337 };
338 let mut inertia = inertia;
339 let mut residuals = residuals;
340 let mut iterations = 0;
341 let mut repairs = 0_u64;
342 let mut termination = KMeansTermination::IterationLimit;
343
344 while iterations < control.max_iterations {
345 let (next_centroids, next_repairs) = update_centroids(points, &model, &residuals, clusters);
346 let displacement = centroid_displacement(&model.centroids, &next_centroids)?;
347 let previous_assignments = model.assignments.clone();
348 let assigned = assign_points(points, &next_centroids, meter);
349 let (next_assignments, next_residuals, next_inertia) = match assigned {
350 Ok(result) => result,
351 Err(ClusteringError::WorkLimit { .. }) => {
352 termination = KMeansTermination::WorkLimit;
353 break;
354 }
355 Err(error) => return Err(error),
356 };
357 model.centroids = next_centroids;
358 model.assignments = next_assignments;
359 residuals = next_residuals;
360 inertia = next_inertia;
361 repairs = repairs.saturating_add(next_repairs);
362 iterations += 1;
363 if displacement <= control.tolerance || model.assignments == previous_assignments {
364 termination = KMeansTermination::Converged;
365 break;
366 }
367 }
368
369 canonicalize_kmeans(&mut model);
370 Ok((
371 model,
372 KMeansRestartEvidence {
373 restart,
374 seed,
375 inertia,
376 iterations,
377 converged: termination == KMeansTermination::Converged,
378 empty_cluster_repairs: repairs,
379 work: meter.used - start_work,
380 termination,
381 },
382 ))
383}
384
385fn update_centroids(
386 points: &[Vec<f64>],
387 model: &KMeansModel,
388 residuals: &[f64],
389 clusters: usize,
390) -> (Vec<Vec<f64>>, u64) {
391 let dimensions = points[0].len();
392 let mut centroids = vec![vec![0.0; dimensions]; clusters];
393 let mut counts = vec![0_usize; clusters];
394 for (point, &cluster) in points.iter().zip(&model.assignments) {
395 counts[cluster] += 1;
396 for (sum, &coordinate) in centroids[cluster].iter_mut().zip(point) {
397 *sum += coordinate;
398 }
399 }
400 for (centroid, &count) in centroids.iter_mut().zip(&counts) {
401 if count > 0 {
402 for coordinate in centroid {
403 *coordinate /= count as f64;
404 }
405 }
406 }
407
408 let mut used_points = vec![false; points.len()];
409 let mut repairs = 0_u64;
410 for cluster in 0..clusters {
411 if counts[cluster] != 0 {
412 continue;
413 }
414 let point = residuals
415 .iter()
416 .enumerate()
417 .filter(|(index, _)| !used_points[*index])
418 .max_by(|(left_index, left), (right_index, right)| {
419 left.total_cmp(right)
420 .then_with(|| right_index.cmp(left_index))
421 })
422 .map(|(index, _)| index)
423 .expect("clusters never exceed points");
424 centroids[cluster].clone_from(&points[point]);
425 used_points[point] = true;
426 repairs += 1;
427 }
428 (centroids, repairs)
429}
430
431fn centroid_displacement(current: &[Vec<f64>], next: &[Vec<f64>]) -> Result<f64, ClusteringError> {
432 let maximum = current
433 .iter()
434 .zip(next)
435 .map(|(left, right)| squared_distance(left, right).sqrt())
436 .max_by(f64::total_cmp)
437 .unwrap_or(0.0);
438 if maximum.is_finite() {
439 Ok(maximum)
440 } else {
441 Err(ClusteringError::NumericalFailure {
442 operation: "centroid displacement",
443 })
444 }
445}
446
447fn canonicalize_kmeans(model: &mut KMeansModel) {
448 let mut order = (0..model.centroids.len()).collect::<Vec<_>>();
449 order.sort_by(|&left, &right| compare_vectors(&model.centroids[left], &model.centroids[right]));
450 let mut remap = vec![0; order.len()];
451 for (new, &old) in order.iter().enumerate() {
452 remap[old] = new;
453 }
454 model.centroids = order
455 .iter()
456 .map(|&index| model.centroids[index].clone())
457 .collect();
458 for assignment in &mut model.assignments {
459 *assignment = remap[*assignment];
460 }
461}
462
463pub(crate) fn validate_points(points: &[Vec<f64>]) -> Result<usize, ClusteringError> {
464 let Some(first) = points.first() else {
465 return Err(ClusteringError::EmptyInput);
466 };
467 if first.is_empty() {
468 return Err(ClusteringError::ZeroDimension);
469 }
470 let dimensions = first.len();
471 for (point_index, point) in points.iter().enumerate() {
472 if point.len() != dimensions {
473 return Err(ClusteringError::DimensionMismatch {
474 expected: dimensions,
475 actual: point.len(),
476 point: point_index,
477 });
478 }
479 for (coordinate, &value) in point.iter().enumerate() {
480 if !value.is_finite() {
481 return Err(ClusteringError::NonFiniteInput {
482 point: point_index,
483 coordinate,
484 value,
485 });
486 }
487 }
488 }
489 Ok(dimensions)
490}
491
492pub(crate) fn validate_components(points: usize, components: usize) -> Result<(), ClusteringError> {
493 if components == 0 || components > points {
494 return Err(ClusteringError::InvalidComponentCount { components, points });
495 }
496 Ok(())
497}
498
499pub(crate) fn kmeans_plus_plus(
500 points: &[Vec<f64>],
501 clusters: usize,
502 random: &mut SplitMix64,
503 meter: &mut WorkMeter,
504) -> Result<Vec<Vec<f64>>, ClusteringError> {
505 let first = random.index_modulo(points.len());
506 let mut selected = vec![first];
507 let mut centroids = vec![points[first].clone()];
508 while centroids.len() < clusters {
509 let work = checked_product(points.len(), centroids.len(), "k-means++ distance work")?;
510 meter.charge(work)?;
511 let distances = points
512 .iter()
513 .map(|point| {
514 centroids
515 .iter()
516 .map(|centroid| squared_distance(point, centroid))
517 .min_by(f64::total_cmp)
518 .unwrap_or(0.0)
519 })
520 .collect::<Vec<_>>();
521 let total = distances.iter().sum::<f64>();
522 if !total.is_finite() {
523 return Err(ClusteringError::NumericalFailure {
524 operation: "k-means++ weighting",
525 });
526 }
527 let next = if total > 0.0 {
528 let threshold = random.unit_interval() * total;
529 let mut cumulative = 0.0;
530 distances
531 .iter()
532 .enumerate()
533 .find_map(|(index, distance)| {
534 cumulative += distance;
535 (cumulative > threshold).then_some(index)
536 })
537 .unwrap_or(points.len() - 1)
538 } else {
539 (0..points.len())
540 .find(|index| !selected.contains(index))
541 .unwrap_or(0)
542 };
543 selected.push(next);
544 centroids.push(points[next].clone());
545 }
546 Ok(centroids)
547}
548
549pub(crate) fn assign_points(
550 points: &[Vec<f64>],
551 centroids: &[Vec<f64>],
552 meter: &mut WorkMeter,
553) -> Result<(Vec<usize>, Vec<f64>, f64), ClusteringError> {
554 meter.charge(checked_product(
555 points.len(),
556 centroids.len(),
557 "k-means assignment work",
558 )?)?;
559 let mut assignments = Vec::with_capacity(points.len());
560 let mut residuals = Vec::with_capacity(points.len());
561 for point in points {
562 let (cluster, distance) = centroids
563 .iter()
564 .enumerate()
565 .map(|(index, centroid)| (index, squared_distance(point, centroid)))
566 .min_by(|(left_index, left), (right_index, right)| {
567 left.total_cmp(right)
568 .then_with(|| left_index.cmp(right_index))
569 })
570 .expect("component count was validated");
571 assignments.push(cluster);
572 residuals.push(distance);
573 }
574 let inertia = residuals.iter().sum::<f64>();
575 if inertia.is_finite() {
576 Ok((assignments, residuals, inertia))
577 } else {
578 Err(ClusteringError::NumericalFailure {
579 operation: "k-means inertia",
580 })
581 }
582}
583
584pub(crate) fn squared_distance(left: &[f64], right: &[f64]) -> f64 {
585 left.iter()
586 .zip(right)
587 .map(|(left, right)| {
588 let difference = left - right;
589 difference * difference
590 })
591 .sum()
592}
593
594pub(crate) fn compare_vectors(left: &[f64], right: &[f64]) -> std::cmp::Ordering {
595 left.iter()
596 .zip(right)
597 .find_map(|(left, right)| {
598 let ordering = left.total_cmp(right);
599 (ordering != std::cmp::Ordering::Equal).then_some(ordering)
600 })
601 .unwrap_or_else(|| left.len().cmp(&right.len()))
602}
603
604pub(crate) fn checked_product(
605 left: usize,
606 right: usize,
607 operation: &'static str,
608) -> Result<u64, ClusteringError> {
609 let left =
610 u64::try_from(left).map_err(|_| ClusteringError::ArithmeticOverflow { operation })?;
611 let right =
612 u64::try_from(right).map_err(|_| ClusteringError::ArithmeticOverflow { operation })?;
613 left.checked_mul(right)
614 .ok_or(ClusteringError::ArithmeticOverflow { operation })
615}
616
617pub(crate) struct WorkMeter {
618 pub(crate) limit: u64,
619 pub(crate) used: u64,
620}
621
622impl WorkMeter {
623 pub(crate) fn new(limit: u64) -> Self {
624 Self { limit, used: 0 }
625 }
626
627 pub(crate) fn charge(&mut self, amount: u64) -> Result<(), ClusteringError> {
628 let Some(next) = self.used.checked_add(amount) else {
629 return Err(ClusteringError::ArithmeticOverflow {
630 operation: "work charge",
631 });
632 };
633 if next > self.limit {
634 return Err(ClusteringError::WorkLimit {
635 limit: self.limit,
636 used: self.used,
637 });
638 }
639 self.used = next;
640 Ok(())
641 }
642}
643
644pub(crate) type SplitMix64 = SeededSampler;
645
646fn derived_seed(seed: u64, restart: usize) -> u64 {
647 let mut random = SplitMix64::new(seed ^ restart as u64);
648 random.next_u64()
649}