Skip to main content

sim_lib_numbers_stats/
clustering.rs

1//! Deterministic, bounded k-means clustering and shared clustering substrate.
2
3use crate::SeededSampler;
4use std::{error::Error, fmt};
5
6/// Errors returned by clustering and mixture-model fitting.
7#[derive(Clone, Debug, PartialEq)]
8pub enum ClusteringError {
9    /// No observations were supplied.
10    EmptyInput,
11    /// An observation had no coordinates.
12    ZeroDimension,
13    /// Observation dimensions were inconsistent.
14    DimensionMismatch {
15        /// Expected coordinate count.
16        expected: usize,
17        /// Actual coordinate count.
18        actual: usize,
19        /// Index of the mismatching observation.
20        point: usize,
21    },
22    /// An observation contained NaN or infinity.
23    NonFiniteInput {
24        /// Observation index.
25        point: usize,
26        /// Coordinate index.
27        coordinate: usize,
28        /// Rejected value.
29        value: f64,
30    },
31    /// A requested component count cannot be fitted to the observations.
32    InvalidComponentCount {
33        /// Requested component count.
34        components: usize,
35        /// Number of supplied observations.
36        points: usize,
37    },
38    /// A control or model field was invalid.
39    InvalidControl {
40        /// Name of the invalid field.
41        field: &'static str,
42        /// Stable reason for rejection.
43        reason: &'static str,
44    },
45    /// The work bound could not admit one complete initial model.
46    WorkLimit {
47        /// Configured work bound.
48        limit: u64,
49        /// Work already charged when admission failed.
50        used: u64,
51    },
52    /// Size or work accounting overflowed.
53    ArithmeticOverflow {
54        /// Operation whose accounting overflowed.
55        operation: &'static str,
56    },
57    /// A component could not be made numerically nonsingular under policy.
58    SingularComponent {
59        /// Component index in stable model order.
60        component: usize,
61    },
62    /// Finite inputs produced a non-finite intermediate result.
63    NumericalFailure {
64        /// Operation that failed.
65        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/// Deterministic initialization, convergence, restart, and work policy for k-means.
123#[derive(Clone, Copy, Debug, PartialEq)]
124pub struct KMeansControl {
125    /// Root seed from which bounded restart seeds are derived.
126    pub seed: u64,
127    /// Maximum number of complete Lloyd updates per restart.
128    pub max_iterations: usize,
129    /// Maximum centroid displacement accepted as convergence.
130    pub tolerance: f64,
131    /// Hard bound on point-to-centroid distance evaluations across all restarts.
132    pub max_work: u64,
133    /// Maximum number of independently seeded candidate results.
134    pub restarts: usize,
135}
136
137impl KMeansControl {
138    /// Builds checked k-means control.
139    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/// Why one k-means restart stopped.
193#[derive(Clone, Copy, Debug, PartialEq, Eq)]
194pub enum KMeansTermination {
195    /// Assignments or centroid displacement met the convergence policy.
196    Converged,
197    /// The restart exhausted its iteration count.
198    IterationLimit,
199    /// The shared work budget could not admit another complete Lloyd step.
200    WorkLimit,
201}
202
203/// Why the bounded multi-restart search stopped.
204#[derive(Clone, Copy, Debug, PartialEq, Eq)]
205pub enum KMeansSearchTermination {
206    /// Every requested restart completed.
207    Completed,
208    /// The shared work budget stopped the current or next restart.
209    WorkLimit,
210}
211
212/// Inspectable k-means centroids and stable cluster assignments.
213#[derive(Clone, Debug, PartialEq)]
214pub struct KMeansModel {
215    /// Lexicographically ordered centroids.
216    pub centroids: Vec<Vec<f64>>,
217    /// Cluster index for every input point, in input order.
218    pub assignments: Vec<usize>,
219}
220
221/// Convergence evidence for one bounded restart.
222#[derive(Clone, Debug, PartialEq)]
223pub struct KMeansRestartEvidence {
224    /// Zero-based restart index.
225    pub restart: usize,
226    /// Derived seed used by k-means++ initialization.
227    pub seed: u64,
228    /// Final within-cluster sum of squared distances.
229    pub inertia: f64,
230    /// Number of complete Lloyd updates.
231    pub iterations: usize,
232    /// Whether convergence, rather than a bound, stopped the restart.
233    pub converged: bool,
234    /// Number of empty centroids deterministically reseeded from worst-fit points.
235    pub empty_cluster_repairs: u64,
236    /// Distance-evaluation work charged by this restart.
237    pub work: u64,
238    /// Concrete termination reason.
239    pub termination: KMeansTermination,
240}
241
242/// Selected model and complete multi-restart evidence.
243#[derive(Clone, Debug, PartialEq)]
244pub struct KMeansReport {
245    /// Lowest-inertia model, with stable restart-index tie breaking.
246    pub model: KMeansModel,
247    /// Index into [`Self::restarts`] selected by the model-selection policy.
248    pub selected_restart: usize,
249    /// Evidence for every candidate that reached a complete assignment.
250    pub restarts: Vec<KMeansRestartEvidence>,
251    /// Number of requested candidates.
252    pub requested_restarts: usize,
253    /// Total charged work across initialization and all candidates.
254    pub work: u64,
255    /// Concrete search-level termination reason.
256    pub termination: KMeansSearchTermination,
257}
258
259/// Fits deterministic seeded k-means with k-means++ initialization.
260///
261/// Empty clusters are reseeded from distinct worst-fit points. Candidate models
262/// are compared by inertia, then by restart order, so identical inputs and
263/// control produce byte-for-byte identical reports.
264pub 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, &centroids, 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}