use crate::algorithms::count_to_f64;
use crate::error::{Error, Result};
use crate::likelihood::{LogLikelihood, MleFit};
use crate::special::ln_gamma;
fn is_count(x: f64) -> bool {
x.is_finite() && x >= 0.0 && (x - x.round()).abs() <= 0.0
}
impl crate::likelihood::PoissonLikelihood {
pub fn fit(&self, data: &[f64]) -> Result<MleFit> {
if data.is_empty() {
return Err(Error::InsufficientData);
}
if data.iter().any(|&x| !is_count(x)) {
return Err(Error::InvalidInput(
"Poisson data must be finite non-negative integer counts".to_owned(),
));
}
let sum: f64 = data.iter().sum();
let lambda_hat = sum / count_to_f64(data.len());
if lambda_hat <= 0.0 {
return Err(Error::DegenerateInput(
"all-zero counts drive lambda_hat to 0, where the log-likelihood is undefined"
.to_owned(),
));
}
let log_likelihood = self.log_likelihood(&[lambda_hat], data);
Ok(MleFit::from_closed_form(
vec![lambda_hat],
log_likelihood,
data.len(),
))
}
}
impl LogLikelihood for crate::likelihood::PoissonLikelihood {
fn n_params(&self) -> usize {
1
}
fn log_likelihood(&self, params: &[f64], data: &[f64]) -> f64 {
let Some(&lambda) = params.first() else {
return f64::NEG_INFINITY;
};
if lambda.is_nan() || lambda <= 0.0 {
return f64::NEG_INFINITY;
}
if data.iter().any(|&x| !is_count(x)) {
return f64::NEG_INFINITY;
}
let ln_lambda = lambda.ln();
data.iter()
.map(|&x| x.mul_add(ln_lambda, -lambda) - ln_gamma(x + 1.0))
.sum()
}
}
#[cfg(test)]
mod tests {
use crate::error::Error;
use crate::likelihood::LogLikelihood;
use crate::likelihood::PoissonLikelihood;
#[test]
fn fit_empty_data_is_insufficient() {
let model = PoissonLikelihood::default();
assert!(
matches!(model.fit(&[]), Err(Error::InsufficientData)),
"empty data must be rejected"
);
}
#[test]
fn fit_negative_count_is_invalid() {
let model = PoissonLikelihood::default();
assert!(
matches!(model.fit(&[1.0, -2.0, 3.0]), Err(Error::InvalidInput(_))),
"a negative count must be rejected"
);
}
#[test]
fn fit_non_integer_count_is_invalid() {
let model = PoissonLikelihood::default();
assert!(
matches!(model.fit(&[1.0, 2.5, 3.0]), Err(Error::InvalidInput(_))),
"a fractional count must be rejected"
);
}
#[test]
fn fit_all_zero_counts_is_degenerate() {
let model = PoissonLikelihood::default();
assert!(
matches!(model.fit(&[0.0, 0.0, 0.0]), Err(Error::DegenerateInput(_))),
"all-zero counts must be rejected"
);
}
const DATA: [f64; 8] = [2.0, 3.0, 1.0, 5.0, 0.0, 4.0, 2.0, 3.0];
#[test]
fn log_likelihood_matches_scipy_at_lambda_3() {
let model = PoissonLikelihood::default();
let got = model.log_likelihood(&[3.0], &DATA);
let want = -14.963_113_099_343_797;
assert!(
(got - want).abs() <= 1e-10 * want.abs(),
"logL(3.0) was {got}, want {want}"
);
}
#[test]
fn log_likelihood_matches_scipy_at_lambda_1() {
let model = PoissonLikelihood::default();
let got = model.log_likelihood(&[1.0], &DATA);
let want = -20.935_358_872_705_99;
assert!(
(got - want).abs() <= 1e-10 * want.abs(),
"logL(1.0) was {got}, want {want}"
);
}
#[test]
fn log_likelihood_non_positive_lambda_is_neg_inf() {
let model = PoissonLikelihood::default();
assert!(
model.log_likelihood(&[0.0], &DATA) == f64::NEG_INFINITY,
"lambda = 0 must give -inf"
);
assert!(
model.log_likelihood(&[-1.0], &DATA) == f64::NEG_INFINITY,
"negative lambda must give -inf"
);
}
#[test]
fn log_likelihood_negative_observation_is_neg_inf() {
let model = PoissonLikelihood::default();
assert!(
model.log_likelihood(&[2.5], &[1.0, -2.0, 3.0]) == f64::NEG_INFINITY,
"a negative observation must give -inf"
);
}
#[test]
fn log_likelihood_non_finite_observation_is_neg_inf() {
let model = PoissonLikelihood::default();
assert!(
model.log_likelihood(&[2.5], &[1.0, f64::NAN, 3.0]) == f64::NEG_INFINITY,
"a NaN observation must give -inf"
);
assert!(
model.log_likelihood(&[2.5], &[1.0, f64::INFINITY, 3.0]) == f64::NEG_INFINITY,
"a +inf observation must give -inf"
);
}
#[test]
fn fit_recovers_sample_mean_and_criteria() -> Result<(), Error> {
let model = PoissonLikelihood::default();
let fit = model.fit(&DATA)?;
let lambda_hat = *fit.params().first().unwrap_or(&f64::NAN);
assert!(
(lambda_hat - 2.5).abs() <= 1e-12,
"lambda_hat was {lambda_hat}"
);
let ll = -14.609_544_235_222_89;
assert!(
(fit.log_likelihood() - ll).abs() <= 1e-10 * ll.abs(),
"logL was {}",
fit.log_likelihood()
);
assert!(fit.converged(), "closed-form fit must report converged");
assert_eq!(
fit.iterations(),
0,
"closed-form fit performs no iterations"
);
assert!(
(fit.aic() - 31.219_088_470_445_78).abs() <= 1e-10,
"aic was {}",
fit.aic()
);
assert!(
(fit.bic() - 31.298_530_012_125_614).abs() <= 1e-10,
"bic was {}",
fit.bic()
);
Ok(())
}
#[test]
fn fit_mle_from_perturbed_init_matches_closed_form() -> Result<(), Error> {
let model = PoissonLikelihood::default();
let closed = model.fit(&DATA)?;
let numeric = crate::likelihood::fit_mle(&model, &DATA, &[1.0], 1e-10)?;
let numeric_lambda = *numeric.params().first().unwrap_or(&f64::NAN);
let closed_lambda = *closed.params().first().unwrap_or(&f64::NAN);
assert!(
(numeric_lambda - closed_lambda).abs() <= 1e-5,
"numeric lambda_hat {numeric_lambda} vs closed {closed_lambda}"
);
Ok(())
}
}