use dry::macro_for;
use ndarray::prelude::*;
use ndarray_linalg::Determinant;
use crate::{
datasets::{GaussIncTable, GaussTable, GaussWtdTable},
estimators::{CPDEstimator, CSSEstimator, MLE, ParCPDEstimator, ParCSSEstimator, SSE},
models::{GaussCPD, GaussCPDP, GaussCPDS, Labelled},
types::{EPSILON, Error, LN_2_PI, Labels, Result, Set},
utils::PseudoInverse,
};
impl MLE<'_, GaussTable> {
fn fit(
labels: &Labels,
x: &Set<usize>,
z: &Set<usize>,
sample_statistics: GaussCPDS,
) -> Result<GaussCPD> {
let (mu_x, mu_z, s_xx, s_xz, s_zz, n) = (
sample_statistics.fitted_response_mean(),
sample_statistics.fitted_design_mean(),
sample_statistics.fitted_response_covariance(),
sample_statistics.fitted_cross_covariance(),
sample_statistics.fitted_design_covariance(),
sample_statistics.fitted_size(),
);
let (a, b, s) = if z.is_empty() {
let a = Array2::zeros((x.len(), 0));
let b = mu_x.clone();
let s = s_xx / n;
(a, b, s)
} else {
let s_zz_pinv = s_zz.pinv()?;
let a = s_xz.dot(&s_zz_pinv);
let b = mu_x - &a.dot(mu_z);
let s = (s_xx - &a.dot(&s_xz.t())) / n;
(a, b, s)
};
let mut s = s;
s.diag_mut().mapv_inplace(|x| x.max(0.));
let t = s.diag().sum() / s.nrows() as f64;
*s.diag_mut() += f64::max(EPSILON, t * EPSILON);
s = (&s + &s.t()) / 2.;
let p = x.len() as f64;
let (_, ln_det) = s
.sln_det()
.map_err(|e| Error::Linalg(&format!("Failed to compute determinant of S: {e}")))?;
let sample_log_likelihood = -0.5 * n * (p * LN_2_PI + ln_det + p);
let parameters = GaussCPDP::new(a, b, s)?;
let conditioning_labels = z.iter().map(|&i| labels[i].clone()).collect();
let labels = x.iter().map(|&i| labels[i].clone()).collect();
let sample_statistics = Some(sample_statistics);
let sample_log_likelihood = Some(sample_log_likelihood);
GaussCPD::with_optionals(
labels,
conditioning_labels,
parameters,
sample_statistics,
sample_log_likelihood,
)
}
}
macro_for!($type in [GaussTable, GaussIncTable, GaussWtdTable] {
impl CPDEstimator<GaussCPD> for MLE<'_, $type> {
fn fit(&self, x: &Set<usize>, z: &Set<usize>) -> Result<GaussCPD> {
let labels = self.dataset.labels();
let sample_statistics = SSE::new(self.dataset);
let sample_statistics = sample_statistics.with_missing_method(
self.missing_method,
self.missing_mechanism.clone()
)?;
let sample_statistics = sample_statistics.fit(x, z)?;
MLE::<'_, GaussTable>::fit(labels, x, z, sample_statistics)
}
}
impl ParCPDEstimator<GaussCPD> for MLE<'_, $type> {
fn par_fit(&self, x: &Set<usize>, z: &Set<usize>) -> Result<GaussCPD> {
let labels = self.dataset.labels();
let sample_statistics = SSE::new(self.dataset);
let sample_statistics = sample_statistics.with_missing_method(
self.missing_method,
self.missing_mechanism.clone()
)?;
let sample_statistics = sample_statistics.par_fit(x, z)?;
MLE::<'_, GaussTable>::fit(labels, x, z, sample_statistics)
}
}
});