use crate::error::Result;
use crate::forecast::spa_common::{
compute_means, compute_spa_pvalues, compute_standardized, compute_variances, find_best_model,
validate_model_data,
};
#[derive(Debug, Clone)]
pub struct MSPEAdjustedResult {
pub statistic: f64,
pub p_value_consistent: f64,
pub p_value_upper: f64,
pub n_bootstrap: usize,
pub best_model_idx: Option<usize>,
}
pub fn mspe_adjusted_spa(
benchmark_errors: &[f64],
model_errors: &[Vec<f64>],
n_bootstrap: usize,
block_length: f64,
seed: Option<u64>,
) -> Result<MSPEAdjustedResult> {
let t = benchmark_errors.len();
validate_model_data(t, model_errors, "alternative model")?;
let f = compute_clark_west_differentials(benchmark_errors, model_errors);
let f_bar = compute_means(&f);
let variances = compute_variances(&f, &f_bar);
let standardized = compute_standardized(&f_bar, &variances, t);
let observed = standardized
.iter()
.cloned()
.fold(f64::NEG_INFINITY, f64::max);
let best_model_idx = find_best_model(&standardized);
let (p_value_consistent, p_value_upper) =
compute_spa_pvalues(&f, &f_bar, observed, n_bootstrap, block_length, seed);
Ok(MSPEAdjustedResult {
statistic: observed,
p_value_consistent,
p_value_upper,
n_bootstrap,
best_model_idx,
})
}
fn compute_clark_west_differentials(benchmark: &[f64], models: &[Vec<f64>]) -> Vec<Vec<f64>> {
models
.iter()
.map(|model_err| {
benchmark
.iter()
.zip(model_err.iter())
.map(|(e_b, e_k)| {
let squared_diff = e_b.powi(2) - e_k.powi(2);
let adjustment = (e_b - e_k).powi(2);
squared_diff + adjustment
})
.collect()
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mspe_adjusted_clearly_better_model() {
let benchmark: Vec<f64> = (0..100)
.map(|i| 2.0 + (i as f64 * 0.1).sin() * 0.1)
.collect();
let model: Vec<Vec<f64>> = vec![(0..100)
.map(|i| 0.5 + (i as f64 * 0.1).cos() * 0.1)
.collect()];
let result = mspe_adjusted_spa(&benchmark, &model, 499, 5.0, Some(42)).unwrap();
assert!(
result.p_value_consistent < 0.05,
"p_value {} should be < 0.05",
result.p_value_consistent
);
assert!(result.best_model_idx == Some(0));
}
#[test]
fn test_mspe_adjusted_no_better_model() {
let benchmark: Vec<f64> = (0..100).map(|i| (i as f64 * 0.1).sin()).collect();
let model1: Vec<f64> = benchmark
.iter()
.enumerate()
.map(|(i, &x)| x + 0.01 * (i as f64).cos())
.collect();
let model2: Vec<f64> = benchmark
.iter()
.enumerate()
.map(|(i, &x)| x - 0.01 * (i as f64).sin())
.collect();
let models = vec![model1, model2];
let result = mspe_adjusted_spa(&benchmark, &models, 499, 5.0, Some(42)).unwrap();
assert!(
result.p_value_upper > 0.01,
"p_value_upper {} should be > 0.01",
result.p_value_upper
);
}
#[test]
fn test_mspe_adjusted_multiple_models() {
let benchmark: Vec<f64> = (0..100)
.map(|i| 1.5 + (i as f64 * 0.1).sin() * 0.1)
.collect();
let models = vec![
(0..100)
.map(|i| 1.2 + (i as f64 * 0.1).cos() * 0.1)
.collect(),
(0..100)
.map(|i| 0.5 + (i as f64 * 0.1).sin() * 0.05)
.collect(),
(0..100)
.map(|i| 2.0 + (i as f64 * 0.1).cos() * 0.1)
.collect(),
];
let result = mspe_adjusted_spa(&benchmark, &models, 499, 3.0, Some(42)).unwrap();
assert_eq!(result.best_model_idx, Some(1));
}
#[test]
fn test_mspe_adjusted_empty_benchmark() {
let benchmark: Vec<f64> = vec![];
let models = vec![vec![1.0, 2.0, 3.0]];
assert!(mspe_adjusted_spa(&benchmark, &models, 100, 3.0, None).is_err());
}
#[test]
fn test_mspe_adjusted_no_models() {
let benchmark: Vec<f64> = vec![1.0, 2.0, 3.0];
let models: Vec<Vec<f64>> = vec![];
assert!(mspe_adjusted_spa(&benchmark, &models, 100, 3.0, None).is_err());
}
#[test]
fn test_mspe_adjusted_length_mismatch() {
let benchmark: Vec<f64> = vec![1.0, 2.0, 3.0];
let models = vec![vec![1.0, 2.0]];
assert!(mspe_adjusted_spa(&benchmark, &models, 100, 3.0, None).is_err());
}
#[test]
fn test_mspe_adjusted_produces_different_results_than_spa() {
let benchmark: Vec<f64> = (0..100)
.map(|i| 1.5 + (i as f64 * 0.1).sin() * 0.2)
.collect();
let model: Vec<Vec<f64>> = vec![(0..100)
.map(|i| 1.0 + (i as f64 * 0.1).cos() * 0.2)
.collect()];
let result = mspe_adjusted_spa(&benchmark, &model, 499, 5.0, Some(42)).unwrap();
assert!(result.statistic.is_finite());
assert!(result.p_value_consistent >= 0.0 && result.p_value_consistent <= 1.0);
assert!(result.p_value_upper >= 0.0 && result.p_value_upper <= 1.0);
}
}