#[cfg(feature = "serde1")]
use serde_derive::{Deserialize, Serialize};
use crate::data::CategoricalDatum;
use crate::data::DataOrSuffStat;
use crate::dist::Categorical;
use crate::traits::SuffStat;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde1", derive(Serialize, Deserialize))]
pub struct CategoricalSuffStat {
n: usize,
counts: Vec<f64>,
}
impl CategoricalSuffStat {
#[inline]
pub fn new(k: usize) -> Self {
CategoricalSuffStat {
n: 0,
counts: vec![0.0; k],
}
}
#[inline]
pub fn n(&self) -> usize {
self.n
}
#[inline]
pub fn counts(&self) -> &Vec<f64> {
&self.counts
}
}
impl<'a, X> Into<DataOrSuffStat<'a, X, Categorical>> for &'a CategoricalSuffStat
where
X: CategoricalDatum,
{
fn into(self) -> DataOrSuffStat<'a, X, Categorical> {
DataOrSuffStat::SuffStat(self)
}
}
impl<'a, X: CategoricalDatum> Into<DataOrSuffStat<'a, X, Categorical>>
for &'a Vec<X>
{
fn into(self) -> DataOrSuffStat<'a, X, Categorical> {
DataOrSuffStat::Data(self)
}
}
impl<X: CategoricalDatum> SuffStat<X> for CategoricalSuffStat {
fn n(&self) -> usize {
self.n
}
fn observe(&mut self, x: &X) {
let ix = x.into_usize();
self.n += 1;
self.counts[ix] += 1.0;
}
fn forget(&mut self, x: &X) {
let ix = x.into_usize();
self.n -= 1;
self.counts[ix] -= 1.0;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new() {
let sf = CategoricalSuffStat::new(4);
assert_eq!(sf.counts.len(), 4);
assert_eq!(sf.n, 0);
assert!(sf.counts.iter().all(|&ct| ct.abs() < 1E-12))
}
}