use num_bigint::BigInt;
use num_traits::{One, Signed, Zero};
use crate::api::context::Context;
use crate::api::expr::Ex;
use crate::base::errors::SymplexError;
use crate::base::interval::Interval;
use super::data::{self, Ddof, Q};
use super::family::Distribution;
const OP: &str = "stats::estimation";
fn invalid(reason: impl Into<String>) -> SymplexError {
SymplexError::invalid_argument(OP, reason)
}
fn qu(n: usize) -> Q {
Q::from_integer(BigInt::from(n))
}
fn q64(n: u64) -> Q {
Q::from_integer(BigInt::from(n))
}
fn nonempty(data: &[Q], what: &str) -> Result<(), SymplexError> {
if data.is_empty() {
return Err(invalid(format!("{what} needs at least one observation")));
}
Ok(())
}
fn require_positive_data(data: &[Q], what: &str) -> Result<(), SymplexError> {
nonempty(data, what)?;
if let Some(x) = data.iter().find(|x| !x.is_positive()) {
return Err(invalid(format!(
"{what} needs positive observations, got {x}"
)));
}
Ok(())
}
fn require_counts(data: &[Q], what: &str) -> Result<(), SymplexError> {
nonempty(data, what)?;
if let Some(x) = data.iter().find(|x| !x.is_integer() || x.is_negative()) {
return Err(invalid(format!(
"{what} needs non-negative integer observations, got {x}"
)));
}
Ok(())
}
fn require_positive_param(x: &Q, what: &str) -> Result<(), SymplexError> {
if !x.is_positive() {
return Err(invalid(format!("{what} must be positive, got {x}")));
}
Ok(())
}
fn sqrt_q(ctx: &Context, v: Q) -> Ex {
ctx.from_ratio(v).sqrt().simplify()
}
fn positive_root(v: &Ex) -> Ex {
let s = v.sqrt().simplify();
if let [u] = s.args().as_slice()
&& s == u.abs()
&& let Ok(value) = u.eval_f64()
{
if value > 0.0 {
return u.clone();
}
if value < 0.0 {
return (-u).simplify();
}
}
s
}
fn ln_q(ctx: &Context, x: &Q) -> Ex {
let side = |n: &BigInt, sign: i64| -> Ex {
let (factors, cofactor) = crate::domains::ntheory::factorint_bounded(n, 64);
let mut acc = ctx.zero();
for (p, e) in factors {
acc += ctx.int(sign * i64::from(e)) * ctx.from_bigint(p).ln();
}
if cofactor > BigInt::one() {
acc += ctx.int(sign) * ctx.from_bigint(cofactor).ln();
}
acc
};
side(x.numer(), 1) + side(x.denom(), -1)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum FamilyKind {
Normal,
Exponential,
Poisson,
Bernoulli,
Binomial {
n: u64,
},
Geometric,
Uniform,
LogNormal,
Gamma,
Beta,
NegativeBinomial,
}
pub fn fit(ctx: &Context, family: FamilyKind, data: &[Q]) -> Result<Distribution, SymplexError> {
match family {
FamilyKind::Normal => fit_normal(ctx, data),
FamilyKind::Exponential => fit_exponential(ctx, data),
FamilyKind::Poisson => fit_poisson(ctx, data),
FamilyKind::Bernoulli => fit_bernoulli(ctx, data),
FamilyKind::Binomial { n } => fit_binomial_p(ctx, n, data),
FamilyKind::Geometric => fit_geometric(ctx, data),
FamilyKind::Uniform => fit_uniform(ctx, data),
FamilyKind::LogNormal => fit_log_normal(ctx, data),
FamilyKind::Gamma | FamilyKind::Beta | FamilyKind::NegativeBinomial => {
Err(invalid(format!(
"the {family:?} maximum-likelihood estimate has no closed form; use method_of_moments"
)))
}
}
}
pub fn method_of_moments(
ctx: &Context,
family: FamilyKind,
data: &[Q],
) -> Result<Distribution, SymplexError> {
match family {
FamilyKind::Normal => fit_normal(ctx, data),
FamilyKind::Exponential => fit_exponential(ctx, data),
FamilyKind::Poisson => fit_poisson(ctx, data),
FamilyKind::Bernoulli => fit_bernoulli(ctx, data),
FamilyKind::Binomial { n } => fit_binomial_p(ctx, n, data),
FamilyKind::Geometric => fit_geometric(ctx, data),
FamilyKind::Uniform => fit_uniform_moments(ctx, data),
FamilyKind::LogNormal => fit_log_normal_moments(ctx, data),
FamilyKind::Gamma => fit_gamma_moments(ctx, data),
FamilyKind::Beta => fit_beta_moments(ctx, data),
FamilyKind::NegativeBinomial => fit_negative_binomial_moments(ctx, data),
}
}
pub fn fit_normal(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
nonempty(data, "fit_normal")?;
let mean = data::mean(data)?;
let var = data::variance(data, Ddof::Population)?;
if var.is_zero() {
return Err(invalid("fit_normal: constant data give σ̂ = 0"));
}
Distribution::try_normal(ctx.from_ratio(mean), sqrt_q(ctx, var))
}
pub fn fit_exponential(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
require_positive_data(data, "fit_exponential")?;
let rate = data::mean(data)?.recip();
Distribution::try_exponential(ctx.from_ratio(rate))
}
pub fn fit_poisson(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
require_counts(data, "fit_poisson")?;
let rate = data::mean(data)?;
if rate.is_zero() {
return Err(invalid("fit_poisson: all-zero counts give λ̂ = 0"));
}
Distribution::try_poisson(ctx.from_ratio(rate))
}
pub fn fit_bernoulli(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
nonempty(data, "fit_bernoulli")?;
if let Some(x) = data.iter().find(|x| !x.is_zero() && !x.is_one()) {
return Err(invalid(format!(
"fit_bernoulli needs observations in {{0, 1}}, got {x}"
)));
}
Distribution::try_bernoulli(ctx.from_ratio(data::mean(data)?))
}
pub fn fit_binomial_p(ctx: &Context, n: u64, data: &[Q]) -> Result<Distribution, SymplexError> {
if n == 0 {
return Err(invalid("fit_binomial_p needs at least one trial"));
}
require_counts(data, "fit_binomial_p")?;
let nq = q64(n);
if let Some(x) = data.iter().find(|x| **x > nq) {
return Err(invalid(format!(
"fit_binomial_p: observation {x} exceeds the number of trials {n}"
)));
}
let p = data::mean(data)? / &nq;
Distribution::try_binomial(ctx.from_ratio(nq), ctx.from_ratio(p))
}
pub fn fit_geometric(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
require_counts(data, "fit_geometric")?;
if let Some(x) = data.iter().find(|x| x.is_zero()) {
return Err(invalid(format!(
"fit_geometric needs observations ≥ 1 (trials up to the first success), got {x}"
)));
}
Distribution::try_geometric(ctx.from_ratio(data::mean(data)?.recip()))
}
pub fn fit_uniform(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
nonempty(data, "fit_uniform")?;
let range = data::min_max(data)?;
if range.lower == range.upper {
return Err(invalid(
"fit_uniform: constant data give a zero-width interval",
));
}
Distribution::try_uniform(ctx.from_ratio(range.lower), ctx.from_ratio(range.upper))
}
pub fn fit_log_normal(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
require_positive_data(data, "fit_log_normal")?;
let range = data::min_max(data)?;
if range.lower == range.upper {
return Err(invalid("fit_log_normal: constant data give σ̂ = 0"));
}
let n = ctx.int(data.len() as i64);
let logs: Vec<Ex> = data.iter().map(|x| ln_q(ctx, x)).collect();
let mu = (logs.iter().fold(ctx.zero(), |acc, l| acc + l) / &n).simplify();
let var = (logs
.iter()
.fold(ctx.zero(), |acc, l| acc + (l - &mu).powi(2))
/ &n)
.simplify();
Distribution::try_log_normal(mu, positive_root(&var))
}
fn two_moments(data: &[Q], what: &str) -> Result<(Q, Q), SymplexError> {
nonempty(data, what)?;
let m = data::mean(data)?;
let v = data::variance(data, Ddof::Population)?;
if v.is_zero() {
return Err(invalid(format!(
"{what}: constant data have zero variance; the moment equations are degenerate"
)));
}
Ok((m, v))
}
pub fn fit_gamma_moments(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
require_positive_data(data, "fit_gamma_moments")?;
let (m, v) = two_moments(data, "fit_gamma_moments")?;
let shape = &m * &m / &v;
let scale = &v / &m;
Distribution::try_gamma(ctx.from_ratio(shape), ctx.from_ratio(scale))
}
pub fn fit_beta_moments(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
nonempty(data, "fit_beta_moments")?;
if let Some(x) = data.iter().find(|x| !x.is_positive() || **x >= Q::one()) {
return Err(invalid(format!(
"fit_beta_moments needs observations in (0, 1), got {x}"
)));
}
let (m, v) = two_moments(data, "fit_beta_moments")?;
let one_minus = Q::one() - &m;
let c = &m * &one_minus / &v - Q::one();
if !c.is_positive() {
return Err(invalid(
"fit_beta_moments: the variance is too large for a Beta law (x̄(1 − x̄)/s² ≤ 1)",
));
}
Distribution::try_beta(ctx.from_ratio(&m * &c), ctx.from_ratio(one_minus * c))
}
pub fn fit_negative_binomial_moments(
ctx: &Context,
data: &[Q],
) -> Result<Distribution, SymplexError> {
require_counts(data, "fit_negative_binomial_moments")?;
let (m, v) = two_moments(data, "fit_negative_binomial_moments")?;
if !m.is_positive() {
return Err(invalid("fit_negative_binomial_moments: all-zero counts"));
}
if v <= m {
return Err(invalid(format!(
"fit_negative_binomial_moments needs over-dispersed counts (s² > x̄), got s² = {v}, x̄ = {m}"
)));
}
let p = &m / &v;
let r = &m * &m / (&v - &m);
Distribution::try_negative_binomial(ctx.from_ratio(r), ctx.from_ratio(p))
}
pub fn fit_uniform_moments(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
let (m, v) = two_moments(data, "fit_uniform_moments")?;
let half_width = sqrt_q(ctx, v * qu(3));
let m = ctx.from_ratio(m);
Distribution::try_uniform((&m - &half_width).simplify(), (m + half_width).simplify())
}
pub fn fit_log_normal_moments(ctx: &Context, data: &[Q]) -> Result<Distribution, SymplexError> {
require_positive_data(data, "fit_log_normal_moments")?;
let (m, v) = two_moments(data, "fit_log_normal_moments")?;
let var = ctx.from_ratio(Q::one() + &v / (&m * &m)).ln();
let mu = (ctx.from_ratio(m).ln() - &var / ctx.int(2)).simplify();
Distribution::try_log_normal(mu, positive_root(&var))
}
pub fn log_likelihood(dist: &Distribution, data: &[Q]) -> Ex {
let ctx = dist.context();
data.iter()
.fold(ctx.zero(), |acc, x| {
acc + dist.density(&ctx.from_ratio(x.clone())).ln()
})
.expand_log()
.simplify()
}
pub fn aic(log_lik: &Ex, k: usize) -> Ex {
let ctx = log_lik.context();
(ctx.int(2 * k as i64) - ctx.int(2) * log_lik).simplify()
}
pub fn bic(log_lik: &Ex, k: usize, n: usize) -> Ex {
let ctx = log_lik.context();
(ctx.int(k as i64) * ctx.int(n as i64).ln() - ctx.int(2) * log_lik).simplify()
}
pub fn beta_binomial_posterior(
ctx: &Context,
alpha: &Q,
beta: &Q,
successes: u64,
failures: u64,
) -> Result<Distribution, SymplexError> {
require_positive_param(alpha, "the prior α")?;
require_positive_param(beta, "the prior β")?;
Distribution::try_beta(
ctx.from_ratio(alpha + q64(successes)),
ctx.from_ratio(beta + q64(failures)),
)
}
pub fn gamma_poisson_posterior(
ctx: &Context,
shape: &Q,
scale: &Q,
counts: &[u64],
) -> Result<Distribution, SymplexError> {
require_positive_param(shape, "the prior shape")?;
require_positive_param(scale, "the prior scale")?;
let total: Q = counts.iter().fold(Q::zero(), |acc, &c| acc + q64(c));
let n = qu(counts.len());
let post_scale = scale / (Q::one() + n * scale);
Distribution::try_gamma(ctx.from_ratio(shape + total), ctx.from_ratio(post_scale))
}
pub fn normal_known_variance_posterior(
ctx: &Context,
prior_mean: &Q,
prior_sd: &Q,
sigma: &Q,
data: &[Q],
) -> Result<Distribution, SymplexError> {
let (mu0, sigma0) = (prior_mean, prior_sd);
require_positive_param(sigma0, "the prior standard deviation")?;
require_positive_param(sigma, "the observation standard deviation")?;
nonempty(data, "normal_known_variance_posterior")?;
let prior_prec = (sigma0 * sigma0).recip();
let obs_prec = (sigma * sigma).recip();
let n = qu(data.len());
let precision = &prior_prec + &n * &obs_prec;
let mean = (mu0 * &prior_prec + data::sum(data) * &obs_prec) / &precision;
Distribution::try_normal(ctx.from_ratio(mean), sqrt_q(ctx, precision.recip()))
}
pub fn dirichlet_posterior_alphas(prior: &[Q], counts: &[u64]) -> Result<Vec<Q>, SymplexError> {
if prior.is_empty() {
return Err(invalid("a Dirichlet prior needs at least one category"));
}
if prior.len() != counts.len() {
return Err(invalid(format!(
"the Dirichlet prior has {} categories but {} counts were given",
prior.len(),
counts.len()
)));
}
for a in prior {
require_positive_param(a, "every prior α")?;
}
Ok(prior.iter().zip(counts).map(|(a, &c)| a + q64(c)).collect())
}
pub fn dirichlet_multinomial_posterior(
prior: &[Q],
counts: &[u64],
) -> Result<Vec<Q>, SymplexError> {
let alphas = dirichlet_posterior_alphas(prior, counts)?;
let total = data::sum(&alphas);
Ok(alphas.into_iter().map(|a| a / &total).collect())
}
pub fn credible_interval(
dist: &Distribution,
confidence: f64,
) -> Result<Interval<f64>, SymplexError> {
if !(confidence > 0.0 && confidence < 1.0) {
return Err(invalid(format!(
"the credible level must lie strictly between 0 and 1, got {confidence}"
)));
}
let tail = (1.0 - confidence) / 2.0;
Ok(Interval::closed(
dist.quantile_f64(tail)?,
dist.quantile_f64(1.0 - tail)?,
))
}
pub fn posterior_predictive_beta_binomial(
ctx: &Context,
alpha: &Q,
beta: &Q,
n: u64,
) -> Result<Distribution, SymplexError> {
require_positive_param(alpha, "α")?;
require_positive_param(beta, "β")?;
let n_usize = usize::try_from(n)
.map_err(|_| invalid(format!("the number of trials {n} is too large")))?;
let denom = rising_factorial(&(alpha + beta), n_usize);
let table: Vec<(Ex, Ex)> = (0..=n_usize)
.map(|k| {
let mass = data::binomial_q(n_usize, k)
* rising_factorial(alpha, k)
* rising_factorial(beta, n_usize - k)
/ &denom;
(ctx.int(k as i64), ctx.from_ratio(mass))
})
.collect();
Distribution::try_finite(ctx, table)
}
fn rising_factorial(a: &Q, m: usize) -> Q {
(0..m).fold(Q::one(), |acc, i| acc * (a + qu(i)))
}
pub fn standard_error_mean(ctx: &Context, sigma: &Q, n: usize) -> Result<Ex, SymplexError> {
require_positive_param(sigma, "σ")?;
if n == 0 {
return Err(invalid("the standard error needs at least one observation"));
}
Ok((ctx.from_ratio(sigma.clone()) / ctx.int(n as i64).sqrt()).simplify())
}
pub fn confidence_interval_mean_z(
ctx: &Context,
data: &[Q],
sigma: &Q,
confidence: f64,
) -> Result<Interval<f64>, SymplexError> {
nonempty(data, "confidence_interval_mean_z")?;
if !(confidence > 0.0 && confidence < 1.0) {
return Err(invalid(format!(
"the confidence level must lie strictly between 0 and 1, got {confidence}"
)));
}
let se = standard_error_mean(ctx, sigma, data.len())?.eval_f64()?;
let mean = ctx.from_ratio(data::mean(data)?).eval_f64()?;
let z =
Distribution::normal(ctx.zero(), ctx.one()).quantile_f64(1.0 - (1.0 - confidence) / 2.0)?;
Ok(Interval::closed(mean - z * se, mean + z * se))
}