use std::marker::PhantomData;
use dry::macro_for;
use crate::{
estimators::CPDEstimator,
models::{CIM, CPD, CatCIM, CatCPD, GaussCPD, HasLabels},
types::{Error, Labels, Result, Set},
};
pub trait HasEstimator {
type Estimator;
fn estimator(&self) -> &Self::Estimator;
}
pub trait ScoringCriterion {
fn call(&self, x: &Set<usize>, z: &Set<usize>) -> Result<f64>;
}
macro_rules! impl_scoring_struct {
($name:ident, $doc:expr) => {
#[doc = $doc]
pub struct $name<'a, E, P> {
_p_marker: PhantomData<P>,
estimator: &'a E,
}
impl<'a, E, P> $name<'a, E, P> {
#[doc = concat!("Creates a new `", stringify!($name), "` instance.")]
#[doc = concat!("A new `", stringify!($name), "` instance.")]
#[inline]
pub const fn new(estimator: &'a E) -> Self {
Self {
_p_marker: PhantomData,
estimator,
}
}
}
impl<E, P> HasEstimator for $name<'_, E, P> {
type Estimator = E;
#[inline]
fn estimator(&self) -> &Self::Estimator {
self.estimator
}
}
impl<'a, E, P> HasLabels for $name<'a, E, P>
where
E: HasLabels,
{
#[inline]
fn labels(&self) -> &Labels {
self.estimator.labels()
}
}
};
}
impl_scoring_struct!(LL, "The Log Likelihood (`LL`).");
impl_scoring_struct!(AIC, "The Akaike Information Criterion (`AIC`).");
impl_scoring_struct!(AICC, "The Akaike Information Criterion Corrected (`AICc`).");
impl_scoring_struct!(BIC, "The Bayesian Information Criterion (`BIC`).");
impl_scoring_struct!(
BICC,
"The Bayesian Information Criterion Corrected (`BICc`)."
);
impl_scoring_struct!(HQC, "The Hannan-Quinn Criterion (`HQC`).");
macro_for!($type in [CatCPD, GaussCPD, CatCIM] {
impl<E> ScoringCriterion for LL<'_, E, $type>
where
E: CPDEstimator<$type>,
{
#[inline]
fn call(&self, x: &Set<usize>, z: &Set<usize>) -> Result<f64> {
let p_xz = self.estimator.fit(x, z)?;
let log_likelihood = p_xz
.fitted_log_likelihood()
.ok_or_else(|| Error::MissingLogLikelihood())?;
Ok(log_likelihood)
}
}
impl<E> ScoringCriterion for AIC<'_, E, $type>
where
E: CPDEstimator<$type>,
{
#[inline]
fn call(&self, x: &Set<usize>, z: &Set<usize>) -> Result<f64> {
let p_xz = self.estimator.fit(x, z)?;
let log_likelihood = p_xz
.fitted_log_likelihood()
.ok_or_else(|| Error::MissingLogLikelihood())?;
let k = p_xz.parameters_size() as f64;
Ok(log_likelihood - k)
}
}
impl<E> ScoringCriterion for AICC<'_, E, $type>
where
E: CPDEstimator<$type>,
{
#[inline]
fn call(&self, x: &Set<usize>, z: &Set<usize>) -> Result<f64> {
let p_xz = self.estimator.fit(x, z)?;
let n = p_xz
.fitted_statistics()
.ok_or_else(|| Error::MissingSufficientStatistics())?
.fitted_size();
let log_likelihood = p_xz
.fitted_log_likelihood()
.ok_or_else(|| Error::MissingLogLikelihood())?;
let k = p_xz.parameters_size() as f64;
let c = n / (f64::max(n - k - 2., 1.));
Ok(log_likelihood - k * c)
}
}
impl<E> ScoringCriterion for BIC<'_, E, $type>
where
E: CPDEstimator<$type>,
{
#[inline]
fn call(&self, x: &Set<usize>, z: &Set<usize>) -> Result<f64> {
let p_xz = self.estimator.fit(x, z)?;
let n = p_xz
.fitted_statistics()
.ok_or_else(|| Error::MissingSufficientStatistics())?
.fitted_size();
let log_likelihood = p_xz
.fitted_log_likelihood()
.ok_or_else(|| Error::MissingLogLikelihood())?;
let k = p_xz.parameters_size() as f64;
Ok(log_likelihood - 0.5 * k * f64::ln(n))
}
}
impl<E> ScoringCriterion for BICC<'_, E, $type>
where
E: CPDEstimator<$type>,
{
#[inline]
fn call(&self, x: &Set<usize>, z: &Set<usize>) -> Result<f64> {
let p_xz = self.estimator.fit(x, z)?;
let n = p_xz
.fitted_statistics()
.ok_or_else(|| Error::MissingSufficientStatistics())?
.fitted_size();
let log_likelihood = p_xz
.fitted_log_likelihood()
.ok_or_else(|| Error::MissingLogLikelihood())?;
let k = p_xz.parameters_size() as f64;
let c = n / (f64::max(n - k - 2., 1.));
Ok(log_likelihood - 0.5 * k * c * f64::ln(n))
}
}
impl<E> ScoringCriterion for HQC<'_, E, $type>
where
E: CPDEstimator<$type>,
{
#[inline]
fn call(&self, x: &Set<usize>, z: &Set<usize>) -> Result<f64> {
let p_xz = self.estimator.fit(x, z)?;
let n = p_xz
.fitted_statistics()
.ok_or_else(|| Error::MissingSufficientStatistics())?
.fitted_size();
let log_likelihood = p_xz
.fitted_log_likelihood()
.ok_or_else(|| Error::MissingLogLikelihood())?;
let k = p_xz.parameters_size() as f64;
Ok(log_likelihood - 0.5 * k * f64::ln(f64::ln(n)))
}
}
});