Skip to main content

sim_lib_numbers_stats/
gmm.rs

1//! Regularized Gaussian-mixture EM with log-domain evidence.
2
3use super::clustering::{
4    ClusteringError, SplitMix64, WorkMeter, checked_product, kmeans_plus_plus, validate_components,
5    validate_points,
6};
7use super::gmm_math::{
8    PreparedCovariance, canonicalize_model, log_sum_exp, model_selection, prepare_covariance,
9    prepare_model, require_finite_covariance, validate_model,
10};
11
12/// Covariance representation fitted for every Gaussian component.
13#[derive(Clone, Copy, Debug, PartialEq, Eq)]
14pub enum CovarianceType {
15    /// One regularized variance per coordinate.
16    Diagonal,
17    /// One regularized symmetric covariance matrix per component.
18    Full,
19}
20
21/// Inspectable covariance parameters for one Gaussian component.
22#[derive(Clone, Debug, PartialEq)]
23pub enum GaussianCovariance {
24    /// Coordinate variances in point-coordinate order.
25    Diagonal(Vec<f64>),
26    /// Symmetric row-major covariance matrix.
27    Full(Vec<Vec<f64>>),
28}
29
30/// Policy for an EM component with negligible responsibility or singular covariance.
31#[derive(Clone, Copy, Debug, PartialEq)]
32pub enum SingularComponentPolicy {
33    /// Reinitialize from a stable worst-fit observation and the global covariance.
34    Reinitialize {
35        /// Smallest accepted fraction of total responsibility mass.
36        minimum_weight: f64,
37    },
38    /// Fail closed instead of changing a singular component.
39    Fail {
40        /// Smallest accepted fraction of total responsibility mass.
41        minimum_weight: f64,
42    },
43}
44
45impl SingularComponentPolicy {
46    fn minimum_weight(self) -> f64 {
47        match self {
48            Self::Reinitialize { minimum_weight } | Self::Fail { minimum_weight } => minimum_weight,
49        }
50    }
51
52    fn validate(self) -> Result<(), ClusteringError> {
53        let weight = self.minimum_weight();
54        if !weight.is_finite() || !(0.0..1.0).contains(&weight) {
55            return Err(ClusteringError::InvalidControl {
56                field: "singular_policy.minimum_weight",
57                reason: "must be finite and in the open interval (0, 1)",
58            });
59        }
60        Ok(())
61    }
62}
63
64impl Default for SingularComponentPolicy {
65    fn default() -> Self {
66        Self::Reinitialize {
67            minimum_weight: 1.0e-8,
68        }
69    }
70}
71
72/// Component count, covariance family, regularization, and singular policy.
73#[derive(Clone, Copy, Debug, PartialEq)]
74pub struct GmmSpec {
75    /// Number of Gaussian components.
76    pub components: usize,
77    /// Covariance representation shared by all components.
78    pub covariance: CovarianceType,
79    /// Positive value added to every fitted covariance diagonal.
80    pub regularization: f64,
81    /// Explicit behavior for empty or numerically singular components.
82    pub singular_policy: SingularComponentPolicy,
83}
84
85impl GmmSpec {
86    /// Builds a checked regularized mixture specification.
87    pub fn new(
88        components: usize,
89        covariance: CovarianceType,
90        regularization: f64,
91        singular_policy: SingularComponentPolicy,
92    ) -> Result<Self, ClusteringError> {
93        let spec = Self {
94            components,
95            covariance,
96            regularization,
97            singular_policy,
98        };
99        spec.validate()?;
100        Ok(spec)
101    }
102
103    fn validate(self) -> Result<(), ClusteringError> {
104        if self.components == 0 {
105            return Err(ClusteringError::InvalidControl {
106                field: "components",
107                reason: "must be greater than zero",
108            });
109        }
110        if !self.regularization.is_finite() || self.regularization <= 0.0 {
111            return Err(ClusteringError::InvalidControl {
112                field: "regularization",
113                reason: "must be finite and greater than zero",
114            });
115        }
116        self.singular_policy.validate()
117    }
118}
119
120/// Deterministic convergence and work policy for Gaussian-mixture EM.
121#[derive(Clone, Copy, Debug, PartialEq)]
122pub struct GmmControl {
123    /// Seed used by k-means++ mean initialization.
124    pub seed: u64,
125    /// Maximum number of accepted EM updates.
126    pub max_iterations: usize,
127    /// Relative log-likelihood convergence tolerance.
128    pub tolerance: f64,
129    /// Hard bound on initialization, responsibility, and parameter-update work.
130    pub max_work: u64,
131}
132
133impl GmmControl {
134    /// Builds checked EM control.
135    pub fn new(
136        seed: u64,
137        max_iterations: usize,
138        tolerance: f64,
139        max_work: u64,
140    ) -> Result<Self, ClusteringError> {
141        let control = Self {
142            seed,
143            max_iterations,
144            tolerance,
145            max_work,
146        };
147        control.validate()?;
148        Ok(control)
149    }
150
151    fn validate(self) -> Result<(), ClusteringError> {
152        for (field, valid, reason) in [
153            (
154                "max_iterations",
155                self.max_iterations > 0,
156                "must be greater than zero",
157            ),
158            ("max_work", self.max_work > 0, "must be greater than zero"),
159            (
160                "tolerance",
161                self.tolerance.is_finite() && self.tolerance >= 0.0,
162                "must be finite and nonnegative",
163            ),
164        ] {
165            if !valid {
166                return Err(ClusteringError::InvalidControl { field, reason });
167            }
168        }
169        Ok(())
170    }
171}
172
173impl Default for GmmControl {
174    fn default() -> Self {
175        Self {
176            seed: 0,
177            max_iterations: 100,
178            tolerance: 1.0e-8,
179            max_work: 1_000_000,
180        }
181    }
182}
183
184/// Why bounded EM stopped.
185#[derive(Clone, Copy, Debug, PartialEq, Eq)]
186pub enum GmmTermination {
187    /// Relative log-likelihood improvement met the tolerance.
188    Converged,
189    /// The configured iteration count was exhausted.
190    IterationLimit,
191    /// The work budget could not admit another complete update or score.
192    WorkLimit,
193    /// A candidate reduced likelihood and was rejected.
194    LikelihoodDecrease,
195}
196
197/// AIC and BIC evidence for comparing fitted component counts.
198#[derive(Clone, Copy, Debug, PartialEq)]
199pub struct ModelSelectionEvidence {
200    /// Maximized natural-log likelihood.
201    pub log_likelihood: f64,
202    /// Number of independently fitted free parameters.
203    pub parameters: usize,
204    /// Akaike information criterion; lower is preferred.
205    pub aic: f64,
206    /// Bayesian information criterion; lower is preferred.
207    pub bic: f64,
208    /// Number of observations used by BIC.
209    pub observations: usize,
210}
211
212/// Inspectable fitted Gaussian-mixture parameters.
213#[derive(Clone, Debug, PartialEq)]
214pub struct GmmModel {
215    /// Normalized component weights.
216    pub weights: Vec<f64>,
217    /// Component means in lexicographic order.
218    pub means: Vec<Vec<f64>>,
219    /// Covariance parameters aligned with [`Self::means`].
220    pub covariances: Vec<GaussianCovariance>,
221}
222
223impl GmmModel {
224    /// Returns posterior component probabilities for every supplied point.
225    pub fn responsibilities(&self, points: &[Vec<f64>]) -> Result<Vec<Vec<f64>>, ClusteringError> {
226        validate_points(points)?;
227        validate_model(self, points[0].len())?;
228        let mut meter = WorkMeter::new(u64::MAX);
229        Ok(expectation(points, self, &mut meter)?.responsibilities)
230    }
231
232    /// Returns the summed natural-log likelihood of supplied points.
233    pub fn log_likelihood(&self, points: &[Vec<f64>]) -> Result<f64, ClusteringError> {
234        validate_points(points)?;
235        validate_model(self, points[0].len())?;
236        let mut meter = WorkMeter::new(u64::MAX);
237        Ok(expectation(points, self, &mut meter)?.log_likelihood)
238    }
239
240    /// Returns the maximum-posterior component index for every point.
241    pub fn predict(&self, points: &[Vec<f64>]) -> Result<Vec<usize>, ClusteringError> {
242        self.responsibilities(points).map(|rows| {
243            rows.iter()
244                .map(|row| {
245                    row.iter()
246                        .enumerate()
247                        .max_by(|(left_index, left), (right_index, right)| {
248                            left.total_cmp(right)
249                                .then_with(|| right_index.cmp(left_index))
250                        })
251                        .map(|(index, _)| index)
252                        .expect("fitted model has components")
253                })
254                .collect()
255        })
256    }
257}
258
259/// Convergence, regularization, work, and selection evidence from EM.
260#[derive(Clone, Debug, PartialEq)]
261pub struct GmmEvidence {
262    /// Initial mixture log likelihood.
263    pub initial_log_likelihood: f64,
264    /// Final accepted mixture log likelihood.
265    pub log_likelihood: f64,
266    /// Initial value followed by every accepted likelihood.
267    pub likelihood_history: Vec<f64>,
268    /// Number of accepted EM updates.
269    pub iterations: usize,
270    /// Whether tolerance caused termination.
271    pub converged: bool,
272    /// Count of components repaired under the singular policy.
273    pub singular_component_repairs: u64,
274    /// Caller-supplied initialization seed.
275    pub seed: u64,
276    /// Charged work, never greater than the configured limit.
277    pub work: u64,
278    /// Concrete termination reason.
279    pub termination: GmmTermination,
280    /// AIC/BIC evidence for component-count selection.
281    pub model_selection: ModelSelectionEvidence,
282}
283
284/// Fitted mixture and complete EM evidence.
285#[derive(Clone, Debug, PartialEq)]
286pub struct GmmReport {
287    /// Last accepted model in canonical component order.
288    pub model: GmmModel,
289    /// Numerical, convergence, and selection evidence.
290    pub evidence: GmmEvidence,
291}
292
293/// Fits a regularized diagonal or full-covariance Gaussian mixture.
294///
295/// Responsibilities and likelihood are computed in the log domain. Candidate
296/// updates that decrease likelihood are rejected, and work exhaustion returns
297/// the last complete model rather than a partially updated parameter set.
298pub fn fit_gmm(
299    points: &[Vec<f64>],
300    spec: GmmSpec,
301    control: GmmControl,
302) -> Result<GmmReport, ClusteringError> {
303    let dimensions = validate_points(points)?;
304    spec.validate()?;
305    validate_components(points.len(), spec.components)?;
306    control.validate()?;
307
308    let mut meter = WorkMeter::new(control.max_work);
309    let global_covariance = global_covariance(points, spec.covariance, spec.regularization)?;
310    let mut random = SplitMix64::new(control.seed);
311    let means = kmeans_plus_plus(points, spec.components, &mut random, &mut meter)?;
312    let mut model = GmmModel {
313        weights: vec![1.0 / spec.components as f64; spec.components],
314        means,
315        covariances: vec![global_covariance.clone(); spec.components],
316    };
317    let mut state = expectation(points, &model, &mut meter)?;
318    let initial_log_likelihood = state.log_likelihood;
319    let mut history = vec![state.log_likelihood];
320    let mut iterations = 0;
321    let mut repairs = 0_u64;
322    let mut termination = GmmTermination::IterationLimit;
323
324    while iterations < control.max_iterations {
325        let candidate = maximize(points, &state, spec, &global_covariance, &mut meter);
326        let (candidate, candidate_repairs) = match candidate {
327            Ok(value) => value,
328            Err(ClusteringError::WorkLimit { .. }) => {
329                termination = GmmTermination::WorkLimit;
330                break;
331            }
332            Err(error) => return Err(error),
333        };
334        let next_state = match expectation(points, &candidate, &mut meter) {
335            Ok(value) => value,
336            Err(ClusteringError::WorkLimit { .. }) => {
337                termination = GmmTermination::WorkLimit;
338                break;
339            }
340            Err(error) => return Err(error),
341        };
342        let previous = state.log_likelihood;
343        let scale = previous.abs().max(1.0);
344        if next_state.log_likelihood + control.tolerance * scale < previous {
345            termination = GmmTermination::LikelihoodDecrease;
346            break;
347        }
348        model = candidate;
349        state = next_state;
350        repairs = repairs.saturating_add(candidate_repairs);
351        history.push(state.log_likelihood);
352        iterations += 1;
353        if (state.log_likelihood - previous).abs() <= control.tolerance * scale {
354            termination = GmmTermination::Converged;
355            break;
356        }
357    }
358
359    canonicalize_model(&mut model);
360    let model_selection = model_selection(
361        state.log_likelihood,
362        points.len(),
363        dimensions,
364        spec.components,
365        spec.covariance,
366    )?;
367    Ok(GmmReport {
368        model,
369        evidence: GmmEvidence {
370            initial_log_likelihood,
371            log_likelihood: state.log_likelihood,
372            likelihood_history: history,
373            iterations,
374            converged: termination == GmmTermination::Converged,
375            singular_component_repairs: repairs,
376            seed: control.seed,
377            work: meter.used,
378            termination,
379            model_selection,
380        },
381    })
382}
383
384struct ExpectationState {
385    responsibilities: Vec<Vec<f64>>,
386    point_log_likelihoods: Vec<f64>,
387    log_likelihood: f64,
388}
389
390fn expectation(
391    points: &[Vec<f64>],
392    model: &GmmModel,
393    meter: &mut WorkMeter,
394) -> Result<ExpectationState, ClusteringError> {
395    let dimensions = points[0].len();
396    let prepared = prepare_model(model, dimensions)?;
397    let component_cost =
398        prepared
399            .iter()
400            .map(PreparedCovariance::work)
401            .try_fold(0_u64, |sum, work| {
402                sum.checked_add(work?)
403                    .ok_or(ClusteringError::ArithmeticOverflow {
404                        operation: "GMM likelihood work",
405                    })
406            })?;
407    let point_count =
408        u64::try_from(points.len()).map_err(|_| ClusteringError::ArithmeticOverflow {
409            operation: "GMM point count",
410        })?;
411    meter.charge(component_cost.checked_mul(point_count).ok_or(
412        ClusteringError::ArithmeticOverflow {
413            operation: "GMM likelihood work",
414        },
415    )?)?;
416
417    let mut responsibilities = Vec::with_capacity(points.len());
418    let mut point_log_likelihoods = Vec::with_capacity(points.len());
419    let mut log_likelihood = 0.0;
420    for point in points {
421        let log_weights = model
422            .weights
423            .iter()
424            .zip(&model.means)
425            .zip(&prepared)
426            .map(|((&weight, mean), covariance)| weight.ln() + covariance.log_density(point, mean))
427            .collect::<Vec<_>>();
428        let normalizer = log_sum_exp(&log_weights);
429        if !normalizer.is_finite() {
430            return Err(ClusteringError::NumericalFailure {
431                operation: "GMM log-domain normalization",
432            });
433        }
434        responsibilities.push(
435            log_weights
436                .iter()
437                .map(|weight| (weight - normalizer).exp())
438                .collect(),
439        );
440        point_log_likelihoods.push(normalizer);
441        log_likelihood += normalizer;
442    }
443    if !log_likelihood.is_finite() {
444        return Err(ClusteringError::NumericalFailure {
445            operation: "GMM log likelihood",
446        });
447    }
448    Ok(ExpectationState {
449        responsibilities,
450        point_log_likelihoods,
451        log_likelihood,
452    })
453}
454
455fn maximize(
456    points: &[Vec<f64>],
457    state: &ExpectationState,
458    spec: GmmSpec,
459    global_covariance: &GaussianCovariance,
460    meter: &mut WorkMeter,
461) -> Result<(GmmModel, u64), ClusteringError> {
462    let dimensions = points[0].len();
463    let point_component_work =
464        checked_product(points.len(), spec.components, "GMM maximization work")?;
465    let dimension_work =
466        u64::try_from(dimensions).map_err(|_| ClusteringError::ArithmeticOverflow {
467            operation: "GMM dimensions",
468        })?;
469    meter.charge(point_component_work.checked_mul(dimension_work).ok_or(
470        ClusteringError::ArithmeticOverflow {
471            operation: "GMM maximization work",
472        },
473    )?)?;
474
475    let mut masses = vec![0.0; spec.components];
476    for row in &state.responsibilities {
477        for (mass, responsibility) in masses.iter_mut().zip(row) {
478            *mass += responsibility;
479        }
480    }
481    let minimum_mass = spec.singular_policy.minimum_weight() * points.len() as f64;
482    let mut means = vec![vec![0.0; dimensions]; spec.components];
483    for (point, row) in points.iter().zip(&state.responsibilities) {
484        for (component, &responsibility) in row.iter().enumerate() {
485            for (sum, &coordinate) in means[component].iter_mut().zip(point) {
486                *sum += responsibility * coordinate;
487            }
488        }
489    }
490    for (mean, &mass) in means.iter_mut().zip(&masses) {
491        if mass > minimum_mass && mass.is_finite() {
492            for coordinate in mean {
493                *coordinate /= mass;
494            }
495        }
496    }
497
498    let mut covariances = (0..spec.components)
499        .map(|component| {
500            component_covariance(
501                points,
502                &state.responsibilities,
503                component,
504                &means[component],
505                masses[component],
506                spec.covariance,
507                spec.regularization,
508            )
509        })
510        .collect::<Result<Vec<_>, _>>()?;
511    let mut repairs = 0_u64;
512    let mut used_points = vec![false; points.len()];
513    for component in 0..spec.components {
514        let singular_mass = !masses[component].is_finite() || masses[component] <= minimum_mass;
515        let singular_covariance = prepare_covariance(&covariances[component], dimensions).is_err();
516        if !singular_mass && !singular_covariance {
517            continue;
518        }
519        match spec.singular_policy {
520            SingularComponentPolicy::Fail { .. } => {
521                return Err(ClusteringError::SingularComponent { component });
522            }
523            SingularComponentPolicy::Reinitialize { .. } => {
524                let point = worst_unused_point(&state.point_log_likelihoods, &used_points);
525                used_points[point] = true;
526                means[component].clone_from(&points[point]);
527                covariances[component] = global_covariance.clone();
528                masses[component] = minimum_mass.max(1.0);
529                repairs += 1;
530            }
531        }
532    }
533    let total_mass = masses.iter().sum::<f64>();
534    if !total_mass.is_finite() || total_mass <= 0.0 {
535        return Err(ClusteringError::NumericalFailure {
536            operation: "GMM component weights",
537        });
538    }
539    let weights = masses.iter().map(|mass| mass / total_mass).collect();
540    Ok((
541        GmmModel {
542            weights,
543            means,
544            covariances,
545        },
546        repairs,
547    ))
548}
549
550fn component_covariance(
551    points: &[Vec<f64>],
552    responsibilities: &[Vec<f64>],
553    component: usize,
554    mean: &[f64],
555    mass: f64,
556    covariance: CovarianceType,
557    regularization: f64,
558) -> Result<GaussianCovariance, ClusteringError> {
559    let dimensions = mean.len();
560    if !mass.is_finite() || mass <= 0.0 {
561        return Ok(match covariance {
562            CovarianceType::Diagonal => GaussianCovariance::Diagonal(vec![0.0; dimensions]),
563            CovarianceType::Full => {
564                GaussianCovariance::Full(vec![vec![0.0; dimensions]; dimensions])
565            }
566        });
567    }
568    match covariance {
569        CovarianceType::Diagonal => {
570            let mut variances = vec![0.0; dimensions];
571            for (point, row) in points.iter().zip(responsibilities) {
572                for coordinate in 0..dimensions {
573                    let difference = point[coordinate] - mean[coordinate];
574                    variances[coordinate] += row[component] * difference * difference;
575                }
576            }
577            for variance in &mut variances {
578                *variance = (*variance / mass) + regularization;
579            }
580            require_finite_covariance(variances.iter().copied())?;
581            Ok(GaussianCovariance::Diagonal(variances))
582        }
583        CovarianceType::Full => {
584            let mut matrix = vec![vec![0.0; dimensions]; dimensions];
585            for (point, row) in points.iter().zip(responsibilities) {
586                for left in 0..dimensions {
587                    let left_difference = point[left] - mean[left];
588                    for right in 0..=left {
589                        let right_difference = point[right] - mean[right];
590                        matrix[left][right] += row[component] * left_difference * right_difference;
591                    }
592                }
593            }
594            for (left, row) in matrix.iter_mut().enumerate() {
595                for value in row.iter_mut().take(left + 1) {
596                    *value /= mass;
597                }
598                row[left] += regularization;
599            }
600            for left in 0..dimensions {
601                let (prior_rows, current_rows) = matrix.split_at_mut(left);
602                let row = &current_rows[0];
603                for (right, prior_row) in prior_rows.iter_mut().enumerate() {
604                    prior_row[left] = row[right];
605                }
606            }
607            require_finite_covariance(matrix.iter().flatten().copied())?;
608            Ok(GaussianCovariance::Full(matrix))
609        }
610    }
611}
612
613fn global_covariance(
614    points: &[Vec<f64>],
615    covariance: CovarianceType,
616    regularization: f64,
617) -> Result<GaussianCovariance, ClusteringError> {
618    let dimensions = points[0].len();
619    let mut mean = vec![0.0; dimensions];
620    for point in points {
621        for (sum, &coordinate) in mean.iter_mut().zip(point) {
622            *sum += coordinate;
623        }
624    }
625    for coordinate in &mut mean {
626        *coordinate /= points.len() as f64;
627    }
628    let responsibilities = vec![vec![1.0]; points.len()];
629    component_covariance(
630        points,
631        &responsibilities,
632        0,
633        &mean,
634        points.len() as f64,
635        covariance,
636        regularization,
637    )
638}
639
640fn worst_unused_point(log_likelihoods: &[f64], used: &[bool]) -> usize {
641    log_likelihoods
642        .iter()
643        .enumerate()
644        .filter(|(index, _)| !used[*index])
645        .min_by(|(left_index, left), (right_index, right)| {
646            left.total_cmp(right)
647                .then_with(|| left_index.cmp(right_index))
648        })
649        .map(|(index, _)| index)
650        .unwrap_or(0)
651}