use ndarray::prelude::*;
use ndarray_stats::QuantileExt;
use rand::{Rng, RngExt};
use rand_distr::Uniform;
use crate::{
datasets::{CatIncTable, CatTable, CatType, Dataset, IncDataset, MissingMechanism},
models::Labelled,
random::Random,
types::{Error, Result},
};
pub struct RngCatIncTable<'a, R> {
rng: &'a mut R,
dataset: &'a CatTable,
missing_mechanism: &'a MissingMechanism,
p_min: f64,
p_max: f64,
}
impl<'a, R: Rng> RngCatIncTable<'a, R> {
pub fn new(
rng: &'a mut R,
dataset: &'a CatTable,
missing_mechanism: &'a MissingMechanism,
p_min: f64,
p_max: f64,
) -> Result<Self> {
if dataset.labels() != missing_mechanism.labels() {
return Err(Error::InvalidParameter(
"missing_mechanism",
"labels do not match dataset labels",
));
}
if !(0.0..=1.0).contains(&p_min) {
return Err(Error::InvalidParameter("p_min", "must be in [0, 1]"));
}
if !(0.0..=1.0).contains(&p_max) {
return Err(Error::InvalidParameter("p_max", "must be in [0, 1]"));
}
if p_min > p_max {
return Err(Error::InvalidParameter(
"p_min",
"must be less than or equal to p_max",
));
}
Ok(Self {
rng,
dataset,
missing_mechanism,
p_min,
p_max,
})
}
}
impl<R: Rng> Random for RngCatIncTable<'_, R> {
type Output = Result<CatIncTable>;
fn random(&mut self) -> Self::Output {
const M: CatType = CatIncTable::MISSING;
let states = self.dataset.states().clone();
let mut values = self.dataset.values().clone();
let p_s = Uniform::new_inclusive(0., 1.)?;
let p_u = Uniform::new_inclusive(self.p_min, self.p_max)?;
for (&x, pa_x) in self.missing_mechanism {
let mut c_x = values.column_mut(x);
if pa_x.is_empty() {
let p_x = self.rng.sample(p_u);
c_x.iter_mut().for_each(|x| {
if self.rng.sample(p_s) < p_x {
*x = M;
}
});
continue;
}
for &z in pa_x {
let c_z = self.dataset.values().column(z);
let mut m_z = Array::from_elem(CatType::MAX as usize, 0);
c_z.iter().for_each(|&z| m_z[z as usize] += 1);
let s_z = match m_z.argmax() {
Ok(s_z) => s_z as CatType,
_ => continue,
};
azip!((x in &mut c_x, z in &c_z) {
if
(self.rng.sample(p_s) < self.p_max && *z == s_z) ||
(self.rng.sample(p_s) < self.p_min && *z != s_z)
{
*x = M;
}
});
}
}
CatIncTable::new(states, values)
}
}