use super::mahalanobis::{AnomalyOutput, MahalanobisConfig, MahalanobisScorer};
use super::parade::Parade;
use crate::core::TimeSeries;
use crate::error::Result;
use crate::models::laplace::dist::GaussianMixture;
use crate::models::laplace::LaplaceForecaster;
pub struct ZBank {
engines: Vec<Vec<Parade>>,
sigmas: Vec<f64>,
strides: Vec<usize>,
k: usize,
t: usize,
last_z: Vec<Option<f64>>,
last_dists: Option<Vec<GaussianMixture>>,
}
impl ZBank {
pub fn effective_k(&self) -> usize {
self.sigmas.len() * self.strides.len() * self.k
}
pub fn n_engines(&self) -> usize {
self.engines.iter().map(|v| v.len()).sum()
}
pub fn z(&self) -> &[Option<f64>] {
&self.last_z
}
pub fn forecast_dist(&self, h: usize) -> Result<Vec<GaussianMixture>> {
self.engines[0][0].forecast_dist(h)
}
pub fn observe(&mut self, y: f64) -> Result<()> {
let k = self.k;
let n_sig = self.sigmas.len();
let n_str = self.strides.len();
for e in self.last_z.iter_mut() {
*e = None;
}
for s_idx in 0..n_sig {
for st_idx in 0..n_str {
let stride = self.strides[st_idx];
let phase = self.t % stride;
let bank_idx = s_idx * n_str + st_idx;
let parade = &mut self.engines[bank_idx][phase];
parade.observe(y)?;
if s_idx == 0 && stride == 1 {
if let Ok(dists) = parade.forecast_dist(k) {
self.last_dists = Some(dists);
}
}
let base = bank_idx * k;
for (h, opt_z) in parade.z().iter().enumerate() {
self.last_z[base + h] = *opt_z;
}
}
}
self.t += 1;
Ok(())
}
}
pub struct ZBankBuilder {
k: usize,
sigmas: Vec<f64>,
strides: Vec<usize>,
}
impl ZBankBuilder {
pub fn new(k: usize) -> Self {
Self {
k,
sigmas: vec![0.03, 0.003],
strides: vec![1, 4, 16],
}
}
pub fn sigmas(mut self, sigmas: Vec<f64>) -> Self {
self.sigmas = sigmas;
self
}
pub fn strides(mut self, strides: Vec<usize>) -> Self {
self.strides = strides;
self
}
pub fn build<F>(self, mut base_factory: F, series: &TimeSeries) -> Result<ZBank>
where
F: FnMut() -> LaplaceForecaster,
{
assert!(self.k >= 1);
assert!(!self.sigmas.is_empty());
assert!(!self.strides.is_empty());
assert!(
self.strides.contains(&1),
"strides must include 1 (the pass-through engine)"
);
let n_sig = self.sigmas.len();
let n_str = self.strides.len();
let mut engines: Vec<Vec<Parade>> = Vec::with_capacity(n_sig * n_str);
for _sig in &self.sigmas {
for &stride in &self.strides {
let mut phase_copies = Vec::with_capacity(stride);
for _ph in 0..stride {
let base = base_factory();
let parade = Parade::fit_and_wrap(base, series, self.k)?;
phase_copies.push(parade);
}
engines.push(phase_copies);
}
}
Ok(ZBank {
engines,
k: self.k,
last_z: vec![None; n_sig * n_str * self.k],
last_dists: None,
sigmas: self.sigmas,
strides: self.strides,
t: 0,
})
}
}
pub struct ZBankDetector {
bank: ZBank,
scorer: MahalanobisScorer,
pend1: Option<GaussianMixture>,
}
impl ZBankDetector {
pub fn wrap(bank: ZBank, cfg: MahalanobisConfig) -> Self {
let k_eff = bank.effective_k();
let scorer = MahalanobisScorer::with_k(k_eff, cfg);
Self {
bank,
scorer,
pend1: None,
}
}
pub fn observe(&mut self, y: f64) -> Result<()> {
if !y.is_finite() {
self.scorer.blank();
return Ok(());
}
let nlp = self.pend1.as_ref().map(|d| {
let lp = d.logpdf(y);
if lp.is_finite() {
-lp
} else {
1e6
}
});
self.bank.observe(y)?;
self.pend1 = self
.bank
.last_dists
.as_ref()
.and_then(|dists| dists.first().cloned());
let z_opt = self.bank.z();
let z: Vec<f64> = if z_opt.iter().all(|v| v.is_some()) {
z_opt.iter().map(|v| v.unwrap()).collect()
} else {
self.scorer.blank();
return Ok(());
};
self.scorer.score_z(&z, nlp);
Ok(())
}
pub fn state(&self) -> &AnomalyOutput {
self.scorer.last()
}
pub fn bank(&self) -> &ZBank {
&self.bank
}
pub fn forecast_dist(&self, h: usize) -> Result<Vec<GaussianMixture>> {
self.bank.forecast_dist(h)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::TimeSeries;
use chrono::{Duration, TimeZone, Utc};
fn synthetic_iid_gaussian(n: usize) -> 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).clamp(1e-12, 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).clamp(1e-12, 1.0 - 1e-12);
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
})
.collect()
}
fn ts_from(vals: Vec<f64>) -> TimeSeries {
let base = Utc.with_ymd_and_hms(2000, 1, 1, 0, 0, 0).unwrap();
let stamps: Vec<_> = (0..vals.len())
.map(|i| base + Duration::hours(i as i64))
.collect();
TimeSeries::univariate(stamps, vals).unwrap()
}
#[test]
fn zbank_effective_k_matches_config() {
let vals = synthetic_iid_gaussian(200);
let train = ts_from(vals[..100].to_vec());
let bank = ZBankBuilder::new(2)
.sigmas(vec![0.03, 0.003])
.strides(vec![1, 4])
.build(|| LaplaceForecaster::new().auto(), &train)
.unwrap();
assert_eq!(bank.effective_k(), 8);
assert_eq!(bank.n_engines(), 2 * (1 + 4)); }
#[test]
fn zbank_detector_flags_spike() {
let mut vals = synthetic_iid_gaussian(400);
vals[350] = 25.0; let train = ts_from(vals[..150].to_vec());
let bank = ZBankBuilder::new(2)
.sigmas(vec![0.03])
.strides(vec![1, 2])
.build(|| LaplaceForecaster::new().auto(), &train)
.unwrap();
let mut det = ZBankDetector::wrap(bank, MahalanobisConfig::new(2));
let mut min_p_before = 1.0;
for &y in &vals[150..350] {
det.observe(y).unwrap();
if let Some(p) = det.state().p_value {
if p < min_p_before {
min_p_before = p;
}
}
}
det.observe(vals[350]).unwrap();
let p_spike = det.state().p_value.unwrap_or(1.0);
assert!(
p_spike < 0.05,
"spike p-value {p_spike} should be < 0.05 (min pre-spike {min_p_before})",
);
}
#[test]
fn zbank_strides_must_include_one() {
let vals = synthetic_iid_gaussian(50);
let train = ts_from(vals.clone());
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
ZBankBuilder::new(2)
.strides(vec![2, 4])
.build(|| LaplaceForecaster::new().auto(), &train)
}));
assert!(result.is_err(), "should panic when strides don't include 1");
}
}