use ndarray::prelude::*;
use ndarray_stats::{Quantile1dExt, interpolate::Nearest};
use noisy_float::prelude::*;
use rand::{Rng, RngExt};
use rand_distr::{Uniform, num_traits::ToPrimitive};
use crate::{
datasets::{Dataset, GaussIncTable, GaussTable, GaussType, IncDataset, MissingMechanism},
models::Labelled,
random::Random,
types::{Error, Result},
};
pub struct RngGaussIncTable<'a, R> {
rng: &'a mut R,
dataset: &'a GaussTable,
missing_mechanism: &'a MissingMechanism,
p_min: f64,
p_max: f64,
}
impl<'a, R: Rng> RngGaussIncTable<'a, R> {
pub fn new(
rng: &'a mut R,
dataset: &'a GaussTable,
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 RngGaussIncTable<'_, R> {
type Output = Result<GaussIncTable>;
fn random(&mut self) -> Self::Output {
const M: GaussType = GaussIncTable::MISSING;
let labels = self.dataset.labels().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 s_z = match c_z
.mapv(n64)
.quantile_mut(n64(0.2), &Nearest)
.map(|q| q.to_f64())
{
Ok(Some(q)) => q,
_ => 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;
}
});
}
}
GaussIncTable::new(labels, values)
}
}