Skip to main content

chronos_ts/
diagnostics.rs

1use crate::decomposition::ProphetDecomposition;
2use crate::errors::{ChronosError, Result};
3use chrono::NaiveDate;
4use ndarray::Array1;
5use serde::{Deserialize, Serialize};
6use std::collections::BTreeMap;
7
8pub struct LjungBoxResult {
9    pub q_stat: f64,
10    pub p_value: f64,
11    pub lags: usize,
12}
13
14pub struct JarqueBeraResult {
15    pub jb_stat: f64,
16    pub p_value: f64,
17    pub skewness: f64,
18    pub kurtosis: f64,
19}
20
21pub struct DiagnosticsResult {
22    pub acf: Array1<f64>,
23    pub pacf: Array1<f64>,
24    pub ljung_box: LjungBoxResult,
25    pub jarque_bera: JarqueBeraResult,
26}
27
28pub struct ResidualDiagnostics;
29
30impl ResidualDiagnostics {
31    /// Sample Autocorrelation Function (ACF) up to `max_lag`
32    pub fn acf(residuals: &Array1<f64>, max_lag: usize) -> Result<Array1<f64>> {
33        let n = residuals.len();
34        if n <= max_lag {
35            return Err(ChronosError::InsufficientData {
36                required: max_lag + 1,
37                found: n,
38            });
39        }
40
41        let mean = residuals.mean().unwrap_or(0.0);
42        let var = residuals.iter().map(|&x| (x - mean).powi(2)).sum::<f64>();
43
44        if var == 0.0 {
45            return Err(ChronosError::InvalidParameters(
46                "Zero variance in residual series".into(),
47            ));
48        }
49
50        let mut acf_vals = Vec::with_capacity(max_lag + 1);
51        acf_vals.push(1.0); // Lag 0 is always 1.0
52
53        for lag in 1..=max_lag {
54            let mut cov = 0.0;
55            for t in lag..n {
56                cov += (residuals[t] - mean) * (residuals[t - lag] - mean);
57            }
58            acf_vals.push(cov / var);
59        }
60
61        Ok(Array1::from_vec(acf_vals))
62    }
63
64    /// Partial Autocorrelation Function (PACF) using Levinson-Durbin Recursion
65    pub fn pacf(residuals: &Array1<f64>, max_lag: usize) -> Result<Array1<f64>> {
66        let acf_vals = Self::acf(residuals, max_lag)?;
67        let mut pacf_vals = Vec::with_capacity(max_lag + 1);
68        pacf_vals.push(1.0);
69
70        if max_lag == 0 {
71            return Ok(Array1::from_vec(pacf_vals));
72        }
73
74        pacf_vals.push(acf_vals[1]);
75
76        let mut phi = vec![vec![0.0; max_lag + 1]; max_lag + 1];
77        phi[1][1] = acf_vals[1];
78
79        for k in 2..=max_lag {
80            let mut num = acf_vals[k];
81            let mut den = 1.0;
82
83            for j in 1..k {
84                num -= phi[k - 1][j] * acf_vals[k - j];
85                den -= phi[k - 1][j] * acf_vals[j];
86            }
87
88            let phi_kk = num / den;
89            phi[k][k] = phi_kk;
90            pacf_vals.push(phi_kk);
91
92            for j in 1..k {
93                phi[k][j] = phi[k - 1][j] - phi_kk * phi[k - 1][k - j];
94            }
95        }
96
97        Ok(Array1::from_vec(pacf_vals))
98    }
99
100    /// Ljung-Box Q-Test for Autocorrelation
101    pub fn ljung_box(residuals: &Array1<f64>, lags: usize) -> Result<LjungBoxResult> {
102        let n = residuals.len() as f64;
103        let acf_vals = Self::acf(residuals, lags)?;
104
105        let mut q_stat = 0.0;
106        for k in 1..=lags {
107            let r_k = acf_vals[k];
108            q_stat += (r_k * r_k) / (n - k as f64);
109        }
110        q_stat *= n * (n + 2.0);
111
112        // Chi-square survival function approximation (1 degree of freedom per lag)
113        let p_value = chi2_sf(q_stat, lags as f64);
114
115        Ok(LjungBoxResult {
116            q_stat,
117            p_value,
118            lags,
119        })
120    }
121
122    /// Jarque-Bera Test for Normality (Skewness & Excess Kurtosis)
123    pub fn jarque_bera(residuals: &Array1<f64>) -> Result<JarqueBeraResult> {
124        let n = residuals.len() as f64;
125        if n < 4.0 {
126            return Err(ChronosError::InsufficientData {
127                required: 4,
128                found: n as usize,
129            });
130        }
131
132        let mean = residuals.mean().unwrap_or(0.0);
133        let m2 = residuals.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / n;
134        let m3 = residuals.iter().map(|&x| (x - mean).powi(3)).sum::<f64>() / n;
135        let m4 = residuals.iter().map(|&x| (x - mean).powi(4)).sum::<f64>() / n;
136
137        if m2 == 0.0 {
138            return Err(ChronosError::InvalidParameters(
139                "Zero residual variance".into(),
140            ));
141        }
142
143        let skewness = m3 / m2.powf(1.5);
144        let kurtosis = m4 / m2.powi(2);
145        let excess_kurtosis = kurtosis - 3.0;
146
147        let jb_stat = (n / 6.0) * (skewness.powi(2) + (0.25 * excess_kurtosis.powi(2)));
148        let p_value = chi2_sf(jb_stat, 2.0);
149
150        Ok(JarqueBeraResult {
151            jb_stat,
152            p_value,
153            skewness,
154            kurtosis,
155        })
156    }
157
158    /// Full Diagnostic Pass across residuals
159    pub fn evaluate(residuals: &Array1<f64>, max_lags: usize) -> Result<DiagnosticsResult> {
160        let acf_vals = Self::acf(residuals, max_lags)?;
161        let pacf_vals = Self::pacf(residuals, max_lags)?;
162        let lb = Self::ljung_box(residuals, max_lags)?;
163        let jb = Self::jarque_bera(residuals)?;
164
165        Ok(DiagnosticsResult {
166            acf: acf_vals,
167            pacf: pacf_vals,
168            ljung_box: lb,
169            jarque_bera: jb,
170        })
171    }
172}
173
174/// Lower incomplete gamma function approximation for Chi-Square distribution survival calculation
175fn chi2_sf(x: f64, df: f64) -> f64 {
176    if x <= 0.0 {
177        return 1.0;
178    }
179    let a = df / 2.0;
180    let z = x / 2.0;
181
182    // Regularized incomplete gamma upper bound approximation
183    let mut sum = 0.0;
184    let mut term = 1.0 / a;
185    sum += term;
186    for i in 1..100 {
187        term *= z / (a + i as f64);
188        sum += term;
189        if term < 1e-10 {
190            break;
191        }
192    }
193
194    let gamma_sf = (-z + a * z.ln() - gamma_log(a)).exp() * sum;
195    gamma_sf.clamp(0.0, 1.0)
196}
197
198fn gamma_log(a: f64) -> f64 {
199    // Lanczos approximation for ln(gamma(a))
200    let coeffs = [
201        76.18009172947146,
202        -86.50532032941677,
203        24.01409824083091,
204        -1.231739572450155,
205        0.1208650973866179e-2,
206        -0.5395239384953e-5,
207    ];
208    let mut y = a;
209    let mut tmp = a + 5.5;
210    tmp -= (a + 0.5) * tmp.ln();
211    let mut ser = 1.000000000190015;
212    for c in &coeffs {
213        y += 1.0;
214        ser += c / y;
215    }
216    -tmp + (2.5066282746310005 * ser / a).ln()
217}
218
219#[derive(Debug, Clone, Serialize, Deserialize)]
220pub struct HorizonMetrics {
221    pub horizon_step: usize,
222    pub mae: f64,
223    pub rmse: f64,
224    pub mape: f64,
225    pub sample_count: usize,
226}
227
228#[derive(Debug, Clone, Serialize, Deserialize)]
229pub struct CrossValidationReport {
230    pub total_folds: usize,
231    pub horizon_metrics: Vec<HorizonMetrics>,
232    pub overall_mae: f64,
233    pub overall_rmse: f64,
234}
235
236pub struct CrossValidationEvaluator {
237    pub horizon: usize,
238    pub initial: usize,
239    pub step: usize,
240}
241
242impl CrossValidationEvaluator {
243    pub fn new(horizon: usize, initial: usize, step: usize) -> Self {
244        Self {
245            horizon,
246            initial,
247            step,
248        }
249    }
250
251    /// Computes horizon-level degradation metrics across rolling cross-validation folds
252    pub fn evaluate(
253        &self,
254        model: &ProphetDecomposition,
255        dates: &[NaiveDate],
256        y: &Array1<f64>,
257    ) -> Result<CrossValidationReport> {
258        let n = y.len();
259        if n < self.initial + self.horizon {
260            return Err(ChronosError::InsufficientData {
261                required: self.initial + self.horizon,
262                found: n,
263            });
264        }
265
266        let mut horizon_errors: BTreeMap<usize, Vec<(f64, f64)>> = BTreeMap::new();
267        let mut fold_count = 0;
268
269        let mut train_end = self.initial;
270        while train_end + self.horizon <= n {
271            let train_dates = &dates[..train_end];
272            let train_y = y.slice(ndarray::s![..train_end]).to_owned();
273
274            let val_dates = &dates[train_end..train_end + self.horizon];
275            let val_y = y.slice(ndarray::s![train_end..train_end + self.horizon]);
276
277            let mut fold_model = model.clone();
278            fold_model.fit(train_dates, &train_y, None, None)?;
279
280            let pred = fold_model.predict(val_dates)?;
281
282            for h in 0..self.horizon {
283                let actual = val_y[h];
284                let predicted = pred.yhat[h];
285                horizon_errors
286                    .entry(h + 1)
287                    .or_default()
288                    .push((actual, predicted));
289            }
290
291            fold_count += 1;
292            train_end += self.step;
293        }
294
295        if fold_count == 0 {
296            return Err(ChronosError::InvalidParameters(
297                "No cross-validation folds were executed".into(),
298            ));
299        }
300
301        let mut horizon_metrics = Vec::with_capacity(self.horizon);
302        let mut total_absolute_error = 0.0;
303        let mut total_squared_error = 0.0;
304        let mut total_points = 0;
305
306        for (h, pairs) in horizon_errors {
307            let count = pairs.len();
308            let mut sum_abs_err = 0.0;
309            let mut sum_sq_err = 0.0;
310            let mut sum_pct_err = 0.0;
311
312            for (actual, predicted) in &pairs {
313                let err = (actual - predicted).abs();
314                sum_abs_err += err;
315                sum_sq_err += err * err;
316                if actual.abs() > 1e-8 {
317                    sum_pct_err += err / actual.abs();
318                }
319
320                total_absolute_error += err;
321                total_squared_error += err * err;
322                total_points += 1;
323            }
324
325            let mae = sum_abs_err / (count as f64);
326            let rmse = (sum_sq_err / (count as f64)).sqrt();
327            let mape = (sum_pct_err / (count as f64)) * 100.0;
328
329            horizon_metrics.push(HorizonMetrics {
330                horizon_step: h,
331                mae,
332                rmse,
333                mape,
334                sample_count: count,
335            });
336        }
337
338        let overall_mae = total_absolute_error / (total_points as f64);
339        let overall_rmse = (total_squared_error / (total_points as f64)).sqrt();
340
341        Ok(CrossValidationReport {
342            total_folds: fold_count,
343            horizon_metrics,
344            overall_mae,
345            overall_rmse,
346        })
347    }
348}