use crate::models::arima::fractional_difference;
use crate::models::laplace::dist::Gaussian;
use crate::models::laplace::leaf::Leaf;
const DEFAULT_WINDOW: usize = 60;
const WEIGHT_THRESHOLD: f64 = 1e-3;
pub struct FractionalDiffLeaf {
d: f64,
alpha_mean: f64,
alpha_diff: f64,
window: Vec<f64>,
max_window: usize,
mean: Option<f64>,
fd_ema: f64,
n: usize,
ss: f64,
mean_resid: f64,
}
impl FractionalDiffLeaf {
pub fn new(d: f64, alpha_mean: f64, alpha_diff: f64) -> Self {
Self {
d: d.clamp(0.05, 0.95),
alpha_mean: alpha_mean.clamp(1e-3, 1.0 - 1e-3),
alpha_diff: alpha_diff.clamp(1e-3, 1.0 - 1e-3),
window: Vec::with_capacity(DEFAULT_WINDOW),
max_window: DEFAULT_WINDOW,
mean: None,
fd_ema: 0.0,
n: 0,
ss: 0.0,
mean_resid: 0.0,
}
}
fn sigma(&self) -> f64 {
if self.n < 2 {
return 1.0;
}
(self.ss / (self.n as f64 - 1.0)).sqrt().max(1e-9)
}
}
impl Leaf for FractionalDiffLeaf {
fn name(&self) -> &'static str {
"frac_diff"
}
fn predict(&self, horizon: usize) -> Vec<Gaussian> {
let level = self.mean.unwrap_or(0.0);
let base = self.sigma();
(1..=horizon)
.map(|h| Gaussian::new(level, base * (h as f64).sqrt()))
.collect()
}
fn observe(&mut self, y: f64) {
let predicted = self.mean.map(|m| m + self.fd_ema).unwrap_or(y);
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.window.push(y);
if self.window.len() > self.max_window {
self.window.remove(0);
}
if self.window.len() >= 8 {
let fd = fractional_difference(&self.window, self.d, WEIGHT_THRESHOLD);
if let Some(&last) = fd.last() {
if last.is_finite() {
self.fd_ema = self.alpha_diff * last + (1.0 - self.alpha_diff) * self.fd_ema;
}
}
}
self.mean = Some(match self.mean {
Some(m) => self.alpha_mean * y + (1.0 - self.alpha_mean) * m,
None => y,
});
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cold_start_produces_finite_predictions() {
let mut leaf = FractionalDiffLeaf::new(0.4, 0.1, 0.1);
for y in [1.0, 2.0, 3.0] {
leaf.observe(y);
}
let preds = leaf.predict(5);
for (h, p) in preds.iter().enumerate() {
assert!(p.mean.is_finite(), "h={}: mean not finite", h + 1);
assert!(p.std.is_finite() && p.std > 0.0, "h={}: std invalid", h + 1);
}
}
#[test]
fn flat_series_forecast_is_finite_and_close_to_level() {
let mut leaf = FractionalDiffLeaf::new(0.4, 0.1, 0.1);
for _ in 0..80 {
leaf.observe(50.0);
}
let preds = leaf.predict(3);
for p in preds {
assert!(p.mean.is_finite(), "mean not finite: {}", p.mean);
assert!(
(p.mean - 50.0).abs() < 20.0,
"forecast drifted too far: {}",
p.mean
);
}
}
#[test]
fn linear_trend_produces_valid_forecast() {
let mut leaf = FractionalDiffLeaf::new(0.4, 0.1, 0.1);
for i in 0..100 {
leaf.observe(10.0 + i as f64);
}
let preds = leaf.predict(3);
for (h, p) in preds.iter().enumerate() {
assert!(p.mean.is_finite(), "h={}: mean not finite", h + 1);
assert!(p.std.is_finite() && p.std > 0.0, "h={}: std invalid", h + 1);
}
}
}