use crate::models::laplace::dist::Gaussian;
use crate::models::laplace::leaf::Leaf;
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct AdidaLeaf {
alpha: f64,
aggregation_period: usize,
buffer_sum: f64,
buffer_count: usize,
aggregated_ema: f64,
initialized: bool,
n: usize,
ss: f64,
mean_resid: f64,
label: String,
}
impl AdidaLeaf {
pub fn new(alpha: f64, k: usize) -> Self {
let k = k.max(1);
let label = format!("adida{k}");
Self {
alpha: alpha.clamp(1e-3, 1.0 - 1e-3),
aggregation_period: k,
buffer_sum: 0.0,
buffer_count: 0,
aggregated_ema: 0.0,
initialized: false,
n: 0,
ss: 0.0,
mean_resid: 0.0,
label,
}
}
fn sigma(&self) -> f64 {
if self.n < 2 {
return 1.0;
}
(self.ss / (self.n as f64 - 1.0)).sqrt().max(1e-9)
}
fn point(&self) -> f64 {
if !self.initialized {
return 0.0;
}
self.aggregated_ema / self.aggregation_period as f64
}
}
impl Leaf for AdidaLeaf {
fn name(&self) -> &'static str {
Box::leak(self.label.clone().into_boxed_str())
}
fn predict(&self, horizon: usize) -> Vec<Gaussian> {
let point = self.point();
let base = self.sigma();
(1..=horizon)
.map(|h| Gaussian::new(point, base * (h as f64).sqrt()))
.collect()
}
fn observe(&mut self, y: f64) {
let predicted = self.point();
let resid = y - predicted;
self.n += 1;
let delta = resid - self.mean_resid;
self.mean_resid += delta / self.n as f64;
self.ss += delta * (resid - self.mean_resid);
self.buffer_sum += y;
self.buffer_count += 1;
if self.buffer_count >= self.aggregation_period {
if !self.initialized {
self.aggregated_ema = self.buffer_sum;
self.initialized = true;
} else {
self.aggregated_ema =
self.alpha * self.buffer_sum + (1.0 - self.alpha) * self.aggregated_ema;
}
self.buffer_sum = 0.0;
self.buffer_count = 0;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn adida_aggregates_reduce_intermittency() {
let mut adida = AdidaLeaf::new(0.3, 7);
for _ in 0..80 {
adida.observe(10.0);
for _ in 0..6 {
adida.observe(0.0);
}
}
let point = adida.predict(1)[0].mean;
let expected = 10.0 / 7.0;
assert!(
(point - expected).abs() < 0.5,
"adida k=7 expected ~{expected:.3}, got {point:.3}"
);
}
#[test]
fn adida_k1_matches_plain_ses() {
let mut adida = AdidaLeaf::new(0.3, 1);
for y in [1.0, 2.0, 3.0, 4.0, 5.0] {
adida.observe(y);
}
let point = adida.predict(1)[0].mean;
assert!(
(1.0..=5.0).contains(&point),
"k=1 adida should be in-range, got {point}"
);
}
}