use itertools::Itertools;
use rand::{Rng, RngExt, seq::index::sample};
use crate::{
datasets::{CatTrj, CatTrjEv, CatTrjEvT, CatTrjs, CatTrjsEv, Dataset},
random::Random,
types::{Error, Result},
};
pub struct RngCatTrjEv<'a, R, D> {
rng: &'a mut R,
dataset: &'a D,
p: f64,
}
impl<'a, R, D> RngCatTrjEv<'a, R, D> {
pub fn new(rng: &'a mut R, dataset: &'a D, p: f64) -> Result<Self> {
if !(0.0..=1.0).contains(&p) {
return Err(Error::InvalidParameter("p", "must be in [0, 1]"));
}
Ok(Self { rng, dataset, p })
}
}
impl<R: Rng> Random for RngCatTrjEv<'_, R, CatTrj> {
type Output = Result<CatTrjEv>;
fn random(&mut self) -> Self::Output {
use CatTrjEvT as E;
let times = self.dataset.times();
let events = self.dataset.values().rows();
let times_events = times.into_iter().zip(events);
let evidence = times_events
.tuple_windows()
.filter_map(|((&start_time, v), (&end_time, _))| {
if !self.rng.random_bool(self.p) {
return None;
}
let n = self.rng.random_range(1..=v.len());
let evidence = sample(self.rng, v.len(), n).into_iter().map(move |index| {
let (event, state) = (index, v[index] as usize);
E::CertainPositiveInterval {
event,
state,
start_time,
end_time,
}
});
Some(evidence)
})
.flatten();
CatTrjEv::new(self.dataset.states().clone(), evidence)
}
}
impl<R: Rng> Random for RngCatTrjEv<'_, R, CatTrjs> {
type Output = Result<CatTrjsEv>;
fn random(&mut self) -> Self::Output {
let evidences = self
.dataset
.values()
.iter()
.map(|trj| RngCatTrjEv::<_, CatTrj>::new(self.rng, trj, self.p)?.random())
.collect::<Result<Vec<_>>>()?;
CatTrjsEv::new(evidences)
}
}