use itertools::Itertools;
use rand::{Rng, seq::index::sample};
use crate::datasets::{CatTrj, CatTrjEv, CatTrjEvT, CatTrjs, CatTrjsEv, Dataset};
pub struct RngEv<'a, R, D> {
rng: &'a mut R,
dataset: &'a D,
p: f64,
}
impl<'a, R, D> RngEv<'a, R, D> {
pub fn new(rng: &'a mut R, dataset: &'a D, p: f64) -> Self {
assert!((0.0..=1.0).contains(&p), "Probability must be in [0, 1]");
Self { rng, dataset, p }
}
}
impl<R: Rng> RngEv<'_, R, CatTrj> {
pub fn random(&mut self) -> CatTrjEv {
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> RngEv<'_, R, CatTrjs> {
pub fn random(&mut self) -> CatTrjsEv {
self.dataset
.values()
.iter()
.map(|trj| RngEv::new(&mut self.rng, trj, self.p).random())
.collect()
}
}