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 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); 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 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 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 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 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 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
174fn 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 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 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 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}