use serde::{Deserialize, Serialize};
use thiserror::Error;
#[derive(Debug, Error)]
pub enum ForecastError {
#[error("no ensemble members added")]
NoMembers,
#[error("member {id} has length {got}, expected {expected}")]
LengthMismatch {
id: usize,
got: usize,
expected: usize,
},
#[error("weight vector length {weights} != n_members {members}")]
WeightLengthMismatch { weights: usize, members: usize },
#[error("quantile {q} must be in [0, 1]")]
InvalidQuantile { q: f64 },
#[error("numerical error: {0}")]
NumericalError(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EnsembleV2Config {
pub n_members: usize,
pub horizon_hours: usize,
pub aggregation_method: AggregationMethod,
pub uncertainty_quantification: UqMethod,
}
impl Default for EnsembleV2Config {
fn default() -> Self {
Self {
n_members: 10,
horizon_hours: 24,
aggregation_method: AggregationMethod::SimpleMean,
uncertainty_quantification: UqMethod::Spread,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum AggregationMethod {
SimpleMean,
WeightedMean {
weights: Vec<f64>,
},
BestN {
n: usize,
},
BayesianModelAveraging,
SuperEnsemble,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum UqMethod {
Spread,
QuantileRegression {
quantiles: Vec<f64>,
},
ConformalPrediction,
BayesianBootstrap,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EnsembleMemberV2 {
pub id: usize,
pub method: String,
pub forecasts: Vec<f64>,
pub historical_rmse: f64,
}
impl EnsembleMemberV2 {
pub fn new(id: usize, method: impl Into<String>, forecasts: Vec<f64>) -> Self {
Self {
id,
method: method.into(),
historical_rmse: 1.0, forecasts,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EnsembleForecast {
pub point_forecast: Vec<f64>,
pub std_dev: Vec<f64>,
pub quantile_10: Vec<f64>,
pub quantile_25: Vec<f64>,
pub quantile_75: Vec<f64>,
pub quantile_90: Vec<f64>,
pub prediction_interval_90: Vec<(f64, f64)>,
pub member_weights: Vec<f64>,
pub skill_scores: Vec<f64>,
}
pub struct EnsembleForecaster {
config: EnsembleV2Config,
members: Vec<EnsembleMemberV2>,
observations: Vec<f64>,
}
impl EnsembleForecaster {
pub fn new(config: EnsembleV2Config) -> Self {
Self {
config,
members: Vec::new(),
observations: Vec::new(),
}
}
pub fn add_member(&mut self, member: EnsembleMemberV2) {
self.members.push(member);
}
pub fn set_observations(&mut self, obs: Vec<f64>) {
self.observations = obs;
}
fn bma_weights(&self) -> Vec<f64> {
if self.members.is_empty() {
return Vec::new();
}
let rmse_vals: Vec<f64> = self.members.iter().map(|m| m.historical_rmse).collect();
let mean_rmse = rmse_vals.iter().sum::<f64>() / rmse_vals.len() as f64;
let sigma2 = if mean_rmse > 1e-12 {
mean_rmse * mean_rmse
} else {
1.0
};
let raw: Vec<f64> = rmse_vals
.iter()
.map(|&r| (-(r * r) / (2.0 * sigma2)).exp())
.collect();
let total: f64 = raw.iter().sum();
if total < 1e-12 {
vec![1.0 / self.members.len() as f64; self.members.len()]
} else {
raw.iter().map(|w| w / total).collect()
}
}
fn empirical_quantile(&self, step: usize, quantile: f64) -> f64 {
let mut values: Vec<f64> = self
.members
.iter()
.filter_map(|m| m.forecasts.get(step).copied())
.collect();
if values.is_empty() {
return 0.0;
}
values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let idx_f = quantile * (values.len() - 1) as f64;
let lo = idx_f.floor() as usize;
let hi = (lo + 1).min(values.len() - 1);
let frac = idx_f - lo as f64;
values[lo] * (1.0 - frac) + values[hi] * frac
}
fn conformal_intervals(&self, quantile: f64) -> Vec<(f64, f64)> {
let n_steps = self.members.first().map(|m| m.forecasts.len()).unwrap_or(0);
(0..n_steps)
.map(|t| {
let mean = self
.members
.iter()
.filter_map(|m| m.forecasts.get(t).copied())
.sum::<f64>()
/ self.members.len().max(1) as f64;
let mut scores: Vec<f64> = self
.members
.iter()
.filter_map(|m| m.forecasts.get(t).map(|&f| (f - mean).abs()))
.collect();
scores.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let idx_f = quantile * (scores.len().saturating_sub(1)) as f64;
let lo = idx_f.floor() as usize;
let hi = (lo + 1).min(scores.len().saturating_sub(1));
let frac = idx_f - lo as f64;
let margin = if scores.is_empty() {
0.0
} else {
scores[lo] * (1.0 - frac) + scores[hi] * frac
};
(mean - margin, mean + margin)
})
.collect()
}
fn crps_skill(&self, member: &EnsembleMemberV2) -> f64 {
if self.observations.is_empty() || member.forecasts.is_empty() {
return member.historical_rmse; }
let n = member.forecasts.len().min(self.observations.len());
if n == 0 {
return f64::INFINITY;
}
let mae: f64 = member
.forecasts
.iter()
.zip(self.observations.iter())
.take(n)
.map(|(&f, &o)| (f - o).abs())
.sum::<f64>()
/ n as f64;
let spread: f64 = if self.members.len() > 1 {
let total_pairs = (self.members.len() * (self.members.len() - 1)) as f64;
let pair_sum: f64 = self
.members
.iter()
.flat_map(|m1| {
self.members.iter().map(move |m2| {
m1.forecasts
.iter()
.zip(m2.forecasts.iter())
.take(n)
.map(|(a, b)| (a - b).abs())
.sum::<f64>()
/ n as f64
})
})
.sum();
pair_sum / total_pairs
} else {
0.0
};
mae - 0.5 * spread
}
pub fn forecast(&self) -> Result<EnsembleForecast, ForecastError> {
if self.members.is_empty() {
return Err(ForecastError::NoMembers);
}
let n_steps = self.members[0].forecasts.len();
for m in &self.members {
if m.forecasts.len() != n_steps {
return Err(ForecastError::LengthMismatch {
id: m.id,
got: m.forecasts.len(),
expected: n_steps,
});
}
}
let n_members = self.members.len();
let weights = match &self.config.aggregation_method {
AggregationMethod::SimpleMean => vec![1.0 / n_members as f64; n_members],
AggregationMethod::WeightedMean { weights } => {
if weights.len() != n_members {
return Err(ForecastError::WeightLengthMismatch {
weights: weights.len(),
members: n_members,
});
}
let total: f64 = weights.iter().sum();
if total < 1e-12 {
return Err(ForecastError::NumericalError("weights sum to zero".into()));
}
weights.iter().map(|w| w / total).collect()
}
AggregationMethod::BestN { n } => {
let mut indexed: Vec<(usize, f64)> = self
.members
.iter()
.map(|m| (m.id, m.historical_rmse))
.enumerate()
.map(|(i, (_, rmse))| (i, rmse))
.collect();
indexed.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
let best_n = (*n).min(n_members);
let best_set: std::collections::HashSet<usize> =
indexed.iter().take(best_n).map(|(i, _)| *i).collect();
let w = 1.0 / best_n.max(1) as f64;
(0..n_members)
.map(|i| if best_set.contains(&i) { w } else { 0.0 })
.collect()
}
AggregationMethod::BayesianModelAveraging | AggregationMethod::SuperEnsemble => {
self.bma_weights()
}
};
let point_forecast: Vec<f64> = (0..n_steps)
.map(|t| {
weights
.iter()
.zip(self.members.iter())
.map(|(&w, m)| w * m.forecasts[t])
.sum()
})
.collect();
let std_dev: Vec<f64> = (0..n_steps)
.map(|t| {
let mean = point_forecast[t];
let var: f64 = self
.members
.iter()
.map(|m| (m.forecasts[t] - mean).powi(2))
.sum::<f64>()
/ n_members as f64;
var.sqrt()
})
.collect();
let quantile_10: Vec<f64> = (0..n_steps)
.map(|t| self.empirical_quantile(t, 0.10))
.collect();
let quantile_25: Vec<f64> = (0..n_steps)
.map(|t| self.empirical_quantile(t, 0.25))
.collect();
let quantile_75: Vec<f64> = (0..n_steps)
.map(|t| self.empirical_quantile(t, 0.75))
.collect();
let quantile_90: Vec<f64> = (0..n_steps)
.map(|t| self.empirical_quantile(t, 0.90))
.collect();
let prediction_interval_90 = match &self.config.uncertainty_quantification {
UqMethod::ConformalPrediction => self.conformal_intervals(0.90),
UqMethod::Spread => (0..n_steps)
.map(|t| {
let z = 1.645; (
point_forecast[t] - z * std_dev[t],
point_forecast[t] + z * std_dev[t],
)
})
.collect(),
UqMethod::QuantileRegression { .. } | UqMethod::BayesianBootstrap => (0..n_steps)
.map(|t| (quantile_10[t], quantile_90[t]))
.collect(),
};
let skill_scores: Vec<f64> = self.members.iter().map(|m| self.crps_skill(m)).collect();
Ok(EnsembleForecast {
point_forecast,
std_dev,
quantile_10,
quantile_25,
quantile_75,
quantile_90,
prediction_interval_90,
member_weights: weights,
skill_scores,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_member(id: usize, forecasts: Vec<f64>, rmse: f64) -> EnsembleMemberV2 {
EnsembleMemberV2 {
id,
method: format!("M{id}"),
forecasts,
historical_rmse: rmse,
}
}
#[test]
fn test_simple_mean_is_average() {
let config = EnsembleV2Config {
n_members: 3,
horizon_hours: 3,
aggregation_method: AggregationMethod::SimpleMean,
uncertainty_quantification: UqMethod::Spread,
};
let mut fc = EnsembleForecaster::new(config);
fc.add_member(make_member(0, vec![10.0, 20.0, 30.0], 1.0));
fc.add_member(make_member(1, vec![20.0, 30.0, 40.0], 1.0));
fc.add_member(make_member(2, vec![30.0, 40.0, 50.0], 1.0));
let result = fc.forecast().expect("must succeed");
let expected = [20.0, 30.0, 40.0];
for (computed, &expected) in result.point_forecast.iter().zip(expected.iter()) {
assert!(
(computed - expected).abs() < 1e-9,
"simple mean mismatch: got {computed:.3}, expected {expected:.3}"
);
}
}
#[test]
fn test_weighted_mean_higher_weight_lower_rmse() {
let weights = vec![0.8, 0.2]; let config = EnsembleV2Config {
n_members: 2,
horizon_hours: 2,
aggregation_method: AggregationMethod::WeightedMean {
weights: weights.clone(),
},
uncertainty_quantification: UqMethod::Spread,
};
let mut fc = EnsembleForecaster::new(config);
fc.add_member(make_member(0, vec![100.0, 100.0], 0.5)); fc.add_member(make_member(1, vec![0.0, 0.0], 5.0));
let result = fc.forecast().expect("must succeed");
for &pf in &result.point_forecast {
assert!(
(pf - 80.0).abs() < 1e-9,
"weighted mean must be 80, got {pf:.3}"
);
}
assert!(
result.member_weights[0] > result.member_weights[1],
"higher explicit weight must dominate"
);
}
#[test]
fn test_bma_weights_sum_to_one() {
let config = EnsembleV2Config {
n_members: 5,
horizon_hours: 4,
aggregation_method: AggregationMethod::BayesianModelAveraging,
uncertainty_quantification: UqMethod::Spread,
};
let mut fc = EnsembleForecaster::new(config);
let rmse_vals = [0.5, 1.0, 2.0, 0.3, 1.5];
for (i, &r) in rmse_vals.iter().enumerate() {
fc.add_member(make_member(i, vec![10.0; 4], r));
}
let bma_w = fc.bma_weights();
let total: f64 = bma_w.iter().sum();
assert!(
(total - 1.0).abs() < 1e-9,
"BMA weights must sum to 1.0, got {total:.10}"
);
let result = fc.forecast().expect("must succeed");
let w_total: f64 = result.member_weights.iter().sum();
assert!((w_total - 1.0).abs() < 1e-9, "forecast weights sum to 1.0");
}
#[test]
fn test_prediction_interval_coverage() {
let config = EnsembleV2Config {
n_members: 3,
horizon_hours: 5,
aggregation_method: AggregationMethod::SimpleMean,
uncertainty_quantification: UqMethod::ConformalPrediction,
};
let mut fc = EnsembleForecaster::new(config);
fc.add_member(make_member(0, vec![10.0; 5], 1.0));
fc.add_member(make_member(1, vec![20.0; 5], 1.0));
fc.add_member(make_member(2, vec![30.0; 5], 1.0));
let observation = 20.0; let result = fc.forecast().expect("must succeed");
let n_covered = result
.prediction_interval_90
.iter()
.filter(|(lo, hi)| observation >= *lo && observation <= *hi)
.count();
assert!(
n_covered > 0,
"prediction interval must contain the observation"
);
}
#[test]
fn test_crps_lower_for_accurate_member() {
let config = EnsembleV2Config {
n_members: 2,
horizon_hours: 5,
aggregation_method: AggregationMethod::SimpleMean,
uncertainty_quantification: UqMethod::Spread,
};
let obs = vec![50.0; 5];
let mut fc = EnsembleForecaster::new(config);
fc.add_member(make_member(0, vec![51.0; 5], 1.0));
fc.add_member(make_member(1, vec![80.0; 5], 30.0));
fc.set_observations(obs);
let crps_accurate = fc.crps_skill(&fc.members[0].clone());
let crps_inaccurate = fc.crps_skill(&fc.members[1].clone());
assert!(
crps_accurate < crps_inaccurate,
"CRPS must be lower for accurate member ({crps_accurate:.4} < {crps_inaccurate:.4})"
);
}
#[test]
fn test_best_n_selection() {
let config = EnsembleV2Config {
n_members: 4,
horizon_hours: 2,
aggregation_method: AggregationMethod::BestN { n: 2 },
uncertainty_quantification: UqMethod::Spread,
};
let mut fc = EnsembleForecaster::new(config);
fc.add_member(make_member(0, vec![100.0; 2], 5.0)); fc.add_member(make_member(1, vec![100.0; 2], 1.0)); fc.add_member(make_member(2, vec![100.0; 2], 2.0)); fc.add_member(make_member(3, vec![100.0; 2], 4.0));
let result = fc.forecast().expect("must succeed");
assert!(
result.member_weights[0] < 1e-9,
"worst member must have weight 0"
);
assert!(
result.member_weights[1] > 0.0,
"best member must have weight > 0"
);
assert!(
result.member_weights[2] > 0.0,
"2nd best must have weight > 0"
);
assert!(
result.member_weights[3] < 1e-9,
"4th member must have weight 0"
);
}
#[test]
fn test_no_members_error() {
let config = EnsembleV2Config::default();
let fc = EnsembleForecaster::new(config);
let result = fc.forecast();
assert!(
matches!(result, Err(ForecastError::NoMembers)),
"must return NoMembers error"
);
}
}