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