use std::collections::VecDeque;
use super::quantile::standard_normal_quantile;
use crate::core::TimeSeries;
use crate::error::Result;
use crate::models::laplace::dist::GaussianMixture;
use crate::models::laplace::LaplaceForecaster;
use crate::models::{DistributionalForecaster, Forecaster};
const PIT_EPS: f64 = 1e-12;
pub struct Parade {
base: LaplaceForecaster,
k: usize,
pending: VecDeque<Vec<GaussianMixture>>,
pit: Vec<Option<f64>>,
z: Vec<Option<f64>>,
last_dists: Option<Vec<GaussianMixture>>,
}
impl Parade {
pub fn wrap(base: LaplaceForecaster, k: usize) -> Result<Self> {
assert!(k >= 1, "parade requires k >= 1");
let mut parade = Self {
base,
k,
pending: VecDeque::with_capacity(k),
pit: vec![None; k],
z: vec![None; k],
last_dists: None,
};
if let Ok(dists) = parade.base.forecast_dist(k) {
if dists.len() == k {
parade.pending.push_back(dists.clone());
parade.last_dists = Some(dists);
}
}
Ok(parade)
}
pub fn fit_and_wrap(
mut base: LaplaceForecaster,
series: &TimeSeries,
k: usize,
) -> Result<Self> {
base.fit(series)?;
Self::wrap(base, k)
}
pub fn observe(&mut self, y: f64) -> Result<()> {
if !y.is_finite() {
self.pit = vec![None; self.k];
self.z = vec![None; self.k];
return Ok(());
}
let n = self.pending.len();
let mut pit = vec![None; self.k];
let mut z = vec![None; self.k];
for m in 1..=self.k {
if m > n {
break;
}
let d = &self.pending[n - m][m - 1];
let u = d.cdf(y);
if !u.is_finite() {
continue;
}
let u = u.clamp(PIT_EPS, 1.0 - PIT_EPS);
pit[m - 1] = Some(u);
z[m - 1] = Some(standard_normal_quantile(u));
}
self.pit = pit;
self.z = z;
let mut y_fed = y;
if n > 0 {
let d1 = &self.pending[n - 1][0];
let (mp, sp) = mixture_moments(d1);
if mp.is_finite() && sp.is_finite() {
let w = 1e12 * (1.0 + mp.abs() + sp);
y_fed = y_fed.clamp(mp - w, mp + w);
}
}
self.base.observe(y_fed)?;
let dists = self.base.forecast_dist(self.k)?;
if dists.len() == self.k {
self.pending.push_back(dists.clone());
if self.pending.len() > self.k {
self.pending.pop_front();
}
self.last_dists = Some(dists);
}
Ok(())
}
pub fn z(&self) -> &[Option<f64>] {
&self.z
}
pub fn pit(&self) -> &[Option<f64>] {
&self.pit
}
pub fn k(&self) -> usize {
self.k
}
pub fn forecast_dist(&self, h: usize) -> Result<Vec<GaussianMixture>> {
self.base.forecast_dist(h)
}
pub fn base(&self) -> &LaplaceForecaster {
&self.base
}
pub fn pending_one_step(&self) -> Option<&GaussianMixture> {
self.pending.back().and_then(|k_vec| k_vec.first())
}
}
fn mixture_moments(m: &GaussianMixture) -> (f64, f64) {
let comps = &m.components;
if comps.is_empty() {
return (0.0, 0.0);
}
let mut mean = 0.0;
let mut second = 0.0;
for (w, g) in comps.iter() {
mean += w * g.mean;
second += w * (g.std * g.std + g.mean * g.mean);
}
let var = (second - mean * mean).max(0.0);
(mean, var.sqrt())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::TimeSeries;
use chrono::{Duration, TimeZone, Utc};
fn synthetic_iid_gaussian(n: usize) -> TimeSeries {
let base = Utc.with_ymd_and_hms(2000, 1, 1, 0, 0, 0).unwrap();
let vals: Vec<f64> = (0..n)
.map(|i| {
let seed = (i as u64)
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
let u1 = ((seed >> 33) as f64 / (1u64 << 31) as f64)
.max(1e-12)
.min(1.0 - 1e-12);
let seed2 = seed.wrapping_mul(6_364_136_223_846_793_005).wrapping_add(1);
let u2 = ((seed2 >> 33) as f64 / (1u64 << 31) as f64)
.max(1e-12)
.min(1.0 - 1e-12);
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
})
.collect();
let stamps: Vec<_> = (0..n).map(|i| base + Duration::hours(i as i64)).collect();
TimeSeries::univariate(stamps, vals).unwrap()
}
#[test]
fn warmup_returns_none_before_maturity() {
let n = 200;
let ts = synthetic_iid_gaussian(n);
let train_len = 150;
let train = TimeSeries::univariate(
ts.timestamps()[..train_len].to_vec(),
ts.primary_values()[..train_len].to_vec(),
)
.unwrap();
let base = LaplaceForecaster::new().auto();
let mut parade = Parade::fit_and_wrap(base, &train, 4).unwrap();
parade.observe(ts.primary_values()[train_len]).unwrap();
assert!(parade.z()[0].is_some());
assert!(parade.z()[1].is_none());
assert!(parade.z()[3].is_none());
for i in 1..4 {
parade.observe(ts.primary_values()[train_len + i]).unwrap();
}
for (i, z) in parade.z().iter().enumerate() {
assert!(z.is_some(), "horizon {} still None after 4 ticks", i + 1);
}
}
#[test]
fn z_is_finite_and_bounded() {
let n = 400;
let ts = synthetic_iid_gaussian(n);
let train_len = 200;
let train = TimeSeries::univariate(
ts.timestamps()[..train_len].to_vec(),
ts.primary_values()[..train_len].to_vec(),
)
.unwrap();
let base = LaplaceForecaster::new().auto();
let mut parade = Parade::fit_and_wrap(base, &train, 4).unwrap();
for i in 0..100 {
parade.observe(ts.primary_values()[train_len + i]).unwrap();
}
for z in parade.z() {
let z = z.unwrap();
assert!(z.is_finite(), "z = {z}");
assert!(z.abs() <= 7.5, "|z| = {} exceeds parade clamp", z.abs());
}
}
#[test]
fn nan_observation_blanks_z_but_survives() {
let n = 100;
let ts = synthetic_iid_gaussian(n);
let base = LaplaceForecaster::new().auto();
let mut parade = Parade::fit_and_wrap(base, &ts, 4).unwrap();
parade.observe(ts.primary_values()[0]).unwrap();
parade.observe(f64::NAN).unwrap();
for z in parade.z() {
assert!(z.is_none());
}
parade.observe(ts.primary_values()[1]).unwrap();
assert!(parade.z()[0].is_some());
}
#[test]
fn forecast_dist_pass_through() {
let n = 100;
let ts = synthetic_iid_gaussian(n);
let base = LaplaceForecaster::new().auto();
let parade = Parade::fit_and_wrap(base, &ts, 4).unwrap();
let dists = parade.forecast_dist(4).unwrap();
assert_eq!(dists.len(), 4);
for d in &dists {
let (mean, std) = mixture_moments(d);
assert!(mean.is_finite() && std.is_finite() && std > 0.0);
}
}
}