Skip to main content

sim_lib_numbers_stats/
implementation.rs

1//! Implementation of the statistics helpers: descriptive statistics,
2//! probability, and fairness metrics over f64 data, with their error type.
3
4use std::{error::Error, fmt};
5
6#[path = "agent_fixtures.rs"]
7mod agent_fixtures;
8#[path = "claim.rs"]
9mod claim;
10#[path = "clustering.rs"]
11mod clustering;
12#[cfg(test)]
13#[path = "clustering_tests.rs"]
14mod clustering_tests;
15#[path = "function.rs"]
16mod function;
17#[path = "gmm.rs"]
18mod gmm;
19#[path = "gmm_math.rs"]
20mod gmm_math;
21#[path = "hmm_baum_welch.rs"]
22mod hmm_baum_welch;
23#[path = "hmm_fit.rs"]
24mod hmm_fit;
25#[path = "hmm_inference.rs"]
26mod hmm_inference;
27#[path = "hmm_model.rs"]
28mod hmm_model;
29#[cfg(test)]
30#[path = "hmm_tests.rs"]
31mod hmm_tests;
32#[path = "markov.rs"]
33mod markov;
34#[cfg(test)]
35#[path = "markov_tests.rs"]
36mod markov_tests;
37#[path = "quantile.rs"]
38mod quantile;
39#[cfg(test)]
40#[path = "quantile_tests.rs"]
41mod quantile_tests;
42#[path = "robust.rs"]
43mod robust;
44#[cfg(test)]
45#[path = "robust_tests.rs"]
46mod robust_tests;
47#[path = "runtime.rs"]
48mod runtime;
49#[path = "runtime_clustering.rs"]
50mod runtime_clustering;
51#[path = "transition.rs"]
52mod transition;
53
54pub use claim::{
55    FairnessClaimValue, FairnessEvidence, StatsClaimEvidence, StatsClaimValue, fairness_claim,
56    fairness_claim_value, stats_result_claim, stats_result_claim_value,
57};
58pub use clustering::{
59    ClusteringError, KMeansControl, KMeansModel, KMeansReport, KMeansRestartEvidence,
60    KMeansSearchTermination, KMeansTermination, fit_kmeans,
61};
62pub use function::{
63    StatsNumbersLib, stats_claims_symbol, stats_disparate_impact_claim_symbol,
64    stats_entropy_claim_symbol, stats_gmm_symbol, stats_kmeans_symbol, stats_mean_claim_symbol,
65    stats_variance_claim_symbol,
66};
67pub use gmm::{
68    CovarianceType, GaussianCovariance, GmmControl, GmmEvidence, GmmModel, GmmReport, GmmSpec,
69    GmmTermination, ModelSelectionEvidence, SingularComponentPolicy, fit_gmm,
70};
71pub use hmm_fit::{
72    HmmFitControl, HmmFitEvidence, HmmFitReport, HmmSpec, HmmTermination, Sequence, StateId,
73    fit_hmm,
74};
75pub use hmm_inference::{
76    ForwardBackward, InferenceEvidence, PosteriorPath, ViterbiPath, forward_backward,
77    posterior_decode, viterbi,
78};
79pub use hmm_model::{EmissionModel, HiddenMarkovModel, HmmError, HmmObservation};
80pub use markov::{
81    CorpusProvenance, MarkovError, MarkovModel, MarkovPolicy, ModelReport, TransitionScore,
82    fit_markov, fnv1a64,
83};
84pub use quantile::{
85    QuantileError, QuantileEstimate, QuantilePolicy, QuantileSketch, exact_quantile,
86};
87pub use robust::{
88    BootstrapControl, BootstrapEffectInterval, bootstrap_mean_difference_interval,
89    median_absolute_deviation,
90};
91pub use transition::{FiniteTransitionMatrix, TransitionError};
92
93const FOUR_FIFTHS_THRESHOLD: f64 = 0.8;
94const PROBABILITY_TOLERANCE: f64 = 1.0e-12;
95
96/// Result alias for the statistics helpers, fixing the error to [`StatsError`].
97pub type StatsResult<T> = Result<T, StatsError>;
98
99/// Errors returned by probability, statistics, and fairness helpers.
100///
101/// Each variant carries the `metric` name of the helper that rejected the
102/// input, so the failure can be reported without the caller tracking context.
103#[derive(Clone, Debug, PartialEq)]
104pub enum StatsError {
105    /// A helper requiring at least one value was given an empty slice.
106    EmptyInput {
107        /// Name of the helper that rejected the input.
108        metric: &'static str,
109    },
110    /// Fewer values were supplied than the helper requires.
111    InsufficientInput {
112        /// Name of the helper that rejected the input.
113        metric: &'static str,
114        /// Smallest number of values the helper accepts.
115        minimum: usize,
116        /// Number of values actually supplied.
117        actual: usize,
118    },
119    /// A value was not finite (`NaN` or infinite).
120    NonFinite {
121        /// Name of the helper that rejected the input.
122        metric: &'static str,
123        /// Position of the offending value, if it came from a slice.
124        index: Option<usize>,
125        /// The offending value.
126        value: f64,
127    },
128    /// A probability fell outside the closed range `0.0..=1.0`.
129    ProbabilityOutOfRange {
130        /// Name of the helper that rejected the input.
131        metric: &'static str,
132        /// Position of the offending value, if it came from a slice.
133        index: Option<usize>,
134        /// The offending value.
135        value: f64,
136    },
137    /// A probability vector did not sum to one within tolerance.
138    ProbabilityMass {
139        /// Name of the helper that rejected the input.
140        metric: &'static str,
141        /// The actual sum of the probabilities.
142        sum: f64,
143    },
144    /// A Bayesian update was given zero total evidence.
145    ZeroEvidence {
146        /// Name of the helper that rejected the input.
147        metric: &'static str,
148    },
149    /// A counts pair was given a zero total.
150    ZeroTotal {
151        /// Label of the count whose total was zero.
152        label: &'static str,
153    },
154    /// A fairness ratio was given a zero reference rate to divide by.
155    ZeroReferenceRate {
156        /// Name of the helper that rejected the input.
157        metric: &'static str,
158    },
159    /// A bounded statistics control contained an invalid field.
160    InvalidControl {
161        /// Name of the rejected control field.
162        field: &'static str,
163        /// Stable explanation of the accepted range.
164        reason: &'static str,
165    },
166    /// A requested computation exceeded its explicit work allowance.
167    WorkLimitExceeded {
168        /// Work units required by the request.
169        required: u64,
170        /// Work units admitted by the caller.
171        limit: u64,
172    },
173}
174
175impl fmt::Display for StatsError {
176    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
177        match self {
178            Self::EmptyInput { metric } => write!(f, "{metric} requires at least one value"),
179            Self::InsufficientInput {
180                metric,
181                minimum,
182                actual,
183            } => write!(
184                f,
185                "{metric} requires at least {minimum} values, got {actual}"
186            ),
187            Self::NonFinite {
188                metric,
189                index,
190                value,
191            } => match index {
192                Some(index) => write!(f, "{metric} value {index} is not finite: {value}"),
193                None => write!(f, "{metric} value is not finite: {value}"),
194            },
195            Self::ProbabilityOutOfRange {
196                metric,
197                index,
198                value,
199            } => match index {
200                Some(index) => write!(
201                    f,
202                    "{metric} probability {index} must be between 0 and 1, got {value}"
203                ),
204                None => write!(
205                    f,
206                    "{metric} probability must be between 0 and 1, got {value}"
207                ),
208            },
209            Self::ProbabilityMass { metric, sum } => {
210                write!(f, "{metric} probabilities must sum to 1, got {sum}")
211            }
212            Self::ZeroEvidence { metric } => write!(f, "{metric} evidence must be nonzero"),
213            Self::ZeroTotal { label } => write!(f, "{label} total must be nonzero"),
214            Self::ZeroReferenceRate { metric } => {
215                write!(f, "{metric} reference rate must be nonzero")
216            }
217            Self::InvalidControl { field, reason } => {
218                write!(f, "invalid statistics control {field}: {reason}")
219            }
220            Self::WorkLimitExceeded { required, limit } => write!(
221                f,
222                "statistics computation requires {required} work units, limit is {limit}"
223            ),
224        }
225    }
226}
227
228impl Error for StatsError {}
229
230/// Counts for a binary selection or outcome table.
231#[derive(Clone, Copy, Debug, PartialEq)]
232pub struct BinaryOutcomeCounts {
233    /// Number of selected (positive-outcome) cases.
234    pub selected: u64,
235    /// Total number of cases; the denominator of the selection rate.
236    pub total: u64,
237}
238
239impl BinaryOutcomeCounts {
240    /// Builds a count pair and rejects impossible or empty totals.
241    pub fn new(selected: u64, total: u64) -> StatsResult<Self> {
242        if total == 0 {
243            return Err(StatsError::ZeroTotal {
244                label: "outcome counts",
245            });
246        }
247        if selected > total {
248            return Err(StatsError::ProbabilityOutOfRange {
249                metric: "outcome counts",
250                index: None,
251                value: selected as f64 / total as f64,
252            });
253        }
254        Ok(Self { selected, total })
255    }
256
257    /// Returns `selected / total` as an f64 rate.
258    pub fn selection_rate(self) -> f64 {
259        self.selected as f64 / self.total as f64
260    }
261}
262
263/// Disparate-impact summary for two binary-outcome groups.
264#[derive(Clone, Copy, Debug, PartialEq)]
265pub struct DisparateImpact {
266    /// Selection rate of the reference group.
267    pub reference_rate: f64,
268    /// Selection rate of the comparison group.
269    pub comparison_rate: f64,
270    /// Comparison rate divided by reference rate.
271    pub ratio: f64,
272    /// Whether the ratio meets the four-fifths (0.8) threshold.
273    pub passes_four_fifths: bool,
274}
275
276/// Computes `prior * likelihood / evidence` and validates the posterior.
277///
278/// Every argument and the result must be a probability in `0.0..=1.0`, and
279/// `evidence` must be nonzero.
280///
281/// # Examples
282///
283/// ```
284/// use sim_lib_numbers_stats::bayesian_update;
285///
286/// let posterior = bayesian_update(0.2, 0.75, 0.3).unwrap();
287/// assert!((posterior - 0.5).abs() < 1e-12);
288/// ```
289pub fn bayesian_update(prior: f64, likelihood: f64, evidence: f64) -> StatsResult<f64> {
290    validate_probability("bayesian_update", None, prior)?;
291    validate_probability("bayesian_update", None, likelihood)?;
292    validate_probability("bayesian_update", None, evidence)?;
293    if evidence == 0.0 {
294        return Err(StatsError::ZeroEvidence {
295            metric: "bayesian_update",
296        });
297    }
298    let posterior = (prior * likelihood) / evidence;
299    validate_probability("bayesian_update", None, posterior)?;
300    Ok(posterior)
301}
302
303/// Computes a binary-test posterior from prior, true-positive, and false-positive rates.
304pub fn bayesian_update_binary(
305    prior: f64,
306    true_positive_rate: f64,
307    false_positive_rate: f64,
308) -> StatsResult<f64> {
309    validate_probability("bayesian_update_binary", None, prior)?;
310    validate_probability("bayesian_update_binary", None, true_positive_rate)?;
311    validate_probability("bayesian_update_binary", None, false_positive_rate)?;
312    let evidence = prior * true_positive_rate + (1.0 - prior) * false_positive_rate;
313    bayesian_update(prior, true_positive_rate, evidence)
314}
315
316/// Computes Shannon entropy in bits for a probability vector that sums to one.
317pub fn entropy(probabilities: &[f64]) -> StatsResult<f64> {
318    if probabilities.is_empty() {
319        return Err(StatsError::EmptyInput { metric: "entropy" });
320    }
321    let mut sum = 0.0;
322    let mut bits = 0.0;
323    for (index, probability) in probabilities.iter().copied().enumerate() {
324        validate_probability("entropy", Some(index), probability)?;
325        sum += probability;
326        if probability > 0.0 {
327            bits -= probability * probability.log2();
328        }
329    }
330    if (sum - 1.0).abs() > PROBABILITY_TOLERANCE {
331        return Err(StatsError::ProbabilityMass {
332            metric: "entropy",
333            sum,
334        });
335    }
336    Ok(bits)
337}
338
339/// Computes the arithmetic mean of finite values.
340///
341/// Returns [`StatsError::EmptyInput`] for an empty slice and
342/// [`StatsError::NonFinite`] if any value is `NaN` or infinite.
343///
344/// # Examples
345///
346/// ```
347/// use sim_lib_numbers_stats::mean;
348///
349/// assert_eq!(mean(&[2.0, 4.0, 6.0]).unwrap(), 4.0);
350/// ```
351pub fn mean(values: &[f64]) -> StatsResult<f64> {
352    validate_values("mean", values)?;
353    Ok(values.iter().sum::<f64>() / values.len() as f64)
354}
355
356/// Computes population variance.
357///
358/// Alias for [`population_variance`].
359///
360/// # Examples
361///
362/// ```
363/// use sim_lib_numbers_stats::variance;
364///
365/// let values = [2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0];
366/// assert!((variance(&values).unwrap() - 4.0).abs() < 1e-12);
367/// ```
368pub fn variance(values: &[f64]) -> StatsResult<f64> {
369    population_variance(values)
370}
371
372/// Computes population variance with divisor `n`.
373pub fn population_variance(values: &[f64]) -> StatsResult<f64> {
374    validate_values("population_variance", values)?;
375    let mean = mean(values)?;
376    Ok(values
377        .iter()
378        .map(|value| {
379            let delta = value - mean;
380            delta * delta
381        })
382        .sum::<f64>()
383        / values.len() as f64)
384}
385
386/// Computes sample variance with divisor `n - 1`.
387pub fn sample_variance(values: &[f64]) -> StatsResult<f64> {
388    validate_values("sample_variance", values)?;
389    if values.len() < 2 {
390        return Err(StatsError::InsufficientInput {
391            metric: "sample_variance",
392            minimum: 2,
393            actual: values.len(),
394        });
395    }
396    let mean = mean(values)?;
397    Ok(values
398        .iter()
399        .map(|value| {
400            let delta = value - mean;
401            delta * delta
402        })
403        .sum::<f64>()
404        / (values.len() - 1) as f64)
405}
406
407/// Computes the comparison/reference rate ratio used by the four-fifths rule.
408pub fn four_fifths_ratio(reference_rate: f64, comparison_rate: f64) -> StatsResult<f64> {
409    validate_probability("four_fifths_ratio", None, reference_rate)?;
410    validate_probability("four_fifths_ratio", None, comparison_rate)?;
411    if reference_rate == 0.0 {
412        return Err(StatsError::ZeroReferenceRate {
413            metric: "four_fifths_ratio",
414        });
415    }
416    Ok(comparison_rate / reference_rate)
417}
418
419/// Computes disparate impact and the four-fifths pass/fail flag.
420///
421/// Divides the comparison group's selection rate by the reference group's and
422/// flags whether the resulting ratio clears the four-fifths (0.8) threshold.
423///
424/// # Examples
425///
426/// ```
427/// use sim_lib_numbers_stats::{BinaryOutcomeCounts, disparate_impact};
428///
429/// let reference = BinaryOutcomeCounts::new(80, 100).unwrap();
430/// let comparison = BinaryOutcomeCounts::new(60, 100).unwrap();
431/// let impact = disparate_impact(reference, comparison).unwrap();
432///
433/// assert!((impact.ratio - 0.75).abs() < 1e-12);
434/// assert!(!impact.passes_four_fifths);
435/// ```
436pub fn disparate_impact(
437    reference: BinaryOutcomeCounts,
438    comparison: BinaryOutcomeCounts,
439) -> StatsResult<DisparateImpact> {
440    let reference_rate = reference.selection_rate();
441    let comparison_rate = comparison.selection_rate();
442    let ratio = four_fifths_ratio(reference_rate, comparison_rate)?;
443    Ok(DisparateImpact {
444        reference_rate,
445        comparison_rate,
446        ratio,
447        passes_four_fifths: ratio >= FOUR_FIFTHS_THRESHOLD,
448    })
449}
450
451pub(super) fn validate_values(metric: &'static str, values: &[f64]) -> StatsResult<()> {
452    if values.is_empty() {
453        return Err(StatsError::EmptyInput { metric });
454    }
455    for (index, value) in values.iter().copied().enumerate() {
456        validate_finite(metric, Some(index), value)?;
457    }
458    Ok(())
459}
460
461fn validate_probability(metric: &'static str, index: Option<usize>, value: f64) -> StatsResult<()> {
462    validate_finite(metric, index, value)?;
463    if !(0.0..=1.0).contains(&value) {
464        return Err(StatsError::ProbabilityOutOfRange {
465            metric,
466            index,
467            value,
468        });
469    }
470    Ok(())
471}
472
473fn validate_finite(metric: &'static str, index: Option<usize>, value: f64) -> StatsResult<()> {
474    if value.is_finite() {
475        Ok(())
476    } else {
477        Err(StatsError::NonFinite {
478            metric,
479            index,
480            value,
481        })
482    }
483}