use std::{error::Error, fmt};
#[path = "agent_fixtures.rs"]
mod agent_fixtures;
#[path = "claim.rs"]
mod claim;
#[path = "clustering.rs"]
mod clustering;
#[cfg(test)]
#[path = "clustering_tests.rs"]
mod clustering_tests;
#[path = "decision.rs"]
mod decision;
#[cfg(test)]
#[path = "decision_tests.rs"]
mod decision_tests;
#[path = "distribution.rs"]
mod distribution;
#[path = "function.rs"]
mod function;
#[path = "gmm.rs"]
mod gmm;
#[path = "gmm_math.rs"]
mod gmm_math;
#[path = "hmm_baum_welch.rs"]
mod hmm_baum_welch;
#[path = "hmm_fit.rs"]
mod hmm_fit;
#[path = "hmm_inference.rs"]
mod hmm_inference;
#[path = "hmm_model.rs"]
mod hmm_model;
#[cfg(test)]
#[path = "hmm_tests.rs"]
mod hmm_tests;
#[path = "markov.rs"]
mod markov;
#[cfg(test)]
#[path = "markov_tests.rs"]
mod markov_tests;
#[path = "parametric_distribution.rs"]
mod parametric_distribution;
#[path = "quantile.rs"]
mod quantile;
#[cfg(test)]
#[path = "quantile_tests.rs"]
mod quantile_tests;
#[path = "robust.rs"]
mod robust;
#[cfg(test)]
#[path = "robust_tests.rs"]
mod robust_tests;
#[path = "runtime.rs"]
mod runtime;
#[path = "runtime_clustering.rs"]
mod runtime_clustering;
#[path = "runtime_decision.rs"]
mod runtime_decision;
#[path = "sampling.rs"]
mod sampling;
#[cfg(test)]
#[path = "sampling_tests.rs"]
mod sampling_tests;
#[path = "transition.rs"]
mod transition;
pub use claim::{
FairnessClaimValue, FairnessEvidence, StatsClaimEvidence, StatsClaimValue, fairness_claim,
fairness_claim_value, stats_result_claim, stats_result_claim_value,
};
pub use clustering::{
ClusteringError, KMeansControl, KMeansModel, KMeansReport, KMeansRestartEvidence,
KMeansSearchTermination, KMeansTermination, fit_kmeans,
};
pub use decision::{
BinaryInterval, ClusterSample, IsotonicFit, IsotonicPoint, RegisteredLook,
RegisteredLookSequence, SequentialInterval, ThresholdReadout, clustered_bootstrap_interval,
exact_binary_interval, fit_isotonic, paired_bootstrap_interval,
};
pub use distribution::{
KsMethod, KsResult, MomentConvention, StandardizedMoments, kolmogorov_smirnov_one_sample,
kolmogorov_smirnov_two_sample, standardized_moments,
};
pub use function::{
StatsNumbersLib, stats_claims_symbol, stats_clustered_bootstrap_symbol,
stats_disparate_impact_claim_symbol, stats_entropy_claim_symbol,
stats_exact_binary_interval_symbol, stats_gmm_symbol, stats_isotonic_symbol,
stats_kmeans_symbol, stats_mean_claim_symbol, stats_paired_bootstrap_symbol,
stats_registered_look_symbol, stats_variance_claim_symbol,
};
pub use gmm::{
CovarianceType, GaussianCovariance, GmmControl, GmmEvidence, GmmModel, GmmReport, GmmSpec,
GmmTermination, ModelSelectionEvidence, SingularComponentPolicy, fit_gmm,
};
pub use hmm_fit::{
HmmFitControl, HmmFitEvidence, HmmFitReport, HmmSpec, HmmTermination, Sequence, StateId,
fit_hmm,
};
pub use hmm_inference::{
ForwardBackward, InferenceEvidence, PosteriorPath, ViterbiPath, forward_backward,
posterior_decode, viterbi,
};
pub use hmm_model::{EmissionModel, HiddenMarkovModel, HmmError, HmmObservation};
pub use markov::{
CorpusProvenance, MarkovError, MarkovModel, MarkovPolicy, ModelReport, TransitionScore,
fit_markov, fnv1a64,
};
pub use parametric_distribution::{
normal_cdf, normal_density, normal_quantile, normal_survival, student_t_cdf, student_t_density,
student_t_quantile, student_t_survival,
};
pub use quantile::{
QuantileError, QuantileEstimate, QuantilePolicy, QuantileSketch, exact_quantile,
};
pub use robust::{
BootstrapControl, BootstrapEffectInterval, bootstrap_mean_difference_interval,
median_absolute_deviation,
};
pub use sampling::{
CoverageEvidence, DesignError, LatinHypercubePlan, SampleDesign, SamplerAlgorithm,
SamplerReceipt, SamplerState, Scramble, SeededSampler, SobolPlan, SweepPlan, UntestedRegion,
};
pub use transition::{FiniteTransitionMatrix, TransitionError};
const FOUR_FIFTHS_THRESHOLD: f64 = 0.8;
const PROBABILITY_TOLERANCE: f64 = 1.0e-12;
pub type StatsResult<T> = Result<T, StatsError>;
#[derive(Clone, Debug, PartialEq)]
pub enum StatsError {
EmptyInput {
metric: &'static str,
},
InsufficientInput {
metric: &'static str,
minimum: usize,
actual: usize,
},
NonFinite {
metric: &'static str,
index: Option<usize>,
value: f64,
},
ProbabilityOutOfRange {
metric: &'static str,
index: Option<usize>,
value: f64,
},
ProbabilityMass {
metric: &'static str,
sum: f64,
},
ZeroEvidence {
metric: &'static str,
},
ZeroTotal {
label: &'static str,
},
ZeroReferenceRate {
metric: &'static str,
},
InvalidControl {
field: &'static str,
reason: &'static str,
},
WorkLimitExceeded {
required: u64,
limit: u64,
},
}
impl fmt::Display for StatsError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::EmptyInput { metric } => write!(f, "{metric} requires at least one value"),
Self::InsufficientInput {
metric,
minimum,
actual,
} => write!(
f,
"{metric} requires at least {minimum} values, got {actual}"
),
Self::NonFinite {
metric,
index,
value,
} => match index {
Some(index) => write!(f, "{metric} value {index} is not finite: {value}"),
None => write!(f, "{metric} value is not finite: {value}"),
},
Self::ProbabilityOutOfRange {
metric,
index,
value,
} => match index {
Some(index) => write!(
f,
"{metric} probability {index} must be between 0 and 1, got {value}"
),
None => write!(
f,
"{metric} probability must be between 0 and 1, got {value}"
),
},
Self::ProbabilityMass { metric, sum } => {
write!(f, "{metric} probabilities must sum to 1, got {sum}")
}
Self::ZeroEvidence { metric } => write!(f, "{metric} evidence must be nonzero"),
Self::ZeroTotal { label } => write!(f, "{label} total must be nonzero"),
Self::ZeroReferenceRate { metric } => {
write!(f, "{metric} reference rate must be nonzero")
}
Self::InvalidControl { field, reason } => {
write!(f, "invalid statistics control {field}: {reason}")
}
Self::WorkLimitExceeded { required, limit } => write!(
f,
"statistics computation requires {required} work units, limit is {limit}"
),
}
}
}
impl Error for StatsError {}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct BinaryOutcomeCounts {
pub selected: u64,
pub total: u64,
}
impl BinaryOutcomeCounts {
pub fn new(selected: u64, total: u64) -> StatsResult<Self> {
if total == 0 {
return Err(StatsError::ZeroTotal {
label: "outcome counts",
});
}
if selected > total {
return Err(StatsError::ProbabilityOutOfRange {
metric: "outcome counts",
index: None,
value: selected as f64 / total as f64,
});
}
Ok(Self { selected, total })
}
pub fn selection_rate(self) -> f64 {
self.selected as f64 / self.total as f64
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct DisparateImpact {
pub reference_rate: f64,
pub comparison_rate: f64,
pub ratio: f64,
pub passes_four_fifths: bool,
}
pub fn bayesian_update(prior: f64, likelihood: f64, evidence: f64) -> StatsResult<f64> {
validate_probability("bayesian_update", None, prior)?;
validate_probability("bayesian_update", None, likelihood)?;
validate_probability("bayesian_update", None, evidence)?;
if evidence == 0.0 {
return Err(StatsError::ZeroEvidence {
metric: "bayesian_update",
});
}
let posterior = (prior * likelihood) / evidence;
validate_probability("bayesian_update", None, posterior)?;
Ok(posterior)
}
pub fn bayesian_update_binary(
prior: f64,
true_positive_rate: f64,
false_positive_rate: f64,
) -> StatsResult<f64> {
validate_probability("bayesian_update_binary", None, prior)?;
validate_probability("bayesian_update_binary", None, true_positive_rate)?;
validate_probability("bayesian_update_binary", None, false_positive_rate)?;
let evidence = prior * true_positive_rate + (1.0 - prior) * false_positive_rate;
bayesian_update(prior, true_positive_rate, evidence)
}
pub fn entropy(probabilities: &[f64]) -> StatsResult<f64> {
if probabilities.is_empty() {
return Err(StatsError::EmptyInput { metric: "entropy" });
}
let mut sum = 0.0;
let mut bits = 0.0;
for (index, probability) in probabilities.iter().copied().enumerate() {
validate_probability("entropy", Some(index), probability)?;
sum += probability;
if probability > 0.0 {
bits -= probability * probability.log2();
}
}
if (sum - 1.0).abs() > PROBABILITY_TOLERANCE {
return Err(StatsError::ProbabilityMass {
metric: "entropy",
sum,
});
}
Ok(bits)
}
pub fn mean(values: &[f64]) -> StatsResult<f64> {
validate_values("mean", values)?;
Ok(values.iter().sum::<f64>() / values.len() as f64)
}
pub fn variance(values: &[f64]) -> StatsResult<f64> {
population_variance(values)
}
pub fn population_variance(values: &[f64]) -> StatsResult<f64> {
validate_values("population_variance", values)?;
let mean = mean(values)?;
Ok(values
.iter()
.map(|value| {
let delta = value - mean;
delta * delta
})
.sum::<f64>()
/ values.len() as f64)
}
pub fn sample_variance(values: &[f64]) -> StatsResult<f64> {
validate_values("sample_variance", values)?;
if values.len() < 2 {
return Err(StatsError::InsufficientInput {
metric: "sample_variance",
minimum: 2,
actual: values.len(),
});
}
let mean = mean(values)?;
Ok(values
.iter()
.map(|value| {
let delta = value - mean;
delta * delta
})
.sum::<f64>()
/ (values.len() - 1) as f64)
}
pub fn four_fifths_ratio(reference_rate: f64, comparison_rate: f64) -> StatsResult<f64> {
validate_probability("four_fifths_ratio", None, reference_rate)?;
validate_probability("four_fifths_ratio", None, comparison_rate)?;
if reference_rate == 0.0 {
return Err(StatsError::ZeroReferenceRate {
metric: "four_fifths_ratio",
});
}
Ok(comparison_rate / reference_rate)
}
pub fn disparate_impact(
reference: BinaryOutcomeCounts,
comparison: BinaryOutcomeCounts,
) -> StatsResult<DisparateImpact> {
let reference_rate = reference.selection_rate();
let comparison_rate = comparison.selection_rate();
let ratio = four_fifths_ratio(reference_rate, comparison_rate)?;
Ok(DisparateImpact {
reference_rate,
comparison_rate,
ratio,
passes_four_fifths: ratio >= FOUR_FIFTHS_THRESHOLD,
})
}
pub(super) fn validate_values(metric: &'static str, values: &[f64]) -> StatsResult<()> {
if values.is_empty() {
return Err(StatsError::EmptyInput { metric });
}
for (index, value) in values.iter().copied().enumerate() {
validate_finite(metric, Some(index), value)?;
}
Ok(())
}
fn validate_probability(metric: &'static str, index: Option<usize>, value: f64) -> StatsResult<()> {
validate_finite(metric, index, value)?;
if !(0.0..=1.0).contains(&value) {
return Err(StatsError::ProbabilityOutOfRange {
metric,
index,
value,
});
}
Ok(())
}
fn validate_finite(metric: &'static str, index: Option<usize>, value: f64) -> StatsResult<()> {
if value.is_finite() {
Ok(())
} else {
Err(StatsError::NonFinite {
metric,
index,
value,
})
}
}