use dry::macro_for;
use ndarray::prelude::*;
use crate::{
datasets::{CatTrj, CatTrjs, CatWtdTrj, CatWtdTrjs},
estimators::{CPDEstimator, CSSEstimator, MLE, ParCPDEstimator, ParCSSEstimator, SSE},
models::{CatCIM, CatCIMS},
types::{Error, Result, Set, States},
};
impl MLE<'_, CatTrj> {
fn fit(
states: &States,
x: &Set<usize>,
z: &Set<usize>,
fitted_statistics: CatCIMS,
) -> Result<CatCIM> {
let n_xz = fitted_statistics.fitted_conditional_counts();
let t_xz = fitted_statistics.fitted_conditional_times();
if !t_xz.iter().all(|&x| x > 0.) {
return Err(Error::Stats("Failed to get non-zero conditional times."));
}
let t_xz = &t_xz.clone().insert_axis(Axis(2));
let mut parameters = n_xz / t_xz;
parameters.outer_iter_mut().for_each(|mut q| {
q.diag_mut().fill(0.);
let q_neg_sum = -q.sum_axis(Axis(1));
q.diag_mut().assign(&q_neg_sum);
});
let eps = f64::MIN_POSITIVE;
let fitted_log_likelihood = {
let ll_q_xz = {
let n_z = n_xz.sum_axis(Axis(2));
let t_z = t_xz.sum_axis(Axis(2));
let mut q_z = Array::zeros(n_z.dim());
parameters
.outer_iter()
.zip(q_z.outer_iter_mut())
.for_each(|(p, mut q)| {
q.assign(&(-&p.diag()));
});
(&n_z * (&q_z + eps).ln()).sum() + (-&q_z * &t_z).sum()
};
let ll_p_xz = {
let mut p_xz = parameters.clone();
p_xz.outer_iter_mut().for_each(|mut p| {
p.diag_mut().fill(0.);
});
p_xz /= &p_xz.sum_axis(Axis(2)).insert_axis(Axis(2));
(n_xz * (p_xz + eps).ln()).sum()
};
ll_q_xz + ll_p_xz
};
let conditioning_states = z
.iter()
.map(|&i| {
let (k, v) = states
.get_index(i)
.ok_or_else(|| Error::IndexOutOfBounds(i))?;
Ok((k.clone(), v.clone()))
})
.collect::<Result<_>>()?;
let states = x
.iter()
.map(|&i| {
let (k, v) = states
.get_index(i)
.ok_or_else(|| Error::IndexOutOfBounds(i))?;
Ok((k.clone(), v.clone()))
})
.collect::<Result<_>>()?;
let fitted_statistics = Some(fitted_statistics);
let fitted_log_likelihood = Some(fitted_log_likelihood);
CatCIM::with_optionals(
states,
conditioning_states,
parameters,
fitted_statistics,
fitted_log_likelihood,
)
}
}
macro_for!($type in [CatTrj, CatWtdTrj, CatTrjs, CatWtdTrjs] {
impl CPDEstimator<CatCIM> for MLE<'_, $type> {
fn fit(&self, x: &Set<usize>, z: &Set<usize>) -> Result<CatCIM> {
let states = self.dataset.states();
let fitted_statistics = SSE::new(self.dataset);
let fitted_statistics = fitted_statistics.with_missing_method(
self.missing_method,
self.missing_mechanism.clone()
)?;
let fitted_statistics = fitted_statistics.fit(x, z)?;
MLE::<'_, CatTrj>::fit(states, x, z, fitted_statistics)
}
}
});
macro_for!($type in [CatTrjs, CatWtdTrjs] {
impl ParCPDEstimator<CatCIM> for MLE<'_, $type> {
fn par_fit(&self, x: &Set<usize>, z: &Set<usize>) -> Result<CatCIM> {
let states = self.dataset.states();
let fitted_statistics = SSE::new(self.dataset);
let fitted_statistics = fitted_statistics.with_missing_method(
self.missing_method,
self.missing_mechanism.clone()
)?;
let fitted_statistics = fitted_statistics.par_fit(x, z)?;
MLE::<'_, CatTrj>::fit(states, x, z, fitted_statistics)
}
}
});