use crate::error::{ForecastError, Result};
use crate::postprocess::PredictionIntervals;
#[derive(Debug, Clone)]
pub struct EnbPiResult {
residuals: Vec<f64>,
quantile_value: f64,
coverage: f64,
n_models: usize,
}
impl EnbPiResult {
pub fn residuals(&self) -> &[f64] {
&self.residuals
}
pub fn quantile_value(&self) -> f64 {
self.quantile_value
}
pub fn coverage(&self) -> f64 {
self.coverage
}
pub fn n_models(&self) -> usize {
self.n_models
}
pub fn predict_intervals(&self, point_forecasts: &[f64]) -> Result<PredictionIntervals> {
let lower: Vec<f64> = point_forecasts
.iter()
.map(|&p| p - self.quantile_value)
.collect();
let upper: Vec<f64> = point_forecasts
.iter()
.map(|&p| p + self.quantile_value)
.collect();
PredictionIntervals::from_bounds(lower, upper, self.coverage)
}
pub fn update(
&mut self,
new_forecasts: &[f64],
new_actuals: &[f64],
max_window: usize,
) -> Result<()> {
if new_forecasts.len() != new_actuals.len() {
return Err(ForecastError::DimensionMismatch {
expected: new_forecasts.len(),
got: new_actuals.len(),
});
}
let mut all: Vec<f64> = self.residuals.clone();
for (f, a) in new_forecasts.iter().zip(new_actuals.iter()) {
all.push((f - a).abs());
}
if all.len() > max_window {
let drop = all.len() - max_window;
all.drain(..drop);
}
all.sort_by(|a, b| a.partial_cmp(b).unwrap());
let n = all.len();
if n == 0 {
return Err(ForecastError::EmptyData);
}
let level = self.coverage * (n as f64 + 1.0) / n as f64;
let level = level.min(1.0);
let idx = ((n as f64) * level).ceil() as usize;
let idx = idx.saturating_sub(1).min(n - 1);
self.quantile_value = all[idx];
self.residuals = all;
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct EnbPiPredictor {
coverage: f64,
}
impl EnbPiPredictor {
pub fn new(coverage: f64) -> Self {
assert!(
coverage > 0.0 && coverage < 1.0,
"coverage must be in (0, 1)"
);
Self { coverage }
}
pub fn coverage(&self) -> f64 {
self.coverage
}
pub fn fit(
&self,
bootstrap_predictions: &[f64],
inclusion_mask: &[bool],
actuals: &[f64],
n_models: usize,
) -> Result<EnbPiResult> {
if n_models == 0 {
return Err(ForecastError::InvalidParameter(
"n_models must be ≥ 1".to_string(),
));
}
let n = actuals.len();
if n == 0 {
return Err(ForecastError::EmptyData);
}
let expected = n * n_models;
if bootstrap_predictions.len() != expected {
return Err(ForecastError::DimensionMismatch {
expected,
got: bootstrap_predictions.len(),
});
}
if inclusion_mask.len() != expected {
return Err(ForecastError::DimensionMismatch {
expected,
got: inclusion_mask.len(),
});
}
let mut residuals: Vec<f64> = Vec::with_capacity(n);
for i in 0..n {
let mut sum = 0.0;
let mut count = 0usize;
let mut full_sum = 0.0;
for k in 0..n_models {
let pred = bootstrap_predictions[i * n_models + k];
full_sum += pred;
if !inclusion_mask[i * n_models + k] {
sum += pred;
count += 1;
}
}
let loo_pred = if count > 0 {
sum / count as f64
} else {
full_sum / n_models as f64
};
residuals.push((loo_pred - actuals[i]).abs());
}
residuals.sort_by(|a, b| a.partial_cmp(b).unwrap());
let level = self.coverage * (n as f64 + 1.0) / n as f64;
let level = level.min(1.0);
let idx = ((n as f64) * level).ceil() as usize;
let idx = idx.saturating_sub(1).min(n - 1);
let quantile_value = residuals[idx];
Ok(EnbPiResult {
residuals,
quantile_value,
coverage: self.coverage,
n_models,
})
}
pub fn ensemble_point_forecast(
&self,
test_predictions: &[f64],
n_models: usize,
) -> Result<Vec<f64>> {
if n_models == 0 {
return Err(ForecastError::InvalidParameter(
"n_models must be ≥ 1".to_string(),
));
}
if test_predictions.len() % n_models != 0 {
return Err(ForecastError::DimensionMismatch {
expected: test_predictions.len() / n_models * n_models,
got: test_predictions.len(),
});
}
let n_test = test_predictions.len() / n_models;
let mut result = Vec::with_capacity(n_test);
for i in 0..n_test {
let mut sum = 0.0;
for k in 0..n_models {
sum += test_predictions[i * n_models + k];
}
result.push(sum / n_models as f64);
}
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fit_recovers_noise_magnitude() {
let n = 100;
let n_models = 5;
let actuals: Vec<f64> = (0..n).map(|i| (i as f64 * 0.1).sin()).collect();
let model_offsets = [-0.2, -0.1, 0.0, 0.1, 0.2];
let mut predictions: Vec<f64> = Vec::with_capacity(n * n_models);
for i in 0..n {
for k in 0..n_models {
predictions.push(actuals[i] + model_offsets[k]);
}
}
let mut mask: Vec<bool> = Vec::with_capacity(n * n_models);
for i in 0..n {
for k in 0..n_models {
mask.push((i + k) % 2 == 0);
}
}
let enbpi = EnbPiPredictor::new(0.90);
let result = enbpi.fit(&predictions, &mask, &actuals, n_models).unwrap();
assert_eq!(result.residuals().len(), n);
assert!(result.quantile_value() >= 0.0);
assert!(
result.quantile_value() < 0.4,
"quantile value {} too large for offsets ≤ 0.2",
result.quantile_value()
);
}
#[test]
fn predict_intervals_widens_around_point_forecast() {
let n = 50;
let n_models = 3;
let actuals: Vec<f64> = (0..n).map(|i| i as f64).collect();
let mut predictions: Vec<f64> = Vec::with_capacity(n * n_models);
for i in 0..n {
for _k in 0..n_models {
predictions.push(actuals[i]);
}
}
let mask: Vec<bool> = vec![false; n * n_models];
let mut perturbed = predictions.clone();
for p in perturbed.iter_mut() {
*p += 0.5;
}
let enbpi = EnbPiPredictor::new(0.90);
let result = enbpi.fit(&perturbed, &mask, &actuals, n_models).unwrap();
assert!((result.quantile_value() - 0.5).abs() < 1e-12);
let test_pf = vec![10.0, 20.0, 30.0];
let intervals = result.predict_intervals(&test_pf).unwrap();
assert!((intervals.lower()[0] - 9.5).abs() < 1e-12);
assert!((intervals.upper()[0] - 10.5).abs() < 1e-12);
}
#[test]
fn ensemble_point_forecast_averages_models() {
let n_test = 3;
let n_models = 4;
let test = vec![
1.0, 2.0, 3.0, 4.0, 10.0, 20.0, 30.0, 40.0, 5.0, 5.0, 5.0, 5.0, ];
let enbpi = EnbPiPredictor::new(0.90);
let pf = enbpi.ensemble_point_forecast(&test, n_models).unwrap();
assert_eq!(pf, vec![2.5, 25.0, 5.0]);
let _ = n_test;
}
#[test]
fn online_update_slides_window() {
let actuals: Vec<f64> = (0..50).map(|i| i as f64).collect();
let preds: Vec<f64> = actuals.iter().map(|&a| a + 0.5).collect();
let mask = vec![false; 50];
let enbpi = EnbPiPredictor::new(0.90);
let mut result = enbpi.fit(&preds, &mask, &actuals, 1).unwrap();
assert!((result.quantile_value() - 0.5).abs() < 1e-12);
let new_actuals = vec![100.0, 101.0, 102.0];
let new_preds = vec![101.5, 102.5, 103.5];
result.update(&new_preds, &new_actuals, 30).unwrap();
assert!(result.quantile_value() >= 0.5);
assert!(result.residuals().len() <= 30);
}
#[test]
fn empty_data_errors() {
let enbpi = EnbPiPredictor::new(0.90);
let err = enbpi.fit(&[], &[], &[], 1).unwrap_err();
assert!(matches!(err, ForecastError::EmptyData));
}
#[test]
fn dimension_mismatch_errors() {
let enbpi = EnbPiPredictor::new(0.90);
let err = enbpi
.fit(&[1.0, 2.0, 3.0, 4.0, 5.0], &[false; 6], &[1.0, 2.0], 3)
.unwrap_err();
assert!(matches!(err, ForecastError::DimensionMismatch { .. }));
}
}