use dry::macro_for;
use ndarray::prelude::*;
use statrs::function::gamma::ln_gamma;
use crate::{
datasets::{CatTrj, CatTrjs, CatWtdTrj, CatWtdTrjs},
estimators::{BE, CIMEstimator, CSSEstimator, ParCIMEstimator, ParCSSEstimator, SSE},
models::{CatCIM, CatCIMS},
types::{Set, States},
};
impl BE<'_, CatTrj, (usize, f64)> {
fn fit(
states: &States,
x: &Set<usize>,
z: &Set<usize>,
sample_statistics: CatCIMS,
prior: (usize, f64),
) -> CatCIM {
let (alpha, tau) = prior;
assert!(alpha > 0, "Alpha must be positive.");
assert!(tau > 0.0, "Tau must be positive.");
let n_xz = sample_statistics.sample_conditional_counts();
let t_xz = sample_statistics.sample_conditional_times();
let t_xz = &t_xz.clone().insert_axis(Axis(2));
let s_z = n_xz.shape()[0] as f64;
let alpha = alpha as f64 / s_z;
let tau = tau / s_z;
let n_xz = n_xz + alpha;
let t_xz = t_xz + tau;
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 sample_log_likelihood = Some({
let n_z = n_xz.sum_axis(Axis(2));
let t_z = t_xz.sum_axis(Axis(2));
let ll_q_xz = {
(&n_z + 1.).mapv(ln_gamma).sum() + (alpha + 1.) * f64::ln(tau) - (ln_gamma(alpha + 1.) + ((&n_z + 1.) * &t_z.ln()).sum())
};
let ll_p_xz = {
(ln_gamma(alpha) - n_z.mapv(ln_gamma).sum()) + (ln_gamma(alpha) - n_xz.mapv(ln_gamma).sum())
};
ll_q_xz + ll_p_xz
});
let conditioning_states = z
.iter()
.map(|&i| {
let (k, v) = states.get_index(i).unwrap();
(k.clone(), v.clone())
})
.collect();
let states = x
.iter()
.map(|&i| {
let (k, v) = states.get_index(i).unwrap();
(k.clone(), v.clone())
})
.collect();
let sample_statistics = Some(sample_statistics);
CatCIM::with_optionals(
states,
conditioning_states,
parameters,
sample_statistics,
sample_log_likelihood,
)
}
}
macro_for!($type in [CatTrj, CatWtdTrj, CatTrjs, CatWtdTrjs] {
impl CIMEstimator<CatCIM> for BE<'_, $type, ()> {
#[inline]
fn fit(&self, x: &Set<usize>, z: &Set<usize>) -> CatCIM {
self.clone().with_prior((1, 1.)).fit(x, z)
}
}
impl CIMEstimator<CatCIM> for BE<'_, $type, (usize, f64)> {
#[inline]
fn fit(&self, x: &Set<usize>, z: &Set<usize>) -> CatCIM {
let (states, prior) = (self.dataset.states(), self.prior);
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);
BE::<'_, CatTrj, _>::fit(states, x, z, sample_statistics, prior)
}
}
});
macro_for!($type in [CatTrjs, CatWtdTrjs] {
impl ParCIMEstimator<CatCIM> for BE<'_, $type, ()> {
#[inline]
fn par_fit(&self, x: &Set<usize>, z: &Set<usize>) -> CatCIM {
self.clone().with_prior((1, 1.)).fit(x, z)
}
}
impl ParCIMEstimator<CatCIM> for BE<'_, $type, (usize, f64)> {
#[inline]
fn par_fit(&self, x: &Set<usize>, z: &Set<usize>) -> CatCIM {
let (states, prior) = (self.dataset.states(), self.prior);
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);
BE::<'_, CatTrj, _>::fit(states, x, z, sample_statistics, prior)
}
}
});