Skip to main content

causal_hub/random/datasets/trajectory/categorical/
evidence.rs

1use 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
10/// A struct representing a random evidence generator.
11pub 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    /// Creates a new `RngCatTrjEv` instance.
19    ///
20    /// # Arguments
21    ///
22    /// * `rng` - A mutable reference to a random number generator.
23    /// * `dataset` - A reference to the dataset.
24    /// * `p` - The probability of selecting an evidence.
25    ///
26    /// # Returns
27    ///
28    /// A new `RngCatTrjEv` instance.
29    pub fn new(rng: &'a mut R, dataset: &'a D, probability: f64) -> Result<Self> {
30        // Check that the probability is in [0, 1].
31        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        // Get shortened variable type.
48        use CatTrjEvT as E;
49
50        // Get times.
51        let times = self.dataset.times();
52        // Get events.
53        let events = self.dataset.values().rows();
54        // Zip times and events.
55        let times_events = times.into_iter().zip(events);
56
57        // Iterate over (time, event) pairs.
58        let evidence = times_events
59            .tuple_windows()
60            .filter_map(|((&start_time, v), (&end_time, _))| {
61                // Choose if the event is selected.
62                if !self.rng.random_bool(self.probability) {
63                    // If the event is not selected, skip it.
64                    return None;
65                }
66                // Select how many events to select.
67                let n = self.rng.random_range(1..=v.len());
68                // Sample the events.
69                let evidence = sample(self.rng, v.len(), n).into_iter().map(move |index| {
70                    // Get label and state.
71                    let (event, state) = (index, v[index] as usize);
72                    // Create the evidence.
73                    E::CertainPositiveInterval {
74                        event,
75                        state,
76                        start_time,
77                        end_time,
78                    }
79                });
80                // Return the evidences.
81                Some(evidence)
82            })
83            .flatten();
84
85        // Collect the evidence.
86        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}