use ndarray::prelude::*;
use rand::{Rng, RngExt};
use rand_distr::Gamma;
use crate::{
models::CatCPD,
random::Random,
types::{Error, Result, States},
};
pub struct RngCatCPD<'a, R>
where
R: Rng,
{
rng: &'a mut R,
states: &'a States,
conditioning_states: &'a States,
alpha: f64,
}
impl<'a, R> RngCatCPD<'a, R>
where
R: Rng,
{
pub fn new(
rng: &'a mut R,
states: &'a States,
conditioning_states: &'a States,
alpha: f64,
) -> Result<Self> {
if alpha <= 0.0 {
return Err(Error::InvalidParameter("alpha", "must be positive"));
}
Ok(Self {
rng,
states,
conditioning_states,
alpha,
})
}
}
impl<R> Random for RngCatCPD<'_, R>
where
R: Rng,
{
type Output = Result<CatCPD>;
fn random(&mut self) -> Self::Output {
let m = self.states.values().map(|v| v.len()).product();
let n = self.conditioning_states.values().map(|v| v.len()).product();
let gamma = Gamma::new(self.alpha, 1.0)
.map_err(|e| Error::InvalidParameter("alpha", &e.to_string()))?;
let mut parameters = Array::from_shape_fn((n, m), |_| self.rng.sample(gamma));
parameters /= ¶meters.sum_axis(Axis(1)).insert_axis(Axis(1));
CatCPD::new(
self.states.clone(),
self.conditioning_states.clone(),
parameters,
)
}
}