1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
use ndarray::{Array1, Array2};
use crate::{Error, Result};
/// Fitted Gaussian Naive Bayes inference parameters.
#[derive(Clone, Debug, PartialEq)]
pub struct GaussianNaiveBayes {
/// Per-class, per-feature means shaped `[classes, features]`.
pub means: Array2<f64>,
/// Per-class, per-feature variances shaped `[classes, features]`.
pub variances: Array2<f64>,
/// Prior probability for each class.
pub priors: Array1<f64>,
/// Integer label corresponding to each class row.
pub class_labels: Vec<i64>,
}
impl GaussianNaiveBayes {
/// Creates validated Gaussian Naive Bayes parameters.
///
/// # Errors
///
/// Returns an error when shapes disagree, variances/priors are not
/// positive, labels are duplicated, or a numeric value is non-finite.
pub fn new(
means: Array2<f64>,
variances: Array2<f64>,
priors: Array1<f64>,
class_labels: Vec<i64>,
) -> Result<Self> {
let classes = means.nrows();
let mut unique_labels = class_labels.clone();
unique_labels.sort_unstable();
unique_labels.dedup();
if classes == 0
|| means.ncols() == 0
|| variances.dim() != means.dim()
|| priors.len() != classes
|| class_labels.len() != classes
|| unique_labels.len() != classes
|| means.iter().any(|value| !value.is_finite())
|| variances
.iter()
.any(|value| !value.is_finite() || *value <= 0.0)
|| priors
.iter()
.any(|value| !value.is_finite() || *value <= 0.0)
{
return Err(Error::InvalidModel(
"invalid Gaussian Naive Bayes parameters".into(),
));
}
Ok(Self {
means,
variances,
priors,
class_labels,
})
}
/// Required number of input features.
#[must_use]
pub fn n_features(&self) -> usize {
self.means.ncols()
}
}