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 = "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
87pub type StatsResult<T> = Result<T, StatsError>;
89
90#[derive(Clone, Debug, PartialEq)]
95pub enum StatsError {
96 EmptyInput {
98 metric: &'static str,
100 },
101 InsufficientInput {
103 metric: &'static str,
105 minimum: usize,
107 actual: usize,
109 },
110 NonFinite {
112 metric: &'static str,
114 index: Option<usize>,
116 value: f64,
118 },
119 ProbabilityOutOfRange {
121 metric: &'static str,
123 index: Option<usize>,
125 value: f64,
127 },
128 ProbabilityMass {
130 metric: &'static str,
132 sum: f64,
134 },
135 ZeroEvidence {
137 metric: &'static str,
139 },
140 ZeroTotal {
142 label: &'static str,
144 },
145 ZeroReferenceRate {
147 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#[derive(Clone, Copy, Debug, PartialEq)]
202pub struct BinaryOutcomeCounts {
203 pub selected: u64,
205 pub total: u64,
207}
208
209impl BinaryOutcomeCounts {
210 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 pub fn selection_rate(self) -> f64 {
229 self.selected as f64 / self.total as f64
230 }
231}
232
233#[derive(Clone, Copy, Debug, PartialEq)]
235pub struct DisparateImpact {
236 pub reference_rate: f64,
238 pub comparison_rate: f64,
240 pub ratio: f64,
242 pub passes_four_fifths: bool,
244}
245
246pub 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
273pub 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
286pub 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
309pub fn mean(values: &[f64]) -> StatsResult<f64> {
322 validate_values("mean", values)?;
323 Ok(values.iter().sum::<f64>() / values.len() as f64)
324}
325
326pub fn variance(values: &[f64]) -> StatsResult<f64> {
339 population_variance(values)
340}
341
342pub 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
356pub 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
377pub 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
389pub 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}