use crate::algorithms::count_to_f64;
use crate::error::{Error, Result};
use crate::likelihood::{LogLikelihood, MleFit};
pub const SIMPLEX_TOLERANCE: f64 = 1e-9;
const INDEX_TOLERANCE: f64 = 1e-9;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CategoricalLikelihoodModel {
pub n_categories: usize,
}
impl LogLikelihood for CategoricalLikelihoodModel {
fn n_params(&self) -> usize {
self.n_categories
}
fn log_likelihood(&self, params: &[f64], data: &[f64]) -> f64 {
log_likelihood_core(self.n_categories, params, data)
}
}
fn is_valid_index(x: f64, k: usize) -> bool {
x.is_finite() && x >= 0.0 && (x - x.round()).abs() < INDEX_TOLERANCE && x < count_to_f64(k)
}
fn category_count(data: &[f64], j: usize) -> f64 {
let target = count_to_f64(j);
let count = data.iter().filter(|&&x| (x - target).abs() < 0.5).count();
count_to_f64(count)
}
fn log_likelihood_core(n_categories: usize, params: &[f64], data: &[f64]) -> f64 {
if params.len() != n_categories {
return f64::NEG_INFINITY;
}
let sum: f64 = params.iter().sum();
if (sum - 1.0).abs() > SIMPLEX_TOLERANCE {
return f64::NEG_INFINITY;
}
if !data.iter().all(|&x| is_valid_index(x, n_categories)) {
return f64::NEG_INFINITY;
}
let mut total = 0.0;
for (j, &p) in params.iter().enumerate() {
let count = category_count(data, j);
if count > 0.0 {
if p <= 0.0 {
return f64::NEG_INFINITY;
}
total = count.mul_add(p.ln(), total);
}
}
total
}
impl crate::likelihood::CategoricalLikelihood {
#[must_use]
pub fn log_likelihood(&self, n_categories: usize, params: &[f64], data: &[f64]) -> f64 {
log_likelihood_core(n_categories, params, data)
}
pub fn fit(&self, n_categories: usize, data: &[f64]) -> Result<MleFit> {
if data.is_empty() {
return Err(Error::InsufficientData);
}
if let Some(&bad) = data.iter().find(|&&x| !is_valid_index(x, n_categories)) {
return Err(Error::InvalidInput(format!(
"datum {bad} is not a non-negative integer below {n_categories} categories"
)));
}
let n = count_to_f64(data.len());
let params: Vec<f64> = (0..n_categories)
.map(|j| category_count(data, j) / n)
.collect();
let log_likelihood = log_likelihood_core(n_categories, ¶ms, data);
Ok(MleFit::from_closed_form(params, log_likelihood, data.len()))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CategoricalLogOdds {
pub n_categories: usize,
}
impl CategoricalLogOdds {
fn log_probabilities(self, params: &[f64]) -> Option<Vec<f64>> {
let k = self.n_categories;
if k == 0 || params.len() + 1 != k {
return None;
}
if !params.iter().all(|z| z.is_finite()) {
return None;
}
let mut logits = Vec::with_capacity(k);
logits.push(0.0_f64);
logits.extend_from_slice(params);
let m = logits.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let sum_exp: f64 = logits.iter().map(|z| (z - m).exp()).sum();
let lse = m + sum_exp.ln();
Some(logits.iter().map(|z| z - lse).collect())
}
pub fn probabilities(&self, params: &[f64]) -> Result<Vec<f64>> {
self.log_probabilities(params)
.map(|log_p| log_p.iter().map(|l| l.exp()).collect())
.ok_or_else(|| {
Error::InvalidInput(format!(
"params must be {} finite logits for {} categories",
self.n_categories.saturating_sub(1),
self.n_categories
))
})
}
pub fn from_probabilities(p: &[f64]) -> Result<Vec<f64>> {
let Some((&p0, rest)) = p.split_first() else {
return Err(Error::InvalidInput(
"probability vector must be non-empty".to_owned(),
));
};
if !p.iter().all(|&pj| pj.is_finite() && pj > 0.0) {
return Err(Error::InvalidInput(format!(
"probabilities must be finite and strictly positive for log-odds, p0 was {p0}"
)));
}
let sum: f64 = p.iter().sum();
if (sum - 1.0).abs() > SIMPLEX_TOLERANCE {
return Err(Error::InvalidInput(format!(
"probabilities must sum to 1 within tolerance, sum was {sum}"
)));
}
let ln_p0 = p0.ln();
Ok(rest.iter().map(|&pj| pj.ln() - ln_p0).collect())
}
}
impl LogLikelihood for CategoricalLogOdds {
fn n_params(&self) -> usize {
self.n_categories.saturating_sub(1)
}
fn log_likelihood(&self, params: &[f64], data: &[f64]) -> f64 {
let Some(log_p) = self.log_probabilities(params) else {
return f64::NEG_INFINITY;
};
let k = self.n_categories;
if !data.iter().all(|&x| is_valid_index(x, k)) {
return f64::NEG_INFINITY;
}
let mut total = 0.0;
for (j, &log_pj) in log_p.iter().enumerate() {
let count = category_count(data, j);
if count > 0.0 {
total = count.mul_add(log_pj, total);
}
}
total
}
}
#[cfg(test)]
mod tests;