use super::{CategoricalLikelihoodModel, CategoricalLogOdds, log_likelihood_core};
use crate::error::{Error, Result};
use crate::likelihood::CategoricalLikelihood;
use crate::likelihood::{LogLikelihood, MleFit, fit_mle};
fn param(fit: &MleFit, j: usize) -> Result<f64> {
fit.params()
.get(j)
.copied()
.ok_or_else(|| Error::InvalidInput(format!("missing parameter {j}")))
}
#[test]
fn fit_empty_is_insufficient_data() {
assert!(
matches!(
CategoricalLikelihood::default().fit(3, &[]),
Err(Error::InsufficientData)
),
"empty data should be InsufficientData"
);
}
#[test]
fn fit_index_out_of_range_is_invalid_input() {
assert!(
matches!(
CategoricalLikelihood::default().fit(3, &[0.0, 3.0]),
Err(Error::InvalidInput(_))
),
"out-of-range index should be InvalidInput"
);
}
#[test]
fn fit_non_integer_is_invalid_input() {
assert!(
matches!(
CategoricalLikelihood::default().fit(3, &[0.0, 1.5]),
Err(Error::InvalidInput(_))
),
"non-integer datum should be InvalidInput"
);
}
#[test]
fn trait_log_likelihood_matches_golden() {
let model = CategoricalLikelihoodModel { n_categories: 3 };
let data = [0.0, 0.0, 1.0, 2.0, 2.0, 2.0];
let ll = model.log_likelihood(&[0.2, 0.3, 0.5], &data);
let expected = -6.502_290_170_873_972;
assert!(
((ll - expected) / expected).abs() < 1e-10,
"ll was {ll}, expected {expected}"
);
}
#[test]
fn trait_zero_prob_on_nonempty_category_is_neg_inf() {
let model = CategoricalLikelihoodModel { n_categories: 3 };
let ll = model.log_likelihood(&[0.0, 0.4, 0.6], &[0.0, 1.0, 2.0]);
assert!(ll.is_infinite() && ll < 0.0, "ll was {ll}");
}
#[test]
fn trait_wrong_params_length_is_neg_inf() {
let model = CategoricalLikelihoodModel { n_categories: 3 };
let ll = model.log_likelihood(&[0.5, 0.5], &[0.0, 1.0, 2.0]);
assert!(ll.is_infinite() && ll < 0.0, "ll was {ll}");
}
#[test]
fn trait_unnormalized_params_is_neg_inf() {
let model = CategoricalLikelihoodModel { n_categories: 3 };
let ll = model.log_likelihood(&[0.3, 0.3, 0.3], &[0.0, 1.0, 2.0]);
assert!(ll.is_infinite() && ll < 0.0, "ll was {ll}");
}
#[test]
fn trait_zero_prob_on_empty_category_is_finite() {
let model = CategoricalLikelihoodModel { n_categories: 3 };
let ll = model.log_likelihood(&[0.5, 0.5, 0.0], &[0.0, 0.0, 1.0, 1.0]);
let expected = -2.772_588_722_239_781;
assert!(ll.is_finite(), "ll was {ll}");
assert!(
((ll - expected) / expected).abs() < 1e-10,
"ll was {ll}, expected {expected}"
);
}
#[test]
fn trait_non_finite_observation_is_neg_inf() {
let model = CategoricalLikelihoodModel { n_categories: 3 };
let ll_nan = model.log_likelihood(&[0.2, 0.3, 0.5], &[0.0, f64::NAN, 2.0]);
assert!(ll_nan.is_infinite() && ll_nan < 0.0, "NaN index: {ll_nan}");
let ll_inf = model.log_likelihood(&[0.2, 0.3, 0.5], &[0.0, f64::INFINITY, 2.0]);
assert!(ll_inf.is_infinite() && ll_inf < 0.0, "+inf index: {ll_inf}");
}
#[test]
fn core_and_trait_agree() {
let model = CategoricalLikelihoodModel { n_categories: 3 };
let params = [0.2, 0.3, 0.5];
let data = [0.0, 1.0, 2.0, 2.0];
let via_trait = model.log_likelihood(¶ms, &data);
let via_core = log_likelihood_core(3, ¶ms, &data);
assert!(
(via_trait - via_core).abs() < 1e-15,
"core disagreed with trait"
);
}
#[test]
fn fit_recovers_empirical_frequencies() -> Result<()> {
let model = CategoricalLikelihood::default();
let fit = model.fit(3, &[0.0, 0.0, 1.0, 2.0, 2.0, 2.0])?;
let expected = [1.0 / 3.0, 1.0 / 6.0, 0.5];
for (j, &want) in expected.iter().enumerate() {
let got = param(&fit, j)?;
assert!((got - want).abs() < 1e-12, "p[{j}] was {got}, want {want}");
}
assert!(fit.converged(), "closed-form fit should report converged");
assert_eq!(fit.iterations(), 0, "closed form does no iterations");
let ll = fit.log_likelihood();
let ll_want = -6.068_425_588_244_111;
assert!(((ll - ll_want) / ll_want).abs() < 1e-12, "ll was {ll}");
assert!(
(fit.aic() - 18.136_851_176_488_22).abs() < 1e-10,
"aic was {}",
fit.aic()
);
assert!(
(fit.bic() - 17.512_129_584_172_385).abs() < 1e-10,
"bic was {}",
fit.bic()
);
Ok(())
}
#[test]
fn fit_allows_zero_count_category() -> Result<()> {
let model = CategoricalLikelihood::default();
let fit = model.fit(3, &[0.0, 0.0, 1.0, 1.0])?;
let expected = [0.5, 0.5, 0.0];
for (j, &want) in expected.iter().enumerate() {
let got = param(&fit, j)?;
assert!((got - want).abs() < 1e-12, "p[{j}] was {got}, want {want}");
}
let ll = fit.log_likelihood();
let ll_want = -2.772_588_722_239_781;
assert!(ll.is_finite(), "fitted ll was {ll}");
assert!(((ll - ll_want) / ll_want).abs() < 1e-12, "ll was {ll}");
Ok(())
}
#[test]
fn from_probabilities_rejects_bad_input() {
assert!(
matches!(
CategoricalLogOdds::from_probabilities(&[]),
Err(Error::InvalidInput(_))
),
"empty p should be InvalidInput"
);
assert!(
matches!(
CategoricalLogOdds::from_probabilities(&[0.0, 0.5, 0.5]),
Err(Error::InvalidInput(_))
),
"p0 = 0 should be InvalidInput"
);
assert!(
matches!(
CategoricalLogOdds::from_probabilities(&[0.5, 0.0, 0.5]),
Err(Error::InvalidInput(_))
),
"an interior zero should be InvalidInput"
);
assert!(
matches!(
CategoricalLogOdds::from_probabilities(&[0.3, 0.3, 0.3]),
Err(Error::InvalidInput(_))
),
"unnormalized p should be InvalidInput"
);
}
#[test]
fn probabilities_from_probabilities_round_trip() -> Result<()> {
let p = [0.2, 0.3, 0.5];
let z = CategoricalLogOdds::from_probabilities(&p)?;
let model = CategoricalLogOdds { n_categories: 3 };
let recovered = model.probabilities(&z)?;
for (j, &want) in p.iter().enumerate() {
let got = *recovered
.get(j)
.ok_or_else(|| Error::InvalidInput(format!("missing p{j}")))?;
assert!((got - want).abs() < 1e-12, "p[{j}] was {got}, want {want}");
}
Ok(())
}
#[test]
fn logodds_invalid_params_are_neg_inf() {
let model = CategoricalLogOdds { n_categories: 3 };
let data = [0.0, 1.0, 2.0];
let ll_nan = model.log_likelihood(&[f64::NAN, 0.0], &data);
assert!(ll_nan.is_infinite() && ll_nan < 0.0, "NaN logit: {ll_nan}");
let ll_len = model.log_likelihood(&[0.0], &data);
assert!(
ll_len.is_infinite() && ll_len < 0.0,
"wrong length: {ll_len}"
);
}
#[test]
fn logodds_fit_mle_recovers_empirical_frequencies() -> Result<()> {
let model = CategoricalLogOdds { n_categories: 3 };
let data = [0.0, 0.0, 1.0, 2.0, 2.0, 2.0];
let fit = fit_mle(&model, &data, &[0.0, 0.0], 1e-8)?;
assert!(
fit.converged(),
"fit_mle should converge, iters {}",
fit.iterations()
);
let p = model.probabilities(fit.params())?;
let expected = [1.0 / 3.0, 1.0 / 6.0, 0.5];
for (j, &want) in expected.iter().enumerate() {
let got = *p
.get(j)
.ok_or_else(|| Error::InvalidInput(format!("missing p{j}")))?;
assert!((got - want).abs() <= 1e-4, "p[{j}] was {got}, want {want}");
}
Ok(())
}
#[test]
fn logodds_log_likelihood_matches_categorical_model() -> Result<()> {
let p = [0.2, 0.3, 0.5];
let data = [0.0, 0.0, 1.0, 2.0, 2.0, 2.0];
let z = CategoricalLogOdds::from_probabilities(&p)?;
let logodds = CategoricalLogOdds { n_categories: 3 };
let model = CategoricalLikelihoodModel { n_categories: 3 };
let ll_logodds = logodds.log_likelihood(&z, &data);
let ll_model = model.log_likelihood(&p, &data);
assert!(
(ll_logodds - ll_model).abs() < 1e-12,
"logodds ll {ll_logodds} vs simplex model ll {ll_model}"
);
Ok(())
}