use ndarray::{Array1, Array2, Array3};
use crate::{Error, Result};
#[derive(Clone, Debug, PartialEq)]
pub struct GaussianMixture {
pub means: Array2<f64>,
pub precisions: Array3<f64>,
pub offsets: Array1<f64>,
}
impl GaussianMixture {
pub fn new(means: Array2<f64>, precisions: Array3<f64>, offsets: Array1<f64>) -> Result<Self> {
let components = means.nrows();
let features = means.ncols();
if components == 0
|| features == 0
|| precisions.dim() != (components, features, features)
|| offsets.len() != components
|| means
.iter()
.chain(precisions.iter())
.chain(offsets.iter())
.any(|value| !value.is_finite())
{
return Err(Error::InvalidModel("invalid Gaussian mixture state".into()));
}
Ok(Self {
means,
precisions,
offsets,
})
}
}