use crate::algorithms::count_to_f64;
use crate::error::{Error, Result};
use crate::likelihood::NormalLikelihood;
use crate::likelihood::{LogLikelihood, MleFit};
use std::f64::consts::PI;
impl LogLikelihood for NormalLikelihood {
fn n_params(&self) -> usize {
2
}
fn log_likelihood(&self, params: &[f64], data: &[f64]) -> f64 {
let mu = *params.first().unwrap_or(&f64::NAN);
let sigma = *params.get(1).unwrap_or(&f64::NAN);
if !mu.is_finite() || !sigma.is_finite() || sigma <= 0.0 {
return f64::NEG_INFINITY;
}
if !data.iter().all(|x| x.is_finite()) {
return f64::NEG_INFINITY;
}
let n = count_to_f64(data.len());
let inv_var = 1.0 / (sigma * sigma);
let sse: f64 = data.iter().map(|x| (x - mu) * (x - mu)).sum();
let neg_half_n = -0.5 * n;
let two_pi_term = neg_half_n * (2.0 * PI).ln();
let sigma_term = (-n).mul_add(sigma.ln(), two_pi_term);
(0.5 * sse).mul_add(-inv_var, sigma_term)
}
}
impl NormalLikelihood {
pub fn fit(&self, data: &[f64]) -> Result<MleFit> {
if data.len() < 2 {
return Err(Error::InsufficientData);
}
let n = count_to_f64(data.len());
let mu_hat = data.iter().sum::<f64>() / n;
let sse: f64 = data.iter().map(|x| (x - mu_hat) * (x - mu_hat)).sum();
let sigma_hat = (sse / n).sqrt();
if sigma_hat <= 0.0 {
return Err(Error::DegenerateInput(
"all observations are identical (zero variance)".to_owned(),
));
}
let params = vec![mu_hat, sigma_hat];
let log_likelihood = self.log_likelihood(¶ms, data);
Ok(MleFit::from_closed_form(params, log_likelihood, data.len()))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::likelihood::fit_mle;
const DATA: [f64; 8] = [2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0];
fn param(fit: &MleFit, i: usize) -> Result<f64> {
fit.params()
.get(i)
.copied()
.ok_or_else(|| Error::InvalidInput(format!("missing parameter {i}")))
}
fn is_neg_inf(x: f64) -> bool {
x.is_infinite() && x.is_sign_negative()
}
#[test]
fn fit_rejects_bad_input() {
let model = NormalLikelihood::default();
assert!(
matches!(model.fit(&[]), Err(Error::InsufficientData)),
"empty data should be InsufficientData"
);
assert!(
matches!(model.fit(&[3.0]), Err(Error::InsufficientData)),
"single point should be InsufficientData"
);
assert!(
matches!(model.fit(&[3.0, 3.0, 3.0]), Err(Error::DegenerateInput(_))),
"zero variance should be DegenerateInput"
);
}
#[test]
fn log_likelihood_matches_scipy() {
let model = NormalLikelihood::default();
let got = model.log_likelihood(&[4.5, 2.3], &DATA);
let want = -17.228_391_835_129_557;
assert!(
((got - want) / want).abs() < 1e-10,
"log_likelihood was {got}, want {want}"
);
}
#[test]
fn log_likelihood_out_of_domain_is_neg_inf() {
let model = NormalLikelihood::default();
assert!(
is_neg_inf(model.log_likelihood(&[1.0, 0.0], &DATA)),
"sigma = 0 should be NEG_INFINITY"
);
assert!(
is_neg_inf(model.log_likelihood(&[1.0, -1.0], &DATA)),
"negative sigma should be NEG_INFINITY"
);
assert!(
is_neg_inf(model.log_likelihood(&[f64::NAN, 2.0], &DATA)),
"non-finite mu should be NEG_INFINITY"
);
assert!(
is_neg_inf(model.log_likelihood(&[1.0, f64::INFINITY], &DATA)),
"non-finite sigma should be NEG_INFINITY"
);
}
#[test]
fn log_likelihood_non_finite_observation_is_neg_inf() {
let model = NormalLikelihood::default();
assert!(
is_neg_inf(model.log_likelihood(&[5.0, 2.0], &[1.0, f64::NAN, 3.0])),
"NaN observation should give NEG_INFINITY, got {}",
model.log_likelihood(&[5.0, 2.0], &[1.0, f64::NAN, 3.0])
);
assert!(
is_neg_inf(model.log_likelihood(&[5.0, 2.0], &[1.0, f64::INFINITY, 3.0])),
"+inf observation should give NEG_INFINITY, got {}",
model.log_likelihood(&[5.0, 2.0], &[1.0, f64::INFINITY, 3.0])
);
}
#[test]
fn fit_recovers_closed_form_and_information_criteria() -> Result<()> {
let model = NormalLikelihood::default();
let fit = model.fit(&DATA)?;
let mu_hat = param(&fit, 0)?;
let sigma_hat = param(&fit, 1)?;
assert!(
(mu_hat - 5.0).abs() < 1e-12,
"mu_hat was {mu_hat}, want 5.0"
);
assert!(
(sigma_hat - 2.0).abs() < 1e-12,
"sigma_hat was {sigma_hat}, want 2.0"
);
let want_ll = -16.896_685_710_116_945;
let ll = fit.log_likelihood();
assert!(
((ll - want_ll) / want_ll).abs() < 1e-10,
"log_likelihood was {ll}, want {want_ll}"
);
assert!(fit.converged(), "closed-form fit must report converged");
assert_eq!(fit.iterations(), 0, "closed-form fit does no iterations");
let n = count_to_f64(DATA.len());
let akaike = 2.0f64.mul_add(2.0, -2.0 * ll);
let bayesian = 2.0f64.mul_add(n.ln(), -2.0 * ll);
assert!(
(fit.aic() - akaike).abs() < 1e-12,
"aic was {}, want {akaike}",
fit.aic()
);
assert!(
(fit.bic() - bayesian).abs() < 1e-12,
"bic was {}, want {bayesian}",
fit.bic()
);
Ok(())
}
#[test]
fn fit_mle_from_perturbed_init_recovers_closed_form() -> Result<()> {
let model = NormalLikelihood::default();
let fit = fit_mle(&model, &DATA, &[3.5, 3.0], 1e-10)?;
let mu_hat = param(&fit, 0)?;
let sigma_hat = param(&fit, 1)?;
assert!(
(mu_hat - 5.0).abs() <= 1e-5,
"mu_hat was {mu_hat}, want 5.0"
);
assert!(
(sigma_hat - 2.0).abs() <= 1e-5,
"sigma_hat was {sigma_hat}, want 2.0"
);
Ok(())
}
}