use crate::algorithms::count_to_f64;
use crate::error::{Error, Result};
use crate::likelihood::BinomialLikelihood;
use crate::likelihood::{LogLikelihood, MleFit};
use crate::special::ln_choose;
const INTEGRALITY_TOL: f64 = 1e-9;
const FLOOR_SEARCH_TOP_BIT: usize = 1 << 40;
impl BinomialLikelihood {
pub fn fit(&self, data: &[f64]) -> Result<MleFit> {
if data.is_empty() {
return Err(Error::InsufficientData);
}
let n = self.trials().ok_or_else(|| {
Error::InvalidInput("number_of_trials must be a positive integer".to_owned())
})?;
let n_f = count_to_f64(n);
let mut sum_x = 0.0;
for &x in data {
let k = success_count(x, n).ok_or_else(|| {
Error::InvalidInput(format!(
"observation {x} is not an integer success count in [0, {n}]"
))
})?;
sum_x += count_to_f64(k);
}
let total = count_to_f64(data.len()) * n_f;
let p_hat = sum_x / total;
if !(p_hat > 0.0 && p_hat < 1.0) {
return Err(Error::DegenerateInput(format!(
"all observations on the boundary (p_hat = {p_hat}); \
the binomial log-likelihood is degenerate"
)));
}
let log_likelihood = self.log_likelihood(&[p_hat], data);
Ok(MleFit::from_closed_form(
vec![p_hat],
log_likelihood,
data.len(),
))
}
fn trials(&self) -> Option<usize> {
usize::try_from(self.number_of_trials)
.ok()
.filter(|&n| n > 0)
}
}
impl LogLikelihood for BinomialLikelihood {
fn n_params(&self) -> usize {
1
}
fn log_likelihood(&self, params: &[f64], data: &[f64]) -> f64 {
let Some(&p) = params.first() else {
return f64::NEG_INFINITY;
};
let Some(n) = self.trials() else {
return f64::NEG_INFINITY;
};
if !(p > 0.0 && p < 1.0) {
return f64::NEG_INFINITY;
}
let ln_p = p.ln();
let ln_q = (1.0 - p).ln();
let n_f = count_to_f64(n);
let mut sum = 0.0;
for &x in data {
let Some(k) = success_count(x, n) else {
return f64::NEG_INFINITY;
};
let k_f = count_to_f64(k);
sum += ln_choose(n, k) + k_f.mul_add(ln_p, (n_f - k_f) * ln_q);
}
sum
}
}
fn success_count(x: f64, n_trials: usize) -> Option<usize> {
if !x.is_finite() || x < 0.0 {
return None;
}
let rounded = x.round();
if (x - rounded).abs() > INTEGRALITY_TOL {
return None;
}
let k = f64_integer_to_usize(rounded)?;
(k <= n_trials).then_some(k)
}
fn f64_integer_to_usize(x: f64) -> Option<usize> {
if x < 0.0 {
return None;
}
let mut acc: usize = 0;
let mut bit: usize = FLOOR_SEARCH_TOP_BIT;
while bit > 0 {
let candidate = acc | bit;
if count_to_f64(candidate) <= x {
acc = candidate;
}
bit >>= 1;
}
((count_to_f64(acc) - x).abs() < 0.5).then_some(acc)
}
#[cfg(test)]
mod tests;