use num::{Float, FromPrimitive};
use std::ops::{AddAssign, SubAssign};
use crate::count::Count;
use crate::stats::{Revertable, RollableUnivariate, Univariate};
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Default, Debug, Serialize, Deserialize)]
pub struct Mean<F: Float + FromPrimitive + AddAssign + SubAssign> {
pub mean: F,
pub n: Count<F>,
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> Mean<F> {
pub fn new() -> Self {
Self {
mean: F::from_f64(0.0).unwrap(),
n: Count::new(),
}
}
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> Univariate<F> for Mean<F> {
fn update(&mut self, x: F) {
self.n.update(x);
self.mean += (F::from_f64(1.).unwrap() / self.n.get()) * (x - self.mean);
}
fn get(&self) -> F {
self.mean
}
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> Revertable<F> for Mean<F> {
fn revert(&mut self, x: F) -> Result<(), &'static str> {
match self.n.revert(x) {
Ok(it) => it,
Err(err) => return Err(err),
};
let count = self.n.get();
if count == F::from_f64(0.).unwrap() {
self.mean = F::from_f64(0.0).unwrap();
} else {
self.mean -= (F::from_f64(1.0).unwrap() / count) * (x - self.mean);
}
Ok(())
}
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> RollableUnivariate<F> for Mean<F> {}