1use 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
96pub type StatsResult<T> = Result<T, StatsError>;
98
99#[derive(Clone, Debug, PartialEq)]
104pub enum StatsError {
105 EmptyInput {
107 metric: &'static str,
109 },
110 InsufficientInput {
112 metric: &'static str,
114 minimum: usize,
116 actual: usize,
118 },
119 NonFinite {
121 metric: &'static str,
123 index: Option<usize>,
125 value: f64,
127 },
128 ProbabilityOutOfRange {
130 metric: &'static str,
132 index: Option<usize>,
134 value: f64,
136 },
137 ProbabilityMass {
139 metric: &'static str,
141 sum: f64,
143 },
144 ZeroEvidence {
146 metric: &'static str,
148 },
149 ZeroTotal {
151 label: &'static str,
153 },
154 ZeroReferenceRate {
156 metric: &'static str,
158 },
159 InvalidControl {
161 field: &'static str,
163 reason: &'static str,
165 },
166 WorkLimitExceeded {
168 required: u64,
170 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#[derive(Clone, Copy, Debug, PartialEq)]
232pub struct BinaryOutcomeCounts {
233 pub selected: u64,
235 pub total: u64,
237}
238
239impl BinaryOutcomeCounts {
240 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 pub fn selection_rate(self) -> f64 {
259 self.selected as f64 / self.total as f64
260 }
261}
262
263#[derive(Clone, Copy, Debug, PartialEq)]
265pub struct DisparateImpact {
266 pub reference_rate: f64,
268 pub comparison_rate: f64,
270 pub ratio: f64,
272 pub passes_four_fifths: bool,
274}
275
276pub 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
303pub 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
316pub 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
339pub fn mean(values: &[f64]) -> StatsResult<f64> {
352 validate_values("mean", values)?;
353 Ok(values.iter().sum::<f64>() / values.len() as f64)
354}
355
356pub fn variance(values: &[f64]) -> StatsResult<f64> {
369 population_variance(values)
370}
371
372pub 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
386pub 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
407pub 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
419pub 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}