use crate::core::TimeSeries;
use crate::error::{ForecastError, Result};
use crate::models::ensemble::Ensemble;
use crate::models::traits::{Forecaster, ModelRegistry};
use crate::utils::cross_validation::CvFoldGenerator;
use crate::utils::metrics::{calculate_metrics, mae, rmse};
use std::fmt;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
#[derive(Debug, Clone)]
pub struct ModelResult {
pub name: String,
pub mae: f64,
pub rmse: f64,
pub mape: f64,
pub fit_succeeded: bool,
}
#[derive(Debug, Clone)]
pub struct ModelComparison {
pub results: Vec<ModelResult>,
}
impl fmt::Display for ModelComparison {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(
f,
"{:<25} {:>10} {:>10} {:>10} Status",
"Model", "MAE", "RMSE", "MAPE"
)?;
writeln!(f, "{}", "-".repeat(70))?;
for r in &self.results {
let status = if r.fit_succeeded { "OK" } else { "FAILED" };
if r.fit_succeeded {
writeln!(
f,
"{:<25} {:>10.4} {:>10.4} {:>9.2}% {}",
r.name, r.mae, r.rmse, r.mape, status
)?;
} else {
writeln!(
f,
"{:<25} {:>10} {:>10} {:>10} {}",
r.name, "N/A", "N/A", "N/A", status
)?;
}
}
Ok(())
}
}
pub fn fit_all_and_compare(
registry: &ModelRegistry,
ts: &TimeSeries,
holdout: usize,
) -> ModelComparison {
if registry.is_empty() {
return ModelComparison {
results: Vec::new(),
};
}
let n = ts.len();
if holdout >= n || holdout == 0 {
let mut results: Vec<ModelResult> = registry
.iter()
.map(|spec| ModelResult {
name: spec.name.to_string(),
mae: f64::NAN,
rmse: f64::NAN,
mape: f64::NAN,
fit_succeeded: false,
})
.collect();
results.sort_by(cmp_model_result);
return ModelComparison { results };
}
let split = n - holdout;
let train = match ts.slice(0, split) {
Ok(t) => t,
Err(_) => {
let mut results: Vec<ModelResult> = registry
.iter()
.map(|spec| ModelResult {
name: spec.name.to_string(),
mae: f64::NAN,
rmse: f64::NAN,
mape: f64::NAN,
fit_succeeded: false,
})
.collect();
results.sort_by(cmp_model_result);
return ModelComparison { results };
}
};
let actual = &ts.primary_values()[split..];
let specs: Vec<_> = registry.iter().collect();
let evaluate = |spec: &&crate::models::traits::ModelSpec| -> ModelResult {
let mut model = spec.create();
let name = spec.name.to_string();
match model.fit(&train) {
Ok(()) => match model.predict(holdout) {
Ok(forecast) => {
let predicted = forecast.primary();
let m = calculate_metrics(actual, predicted, None);
match m {
Ok(metrics) => ModelResult {
name,
mae: metrics.mae,
rmse: metrics.rmse,
mape: metrics.mape.unwrap_or(f64::NAN),
fit_succeeded: true,
},
Err(_) => ModelResult {
name,
mae: mae(actual, predicted),
rmse: rmse(actual, predicted),
mape: f64::NAN,
fit_succeeded: true,
},
}
}
Err(_) => ModelResult {
name,
mae: f64::NAN,
rmse: f64::NAN,
mape: f64::NAN,
fit_succeeded: false,
},
},
Err(_) => ModelResult {
name,
mae: f64::NAN,
rmse: f64::NAN,
mape: f64::NAN,
fit_succeeded: false,
},
}
};
#[cfg(feature = "parallel")]
let mut results: Vec<ModelResult> = specs.par_iter().map(evaluate).collect();
#[cfg(not(feature = "parallel"))]
let mut results: Vec<ModelResult> = specs.iter().map(evaluate).collect();
results.sort_by(cmp_model_result);
ModelComparison { results }
}
fn cmp_model_result(a: &ModelResult, b: &ModelResult) -> std::cmp::Ordering {
match (a.fit_succeeded, b.fit_succeeded) {
(true, false) => std::cmp::Ordering::Less,
(false, true) => std::cmp::Ordering::Greater,
(false, false) => a.name.cmp(&b.name),
(true, true) => a
.mae
.partial_cmp(&b.mae)
.unwrap_or(std::cmp::Ordering::Equal),
}
}
#[derive(Debug, Clone)]
pub struct CVModelResult {
pub name: String,
pub mean_mae: f64,
pub mean_rmse: f64,
pub std_mae: f64,
}
#[derive(Debug, Clone)]
pub struct CVComparison {
pub results: Vec<CVModelResult>,
}
impl fmt::Display for CVComparison {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
writeln!(
f,
"{:<25} {:>10} {:>10} {:>10}",
"Model", "Mean MAE", "Mean RMSE", "Std MAE"
)?;
writeln!(f, "{}", "-".repeat(58))?;
for r in &self.results {
if r.mean_mae.is_finite() {
writeln!(
f,
"{:<25} {:>10.4} {:>10.4} {:>10.4}",
r.name, r.mean_mae, r.mean_rmse, r.std_mae
)?;
} else {
writeln!(
f,
"{:<25} {:>10} {:>10} {:>10}",
r.name, "N/A", "N/A", "N/A"
)?;
}
}
Ok(())
}
}
pub fn cross_validate_all(
registry: &ModelRegistry,
ts: &TimeSeries,
n_folds: usize,
horizon: usize,
) -> CVComparison {
if registry.is_empty() {
return CVComparison {
results: Vec::new(),
};
}
let series_len = ts.len();
let (initial_window, step_size) = if n_folds == 0 || horizon == 0 || series_len < horizon + 2 {
(series_len, 1) } else {
let needed_for_folds = n_folds * horizon;
if needed_for_folds >= series_len {
let iw = series_len.saturating_sub(n_folds + horizon - 1).max(2);
(iw, 1)
} else {
let iw = series_len.saturating_sub(needed_for_folds).max(2);
(iw, horizon)
}
};
let generator = CvFoldGenerator::new()
.initial_window(initial_window)
.horizon(horizon)
.step_size(step_size);
let folds = generator.generate(series_len);
let specs: Vec<_> = registry.iter().collect();
let evaluate = |spec: &&crate::models::traits::ModelSpec| -> CVModelResult {
let name = spec.name.to_string();
if folds.is_empty() {
return CVModelResult {
name,
mean_mae: f64::NAN,
mean_rmse: f64::NAN,
std_mae: f64::NAN,
};
}
let mut fold_maes = Vec::with_capacity(folds.len());
let mut fold_rmses = Vec::with_capacity(folds.len());
for fold in &folds {
let train = match ts.slice(fold.train_start, fold.train_end) {
Ok(t) => t,
Err(_) => continue,
};
let mut model = spec.create();
if model.fit(&train).is_err() {
continue;
}
let forecast = match model.predict(fold.test_size()) {
Ok(f) => f,
Err(_) => continue,
};
let actual_slice = &ts.primary_values()[fold.test_start..fold.test_end];
let predicted = forecast.primary();
let fold_mae = mae(actual_slice, predicted);
let fold_rmse = rmse(actual_slice, predicted);
if fold_mae.is_finite() && fold_rmse.is_finite() {
fold_maes.push(fold_mae);
fold_rmses.push(fold_rmse);
}
}
if fold_maes.is_empty() {
return CVModelResult {
name,
mean_mae: f64::NAN,
mean_rmse: f64::NAN,
std_mae: f64::NAN,
};
}
let n = fold_maes.len() as f64;
let mean_mae = fold_maes.iter().sum::<f64>() / n;
let mean_rmse = fold_rmses.iter().sum::<f64>() / n;
let std_mae = if fold_maes.len() > 1 {
let variance = fold_maes
.iter()
.map(|v| (v - mean_mae).powi(2))
.sum::<f64>()
/ (n - 1.0);
variance.sqrt()
} else {
0.0
};
CVModelResult {
name,
mean_mae,
mean_rmse,
std_mae,
}
};
#[cfg(feature = "parallel")]
let mut results: Vec<CVModelResult> = specs.par_iter().map(evaluate).collect();
#[cfg(not(feature = "parallel"))]
let mut results: Vec<CVModelResult> = specs.iter().map(evaluate).collect();
results.sort_by(|a, b| {
let fa = a.mean_mae.is_finite();
let fb = b.mean_mae.is_finite();
match (fa, fb) {
(true, false) => std::cmp::Ordering::Less,
(false, true) => std::cmp::Ordering::Greater,
(false, false) => a.name.cmp(&b.name),
(true, true) => a
.mean_mae
.partial_cmp(&b.mean_mae)
.unwrap_or(std::cmp::Ordering::Equal),
}
});
CVComparison { results }
}
pub fn ensemble_best_k(
registry: &ModelRegistry,
ts: &TimeSeries,
k: usize,
holdout: usize,
) -> Result<Box<dyn Forecaster>> {
if registry.is_empty() || k == 0 {
return Err(ForecastError::EmptyData);
}
let comparison = fit_all_and_compare(registry, ts, holdout);
let top_names: Vec<String> = comparison
.results
.iter()
.filter(|r| r.fit_succeeded)
.take(k)
.map(|r| r.name.clone())
.collect();
if top_names.is_empty() {
return Err(ForecastError::EmptyData);
}
let models: Vec<Box<dyn Forecaster>> = registry
.iter()
.filter(|spec| top_names.contains(&spec.name.to_string()))
.map(|spec| spec.create())
.collect();
if models.is_empty() {
return Err(ForecastError::EmptyData);
}
let mut ensemble = Ensemble::new(models);
ensemble.fit(ts)?;
Ok(Box::new(ensemble))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::TimeSeries;
use crate::models::baseline::{Naive, RandomWalkWithDrift, WindowAverage};
use crate::models::{ModelRegistry, ModelSpec};
use chrono::{TimeZone, Utc};
fn make_timestamps(n: usize) -> Vec<chrono::DateTime<Utc>> {
(0..n)
.map(|i| {
Utc.with_ymd_and_hms(2020, 1, 1, 0, 0, 0).unwrap()
+ chrono::Duration::days(i as i64)
})
.collect()
}
fn make_test_series(n: usize) -> TimeSeries {
let timestamps = make_timestamps(n);
let values: Vec<f64> = (1..=n).map(|i| i as f64).collect();
TimeSeries::univariate(timestamps, values).unwrap()
}
fn make_registry() -> ModelRegistry {
let mut registry = ModelRegistry::new();
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
registry.register(ModelSpec::new(
"RWD",
|| Box::new(RandomWalkWithDrift::new()),
true,
));
registry.register(ModelSpec::with_period(
"WindowAvg5",
|p| Box::new(WindowAverage::new(p)),
5,
false,
));
registry
}
#[test]
fn test_fit_all_and_compare_basic() {
let registry = make_registry();
let ts = make_test_series(50);
let comparison = fit_all_and_compare(®istry, &ts, 10);
assert_eq!(comparison.results.len(), 3);
for r in &comparison.results {
assert!(r.fit_succeeded, "model {} should succeed", r.name);
assert!(r.mae.is_finite(), "model {} MAE should be finite", r.name);
assert!(r.rmse.is_finite(), "model {} RMSE should be finite", r.name);
}
for w in comparison.results.windows(2) {
assert!(
w[0].mae <= w[1].mae,
"results should be sorted by MAE: {} ({}) <= {} ({})",
w[0].name,
w[0].mae,
w[1].name,
w[1].mae
);
}
}
#[test]
fn test_fit_all_and_compare_display() {
let registry = make_registry();
let ts = make_test_series(50);
let comparison = fit_all_and_compare(®istry, &ts, 10);
let display = format!("{}", comparison);
assert!(display.contains("Model"));
assert!(display.contains("MAE"));
assert!(display.contains("RMSE"));
}
#[test]
fn test_fit_all_and_compare_empty_registry() {
let registry = ModelRegistry::new();
let ts = make_test_series(50);
let comparison = fit_all_and_compare(®istry, &ts, 10);
assert!(comparison.results.is_empty());
}
#[test]
fn test_fit_all_and_compare_bad_holdout() {
let registry = make_registry();
let ts = make_test_series(50);
let comparison = fit_all_and_compare(®istry, &ts, 50);
assert_eq!(comparison.results.len(), 3);
for r in &comparison.results {
assert!(!r.fit_succeeded);
}
}
#[test]
fn test_cross_validate_all_basic() {
let registry = make_registry();
let ts = make_test_series(60);
let cv = cross_validate_all(®istry, &ts, 3, 5);
assert_eq!(cv.results.len(), 3);
for r in &cv.results {
assert!(
r.mean_mae.is_finite(),
"model {} mean_mae should be finite",
r.name
);
assert!(
r.mean_rmse.is_finite(),
"model {} mean_rmse should be finite",
r.name
);
assert!(
r.std_mae.is_finite(),
"model {} std_mae should be finite",
r.name
);
}
for w in cv.results.windows(2) {
assert!(
w[0].mean_mae <= w[1].mean_mae,
"results should be sorted by mean MAE"
);
}
}
#[test]
fn test_cross_validate_all_display() {
let registry = make_registry();
let ts = make_test_series(60);
let cv = cross_validate_all(®istry, &ts, 3, 5);
let display = format!("{}", cv);
assert!(display.contains("Model"));
assert!(display.contains("Mean MAE"));
}
#[test]
fn test_cross_validate_all_empty_registry() {
let registry = ModelRegistry::new();
let ts = make_test_series(60);
let cv = cross_validate_all(®istry, &ts, 3, 5);
assert!(cv.results.is_empty());
}
#[test]
fn test_ensemble_best_k_basic() {
let registry = make_registry();
let ts = make_test_series(50);
let ensemble = ensemble_best_k(®istry, &ts, 2, 10);
assert!(ensemble.is_ok(), "ensemble_best_k should succeed");
let model = ensemble.unwrap();
assert!(model.is_fitted());
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
for &v in forecast.primary() {
assert!(v.is_finite(), "forecast value should be finite");
}
}
#[test]
fn test_ensemble_best_k_single() {
let registry = make_registry();
let ts = make_test_series(50);
let ensemble = ensemble_best_k(®istry, &ts, 1, 10);
assert!(ensemble.is_ok());
let model = ensemble.unwrap();
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn test_ensemble_best_k_empty_registry() {
let registry = ModelRegistry::new();
let ts = make_test_series(50);
let result = ensemble_best_k(®istry, &ts, 2, 10);
assert!(result.is_err());
}
#[test]
fn test_ensemble_best_k_k_zero() {
let registry = make_registry();
let ts = make_test_series(50);
let result = ensemble_best_k(®istry, &ts, 0, 10);
assert!(result.is_err());
}
#[test]
fn test_ensemble_best_k_k_larger_than_registry() {
let registry = make_registry();
let ts = make_test_series(50);
let ensemble = ensemble_best_k(®istry, &ts, 10, 10);
assert!(ensemble.is_ok());
let model = ensemble.unwrap();
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
}