causal_hub/random/models/bayesian_network/categorical/
parameters.rs1use ndarray::prelude::*;
2use rand::{Rng, RngExt};
3use rand_distr::Gamma;
4
5use crate::{
6 models::{CatCPD, CatSupport},
7 random::Random,
8 types::{Error, Result},
9};
10
11pub struct RngCatCPD<'a, R>
13where
14 R: Rng,
15{
16 rng: &'a mut R,
17 support: &'a CatSupport,
18 conditioning_support: &'a CatSupport,
19 alpha: f64,
20}
21
22impl<'a, R> RngCatCPD<'a, R>
23where
24 R: Rng,
25{
26 pub fn new(
44 rng: &'a mut R,
45 support: &'a CatSupport,
46 conditioning_support: &'a CatSupport,
47 alpha: f64,
48 ) -> Result<Self> {
49 if alpha <= 0.0 {
51 return Err(Error::InvalidParameter("alpha", "must be positive"));
52 }
53
54 Ok(Self {
55 rng,
56 support,
57 conditioning_support,
58 alpha,
59 })
60 }
61}
62
63impl<R> Random for RngCatCPD<'_, R>
64where
65 R: Rng,
66{
67 type Output = Result<CatCPD>;
68
69 fn random(&mut self) -> Self::Output {
70 let model = self.support.values().map(|v| v.len()).product();
72 let n = self
73 .conditioning_support
74 .values()
75 .map(|v| v.len())
76 .product();
77
78 let gamma = Gamma::new(self.alpha, 1.0)
80 .map_err(|evidence| Error::InvalidParameter("alpha", &evidence.to_string()))?;
81
82 let mut parameters = Array::from_shape_fn((n, model), |_| self.rng.sample(gamma));
84 parameters /= ¶meters.sum_axis(Axis(1)).insert_axis(Axis(1));
86
87 CatCPD::new(
89 self.support.clone(),
90 self.conditioning_support.clone(),
91 parameters,
92 )
93 }
94}