use super::{argmax, classification_result_from, normalize_log, sorted_unique, validate_dims};
use crate::algorithms::classification::ClassificationResult;
use crate::algorithms::count_to_f64;
use crate::error::{Error, Result};
#[derive(Debug, Clone)]
pub struct CategoricalNbModel {
classes: Vec<usize>,
log_priors: Vec<f64>,
log_probs: Vec<Vec<Vec<f64>>>,
cardinalities: Vec<usize>,
n_features: usize,
}
pub fn categorical_nb_fit(x: &[Vec<usize>], y: &[usize], alpha: f64) -> Result<CategoricalNbModel> {
if alpha < 0.0 {
return Err(Error::InvalidInput("alpha must be >= 0".to_owned()));
}
let n_features = validate_dims(x, y)?;
let classes = sorted_unique(y);
if classes.len() < 2 {
return Err(Error::InsufficientData);
}
let n_total = count_to_f64(y.len());
let mut cardinalities = vec![0_usize; n_features];
for row in x {
for (card, &value) in cardinalities.iter_mut().zip(row.iter()) {
*card = (*card).max(value + 1);
}
}
let mut counts: Vec<Vec<Vec<f64>>> = cardinalities
.iter()
.map(|&card| vec![vec![0.0_f64; card]; classes.len()])
.collect();
let mut class_totals = vec![0.0_f64; classes.len()];
for (row, &label) in x.iter().zip(y) {
let class_idx = classes.iter().position(|&c| c == label).unwrap_or(0);
if let Some(total) = class_totals.get_mut(class_idx) {
*total += 1.0;
}
for (feature, &value) in counts.iter_mut().zip(row.iter()) {
if let Some(cell) = feature.get_mut(class_idx).and_then(|c| c.get_mut(value)) {
*cell += 1.0;
}
}
}
let log_priors: Vec<f64> = class_totals
.iter()
.map(|&total| (total / n_total).ln())
.collect();
let log_probs: Vec<Vec<Vec<f64>>> = counts
.iter()
.zip(cardinalities.iter())
.map(|(feature, &card)| {
let card_f = count_to_f64(card);
feature
.iter()
.zip(class_totals.iter())
.map(|(class_counts, &total)| {
let denom = alpha.mul_add(card_f, total);
class_counts
.iter()
.map(|&count| ((count + alpha) / denom).ln())
.collect()
})
.collect()
})
.collect();
Ok(CategoricalNbModel {
classes,
log_priors,
log_probs,
cardinalities,
n_features,
})
}
impl CategoricalNbModel {
#[must_use]
pub fn classes(&self) -> &[usize] {
&self.classes
}
fn joint_log_likelihoods(&self, x: &[Vec<usize>]) -> Result<Vec<Vec<f64>>> {
let mut out = Vec::with_capacity(x.len());
for row in x {
if row.len() != self.n_features {
return Err(Error::InvalidInput(
"sample feature count differs from the fitted model".to_owned(),
));
}
for (&value, &card) in row.iter().zip(self.cardinalities.iter()) {
if value >= card {
return Err(Error::InvalidInput(
"category index exceeds the trained cardinality".to_owned(),
));
}
}
let mut scores = self.log_priors.clone();
for (feature, &value) in self.log_probs.iter().zip(row.iter()) {
for (score, class_probs) in scores.iter_mut().zip(feature.iter()) {
*score += class_probs.get(value).copied().unwrap_or(f64::NEG_INFINITY);
}
}
out.push(scores);
}
Ok(out)
}
pub fn predict(&self, x: &[Vec<usize>]) -> Result<Vec<usize>> {
let joints = self.joint_log_likelihoods(x)?;
Ok(joints
.iter()
.map(|row| self.classes.get(argmax(row)).copied().unwrap_or(0))
.collect())
}
pub fn predict_log_proba(&self, x: &[Vec<usize>]) -> Result<Vec<Vec<f64>>> {
let joints = self.joint_log_likelihoods(x)?;
Ok(joints.iter().map(|row| normalize_log(row)).collect())
}
pub fn classification_result(
&self,
x: &[Vec<usize>],
y_true: &[usize],
) -> Result<ClassificationResult> {
let predictions = self.predict(x)?;
classification_result_from(
&self.classes,
&predictions,
y_true,
"Categorical Naive Bayes",
)
}
}