causal_hub/random/datasets/trajectory/categorical/
evidence.rs1use itertools::Itertools;
2use rand::{Rng, RngExt, seq::index::sample};
3
4use crate::{
5 datasets::{CatTrj, CatTrjEv, CatTrjEvT, CatTrjs, CatTrjsEv, Dataset},
6 random::Random,
7 types::{Error, Result},
8};
9
10pub struct RngCatTrjEv<'a, R, D> {
12 rng: &'a mut R,
13 dataset: &'a D,
14 probability: f64,
15}
16
17impl<'a, R, D> RngCatTrjEv<'a, R, D> {
18 pub fn new(rng: &'a mut R, dataset: &'a D, probability: f64) -> Result<Self> {
30 if !(0.0..=1.0).contains(&probability) {
32 return Err(Error::InvalidParameter("p", "must be in [0, 1]"));
33 }
34
35 Ok(Self {
36 rng,
37 dataset,
38 probability,
39 })
40 }
41}
42
43impl<R: Rng> Random for RngCatTrjEv<'_, R, CatTrj> {
44 type Output = Result<CatTrjEv>;
45
46 fn random(&mut self) -> Self::Output {
47 use CatTrjEvT as E;
49
50 let times = self.dataset.times();
52 let events = self.dataset.values().rows();
54 let times_events = times.into_iter().zip(events);
56
57 let evidence = times_events
59 .tuple_windows()
60 .filter_map(|((&start_time, v), (&end_time, _))| {
61 if !self.rng.random_bool(self.probability) {
63 return None;
65 }
66 let n = self.rng.random_range(1..=v.len());
68 let evidence = sample(self.rng, v.len(), n).into_iter().map(move |index| {
70 let (event, state) = (index, v[index] as usize);
72 E::CertainPositiveInterval {
74 event,
75 state,
76 start_time,
77 end_time,
78 }
79 });
80 Some(evidence)
82 })
83 .flatten();
84
85 CatTrjEv::new(self.dataset.support().clone(), evidence)
87 }
88}
89
90impl<R: Rng> Random for RngCatTrjEv<'_, R, CatTrjs> {
91 type Output = Result<CatTrjsEv>;
92
93 fn random(&mut self) -> Self::Output {
94 let evidences = self
95 .dataset
96 .values()
97 .iter()
98 .map(|trj| RngCatTrjEv::<_, CatTrj>::new(self.rng, trj, self.probability)?.random())
99 .collect::<Result<Vec<_>>>()?;
100
101 CatTrjsEv::new(evidences)
102 }
103}