1use 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
14pub enum CovarianceType {
15 Diagonal,
17 Full,
19}
20
21#[derive(Clone, Debug, PartialEq)]
23pub enum GaussianCovariance {
24 Diagonal(Vec<f64>),
26 Full(Vec<Vec<f64>>),
28}
29
30#[derive(Clone, Copy, Debug, PartialEq)]
32pub enum SingularComponentPolicy {
33 Reinitialize {
35 minimum_weight: f64,
37 },
38 Fail {
40 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#[derive(Clone, Copy, Debug, PartialEq)]
74pub struct GmmSpec {
75 pub components: usize,
77 pub covariance: CovarianceType,
79 pub regularization: f64,
81 pub singular_policy: SingularComponentPolicy,
83}
84
85impl GmmSpec {
86 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#[derive(Clone, Copy, Debug, PartialEq)]
122pub struct GmmControl {
123 pub seed: u64,
125 pub max_iterations: usize,
127 pub tolerance: f64,
129 pub max_work: u64,
131}
132
133impl GmmControl {
134 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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
186pub enum GmmTermination {
187 Converged,
189 IterationLimit,
191 WorkLimit,
193 LikelihoodDecrease,
195}
196
197#[derive(Clone, Copy, Debug, PartialEq)]
199pub struct ModelSelectionEvidence {
200 pub log_likelihood: f64,
202 pub parameters: usize,
204 pub aic: f64,
206 pub bic: f64,
208 pub observations: usize,
210}
211
212#[derive(Clone, Debug, PartialEq)]
214pub struct GmmModel {
215 pub weights: Vec<f64>,
217 pub means: Vec<Vec<f64>>,
219 pub covariances: Vec<GaussianCovariance>,
221}
222
223impl GmmModel {
224 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 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 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#[derive(Clone, Debug, PartialEq)]
261pub struct GmmEvidence {
262 pub initial_log_likelihood: f64,
264 pub log_likelihood: f64,
266 pub likelihood_history: Vec<f64>,
268 pub iterations: usize,
270 pub converged: bool,
272 pub singular_component_repairs: u64,
274 pub seed: u64,
276 pub work: u64,
278 pub termination: GmmTermination,
280 pub model_selection: ModelSelectionEvidence,
282}
283
284#[derive(Clone, Debug, PartialEq)]
286pub struct GmmReport {
287 pub model: GmmModel,
289 pub evidence: GmmEvidence,
291}
292
293pub 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 = ¤t_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}