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