use ndarray::{Array1, Array2};
use crate::{Error, Result};
#[derive(Clone, Debug, PartialEq)]
pub struct CategoricalNaiveBayes {
pub feature_log_probabilities: Vec<Array2<f64>>,
pub log_priors: Array1<f64>,
pub class_labels: Vec<i64>,
}
impl CategoricalNaiveBayes {
pub fn new(
feature_log_probabilities: Vec<Array2<f64>>,
log_priors: Array1<f64>,
class_labels: Vec<i64>,
) -> Result<Self> {
let classes = class_labels.len();
if classes == 0
|| log_priors.len() != classes
|| feature_log_probabilities.is_empty()
|| feature_log_probabilities.iter().any(|table| {
table.nrows() == 0
|| table.ncols() != classes
|| table
.iter()
.any(|value| value.is_nan() || *value == f64::INFINITY)
})
|| log_priors
.iter()
.any(|value| value.is_nan() || *value == f64::INFINITY)
{
return Err(Error::InvalidModel(
"invalid categorical Naive Bayes state".into(),
));
}
Ok(Self {
feature_log_probabilities,
log_priors,
class_labels,
})
}
}