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