Skip to main content

sim_lib_numbers_stats/
clustering.rs

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