use crate::core::{Forecast, TimeSeries};
use crate::error::Result;
use crate::models::traits::{Forecaster, ModelRegistry};
use crate::utils::metrics::{calculate_metrics, AccuracyMetrics};
#[cfg(feature = "parallel")]
use rayon::prelude::*;
#[derive(Debug)]
pub struct BatchResult<M> {
pub model: M,
pub series_index: usize,
pub error: Option<String>,
}
#[derive(Debug)]
pub struct FitPredictResult {
pub forecast: Option<Forecast>,
pub model_name: String,
pub series_index: usize,
pub error: Option<String>,
}
#[derive(Debug)]
pub struct RegistryResult {
pub model_name: String,
pub forecast: Option<Forecast>,
pub metrics: Option<AccuracyMetrics>,
pub error: Option<String>,
}
pub fn fit_many<F, M>(factory: F, series: &[&TimeSeries]) -> Vec<Result<M>>
where
F: Fn() -> M + Send + Sync,
M: Forecaster + Send,
{
#[cfg(feature = "parallel")]
{
series
.par_iter()
.map(|ts| {
let mut model = factory();
model.fit(ts)?;
Ok(model)
})
.collect()
}
#[cfg(not(feature = "parallel"))]
{
series
.iter()
.map(|ts| {
let mut model = factory();
model.fit(ts)?;
Ok(model)
})
.collect()
}
}
pub fn predict_many(models: &[impl Forecaster], horizon: usize) -> Vec<Result<Forecast>> {
models.iter().map(|m| m.predict(horizon)).collect()
}
pub fn fit_predict_many<F, M>(
factory: F,
series: &[&TimeSeries],
horizon: usize,
) -> Vec<FitPredictResult>
where
F: Fn() -> M + Send + Sync,
M: Forecaster + Send,
{
let process = |idx: usize, ts: &&TimeSeries| -> FitPredictResult {
let mut model = factory();
let model_name = model.name().to_string();
match model.fit(ts) {
Ok(()) => match model.predict(horizon) {
Ok(forecast) => FitPredictResult {
forecast: Some(forecast),
model_name,
series_index: idx,
error: None,
},
Err(e) => FitPredictResult {
forecast: None,
model_name,
series_index: idx,
error: Some(format!("predict failed: {e}")),
},
},
Err(e) => FitPredictResult {
forecast: None,
model_name,
series_index: idx,
error: Some(format!("fit failed: {e}")),
},
}
};
#[cfg(feature = "parallel")]
{
series
.par_iter()
.enumerate()
.map(|(idx, ts)| process(idx, ts))
.collect()
}
#[cfg(not(feature = "parallel"))]
{
series
.iter()
.enumerate()
.map(|(idx, ts)| process(idx, ts))
.collect()
}
}
pub fn fit_registry(registry: &ModelRegistry, series: &TimeSeries) -> Vec<RegistryResult> {
let actual = series.primary_values();
let specs: Vec<_> = registry.iter().collect();
let process = |spec: &&crate::models::traits::ModelSpec| -> RegistryResult {
let model_name = spec.name.to_string();
let mut model = spec.create();
match model.fit(series) {
Ok(()) => {
let metrics = model.fitted_values().and_then(|fitted| {
let (a, p): (Vec<f64>, Vec<f64>) = actual
.iter()
.zip(fitted.iter())
.filter(|(a, f)| a.is_finite() && f.is_finite())
.map(|(&a, &f)| (a, f))
.unzip();
if a.is_empty() {
return None;
}
calculate_metrics(&a, &p, None).ok()
});
match model.predict(1) {
Ok(forecast) => RegistryResult {
model_name,
forecast: Some(forecast),
metrics,
error: None,
},
Err(e) => RegistryResult {
model_name,
forecast: None,
metrics,
error: Some(format!("predict failed: {e}")),
},
}
}
Err(e) => RegistryResult {
model_name,
forecast: None,
metrics: None,
error: Some(format!("fit failed: {e}")),
},
}
};
#[cfg(feature = "parallel")]
{
specs.par_iter().map(process).collect()
}
#[cfg(not(feature = "parallel"))]
{
specs.iter().map(process).collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::TimeSeries;
use crate::models::baseline::{Naive, RandomWalkWithDrift, SeasonalNaive, 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_seasonal_series(n: usize, period: usize) -> TimeSeries {
let timestamps = make_timestamps(n);
let values: Vec<f64> = (0..n)
.map(|i| 10.0 + 5.0 * ((2.0 * std::f64::consts::PI * i as f64) / period as f64).sin())
.collect();
TimeSeries::univariate(timestamps, values).unwrap()
}
#[test]
fn test_fit_many_single_series() {
let ts = make_test_series(30);
let results = fit_many(Naive::new, &[&ts]);
assert_eq!(results.len(), 1);
assert!(results[0].is_ok());
assert!(results[0].as_ref().unwrap().is_fitted());
}
#[test]
fn test_fit_many_multiple_series() {
let ts1 = make_test_series(30);
let ts2 = make_test_series(50);
let ts3 = make_seasonal_series(40, 7);
let results = fit_many(Naive::new, &[&ts1, &ts2, &ts3]);
assert_eq!(results.len(), 3);
for r in &results {
assert!(r.is_ok());
}
}
#[test]
fn test_fit_many_preserves_order() {
let ts1 = make_test_series(20);
let ts2 = make_test_series(40);
let results = fit_many(Naive::new, &[&ts1, &ts2]);
let m1 = results[0].as_ref().unwrap();
let m2 = results[1].as_ref().unwrap();
assert_eq!(m1.residuals().unwrap().len(), 20);
assert_eq!(m2.residuals().unwrap().len(), 40);
}
#[test]
fn test_fit_many_empty_series_errors() {
let ts = make_test_series(30);
let short = make_test_series(5);
let results = fit_many(|| SeasonalNaive::new(12), &[&ts, &short]);
assert!(results[0].is_ok());
assert!(results[1].is_err());
}
#[test]
fn test_fit_many_empty_slice() {
let results = fit_many(Naive::new, &[] as &[&TimeSeries]);
assert!(results.is_empty());
}
#[test]
fn test_predict_many_basic() {
let ts = make_test_series(30);
let fitted: Vec<Naive> = fit_many(Naive::new, &[&ts])
.into_iter()
.map(|r| r.unwrap())
.collect();
let forecasts = predict_many(&fitted, 5);
assert_eq!(forecasts.len(), 1);
let fc = forecasts[0].as_ref().unwrap();
assert_eq!(fc.horizon(), 5);
}
#[test]
fn test_predict_many_multiple_models() {
let ts1 = make_test_series(30);
let ts2 = make_test_series(40);
let fitted: Vec<Naive> = fit_many(Naive::new, &[&ts1, &ts2])
.into_iter()
.map(|r| r.unwrap())
.collect();
let forecasts = predict_many(&fitted, 3);
assert_eq!(forecasts.len(), 2);
for fc in &forecasts {
assert!(fc.is_ok());
assert_eq!(fc.as_ref().unwrap().horizon(), 3);
}
}
#[test]
fn test_predict_many_unfitted_model_errors() {
let unfitted = Naive::new();
let forecasts = predict_many(&[unfitted], 5);
assert_eq!(forecasts.len(), 1);
assert!(forecasts[0].is_err());
}
#[test]
fn test_fit_predict_many_basic() {
let ts = make_test_series(30);
let results = fit_predict_many(Naive::new, &[&ts], 5);
assert_eq!(results.len(), 1);
let r = &results[0];
assert!(r.forecast.is_some());
assert!(r.error.is_none());
assert_eq!(r.series_index, 0);
assert_eq!(r.model_name, "Naive");
assert_eq!(r.forecast.as_ref().unwrap().horizon(), 5);
}
#[test]
fn test_fit_predict_many_multiple() {
let ts1 = make_test_series(30);
let ts2 = make_test_series(50);
let results = fit_predict_many(Naive::new, &[&ts1, &ts2], 3);
assert_eq!(results.len(), 2);
for (idx, r) in results.iter().enumerate() {
assert!(r.forecast.is_some(), "series {idx} should succeed");
assert!(r.error.is_none());
assert_eq!(r.series_index, idx);
}
}
#[test]
fn test_fit_predict_many_records_fit_error() {
let short = make_test_series(5);
let results = fit_predict_many(|| SeasonalNaive::new(12), &[&short], 3);
assert_eq!(results.len(), 1);
let r = &results[0];
assert!(r.forecast.is_none());
assert!(r.error.is_some());
assert!(r.error.as_ref().unwrap().contains("fit failed"));
}
#[test]
fn test_fit_predict_many_mixed_success_failure() {
let good = make_test_series(30);
let bad = make_test_series(3);
let results = fit_predict_many(|| SeasonalNaive::new(12), &[&good, &bad], 5);
assert_eq!(results.len(), 2);
assert!(results[0].forecast.is_some());
assert!(results[0].error.is_none());
assert!(results[1].forecast.is_none());
assert!(results[1].error.is_some());
}
#[test]
fn test_fit_predict_many_empty_slice() {
let results = fit_predict_many(Naive::new, &[] as &[&TimeSeries], 5);
assert!(results.is_empty());
}
#[test]
fn test_fit_predict_many_model_name_preserved() {
let ts = make_test_series(30);
let results = fit_predict_many(|| WindowAverage::new(5), &[&ts], 3);
assert_eq!(results[0].model_name, "WindowAverage");
}
#[test]
fn test_fit_registry_basic() {
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,
));
let ts = make_test_series(30);
let results = fit_registry(®istry, &ts);
assert_eq!(results.len(), 2);
for r in &results {
assert!(r.error.is_none(), "{} failed: {:?}", r.model_name, r.error);
assert!(r.forecast.is_some());
assert!(r.metrics.is_some());
}
}
#[test]
fn test_fit_registry_has_metrics() {
let mut registry = ModelRegistry::new();
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
let ts = make_test_series(30);
let results = fit_registry(®istry, &ts);
let metrics = results[0].metrics.as_ref().unwrap();
assert!(metrics.mae >= 0.0);
assert!(metrics.rmse >= 0.0);
assert!(metrics.smape >= 0.0);
}
#[test]
fn test_fit_registry_model_names() {
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,
));
let ts = make_test_series(30);
let results = fit_registry(®istry, &ts);
let names: Vec<&str> = results.iter().map(|r| r.model_name.as_str()).collect();
assert_eq!(names, vec!["Naive", "RWD"]);
}
#[test]
fn test_fit_registry_empty_registry() {
let registry = ModelRegistry::new();
let ts = make_test_series(30);
let results = fit_registry(®istry, &ts);
assert!(results.is_empty());
}
#[test]
fn test_fit_registry_records_error_for_bad_model() {
let mut registry = ModelRegistry::new();
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
registry.register(ModelSpec::with_period(
"SeasonalNaive",
|p| Box::new(SeasonalNaive::new(p)),
12,
true,
));
let ts = make_test_series(5);
let results = fit_registry(®istry, &ts);
assert_eq!(results.len(), 2);
assert!(results[0].error.is_none(), "Naive should succeed");
assert!(
results[1].error.is_some(),
"SeasonalNaive(12) should fail with 5 points"
);
}
#[test]
fn test_fit_registry_forecast_horizon_is_one() {
let mut registry = ModelRegistry::new();
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
let ts = make_test_series(30);
let results = fit_registry(®istry, &ts);
let fc = results[0].forecast.as_ref().unwrap();
assert_eq!(fc.horizon(), 1);
}
#[test]
fn test_fit_registry_metrics_reasonable_for_trend() {
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,
));
let ts = make_test_series(30);
let results = fit_registry(®istry, &ts);
let naive_rmse = results[0].metrics.as_ref().unwrap().rmse;
let rwd_rmse = results[1].metrics.as_ref().unwrap().rmse;
assert!(
rwd_rmse <= naive_rmse + 0.01,
"RWD rmse={rwd_rmse} should be <= Naive rmse={naive_rmse}"
);
}
#[test]
fn test_batch_result_success() {
let ts = make_test_series(30);
let mut model = Naive::new();
model.fit(&ts).unwrap();
let result = BatchResult {
model,
series_index: 0,
error: None,
};
assert!(result.error.is_none());
assert!(result.model.is_fitted());
assert_eq!(result.series_index, 0);
}
#[test]
fn test_batch_result_failure() {
let model = Naive::new();
let result = BatchResult {
model,
series_index: 3,
error: Some("insufficient data".to_string()),
};
assert!(result.error.is_some());
assert!(!result.model.is_fitted());
}
}