use rand::prelude::*;
use crate::{
labels,
models::{BN, GaussBN},
random::{Random, RngDag, RngGaussCPD},
set,
types::{Error, Labels, Result},
};
pub struct RngGaussBN<'a, R>
where
R: Rng,
{
rng: &'a mut R,
labels: &'a Labels,
s_a: f64,
s_b: f64,
e: f64,
p: f64,
}
impl<'a, R> RngGaussBN<'a, R>
where
R: Rng,
{
pub fn new(
rng: &'a mut R,
labels: &'a Labels,
s_a: f64,
s_b: f64,
e: f64,
p: f64,
) -> Result<Self> {
if s_a <= 0.0 {
return Err(Error::InvalidParameter("s_a", "must be positive"));
}
if s_b <= 0.0 {
return Err(Error::InvalidParameter("s_b", "must be positive"));
}
if e <= 0.0 {
return Err(Error::InvalidParameter("e", "must be positive"));
}
if !(0.0..=1.0).contains(&p) {
return Err(Error::InvalidParameter("p", "must be in [0, 1]"));
}
Ok(Self {
rng,
labels,
s_a,
s_b,
e,
p,
})
}
}
impl<R> Random for RngGaussBN<'_, R>
where
R: Rng,
{
type Output = Result<GaussBN>;
fn random(&mut self) -> Self::Output {
let graph = RngDag::new(self.rng, self.labels, self.p)?.random()?;
let cpds = self
.labels
.iter()
.enumerate()
.map(|(i, x)| {
let pa_i = graph.parents(&set![i])?;
let labels = labels![x.clone()];
let conditioning_labels =
pa_i.into_iter().map(|j| self.labels[j].clone()).collect();
RngGaussCPD::new(
self.rng,
&labels,
&conditioning_labels,
self.s_a,
self.s_b,
self.e,
)?
.random()
})
.collect::<Result<Vec<_>>>()?;
GaussBN::new(graph, cpds)
}
}