use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::utils::ols::OLSResult;
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub struct FittedParams {
pub params: HashMap<String, f64>,
pub seasonal: Option<Vec<f64>>,
}
pub fn validate_series_complete(series: &TimeSeries) -> Result<()> {
if series.has_missing_values() {
return Err(ForecastError::MissingValues);
}
Ok(())
}
pub trait Forecaster {
fn fit(&mut self, series: &TimeSeries) -> Result<()>;
fn predict(&self, horizon: usize) -> Result<Forecast>;
fn predict_with_intervals(&self, horizon: usize, level: f64) -> Result<Forecast> {
let _ = level;
self.predict(horizon)
}
fn fit_predict(&mut self, series: &TimeSeries, horizon: usize) -> Result<Forecast> {
self.fit(series)?;
self.predict(horizon)
}
fn fit_predict_with_intervals(
&mut self,
series: &TimeSeries,
horizon: usize,
level: f64,
) -> Result<Forecast> {
self.fit(series)?;
self.predict_with_intervals(horizon, level)
}
fn fitted_values(&self) -> Option<&[f64]>;
fn fitted_values_with_intervals(&self, level: f64) -> Option<Forecast> {
let _ = level;
None
}
fn residuals(&self) -> Option<&[f64]>;
fn trend_component(&self) -> Result<&[f64]> {
Err(ForecastError::InvalidParameter(format!(
"{} does not expose a trend component",
self.name()
)))
}
fn seasonal_component(&self) -> Result<&[f64]> {
Err(ForecastError::InvalidParameter(format!(
"{} does not expose a seasonal component",
self.name()
)))
}
fn residual_component(&self) -> Result<Vec<f64>> {
self.residuals().map(|r| r.to_vec()).ok_or_else(|| {
ForecastError::InvalidParameter(format!("{} does not expose residuals", self.name()))
})
}
fn training_values(&self) -> Result<&[f64]> {
Err(ForecastError::InvalidParameter(format!(
"{} does not retain training values",
self.name()
)))
}
fn training_regressors(&self) -> Option<&HashMap<String, Vec<f64>>> {
None
}
fn name(&self) -> &str;
fn is_fitted(&self) -> bool {
self.fitted_values().is_some()
}
fn fitted_params(&self) -> Option<FittedParams> {
None
}
fn supports_exog(&self) -> bool {
false
}
fn has_exog(&self) -> bool {
false
}
fn exog_names(&self) -> Option<&[String]> {
None
}
fn exog_coefficients(&self) -> Option<&OLSResult> {
None
}
fn predict_with_exog(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
) -> Result<Forecast> {
if !self.supports_exog() {
return Err(ForecastError::InvalidParameter(format!(
"{} does not support exogenous variables",
self.name()
)));
}
if !self.has_exog() {
if !future_regressors.is_empty() {
return Err(ForecastError::InvalidParameter(
"Model was not fit with exogenous regressors".into(),
));
}
return self.predict(horizon);
}
Err(ForecastError::InvalidParameter(
"Model was fit with exogenous regressors but predict_with_exog not implemented".into(),
))
}
fn predict_with_exog_intervals(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
level: f64,
) -> Result<Forecast> {
let _ = level;
self.predict_with_exog(horizon, future_regressors)
}
}
pub type BoxedForecaster = Box<dyn Forecaster>;
pub struct ModelSpec {
pub name: String,
pub model_type: String,
factory: Box<dyn Fn() -> BoxedForecaster + Send + Sync>,
pub has_intervals: bool,
}
impl ModelSpec {
pub fn new<F>(name: impl Into<String>, factory: F, has_intervals: bool) -> Self
where
F: Fn() -> BoxedForecaster + Send + Sync + 'static,
{
let name = name.into();
Self {
model_type: name.clone(),
name,
factory: Box::new(factory),
has_intervals,
}
}
pub fn with_type<F>(
name: impl Into<String>,
model_type: impl Into<String>,
factory: F,
has_intervals: bool,
) -> Self
where
F: Fn() -> BoxedForecaster + Send + Sync + 'static,
{
Self {
name: name.into(),
model_type: model_type.into(),
factory: Box::new(factory),
has_intervals,
}
}
pub fn with_period<F>(
name: impl Into<String>,
factory: F,
period: usize,
has_intervals: bool,
) -> Self
where
F: Fn(usize) -> BoxedForecaster + Send + Sync + 'static,
{
let name = name.into();
Self {
model_type: name.clone(),
name,
factory: Box::new(move || factory(period)),
has_intervals,
}
}
pub fn create(&self) -> BoxedForecaster {
(self.factory)()
}
}
pub struct ModelRegistry {
models: Vec<ModelSpec>,
}
impl ModelRegistry {
pub fn new() -> Self {
Self { models: Vec::new() }
}
pub fn register(&mut self, spec: ModelSpec) {
self.models.push(spec);
}
pub fn len(&self) -> usize {
self.models.len()
}
pub fn is_empty(&self) -> bool {
self.models.is_empty()
}
pub fn iter(&self) -> impl Iterator<Item = &ModelSpec> {
self.models.iter()
}
pub fn names(&self) -> Vec<&str> {
self.models.iter().map(|s| s.name.as_str()).collect()
}
pub fn remove(&mut self, name: &str) -> bool {
let before = self.models.len();
self.models.retain(|s| s.name != name);
self.models.len() < before
}
pub fn by_type(&self, model_type: &str) -> Vec<&ModelSpec> {
self.models
.iter()
.filter(|s| s.model_type == model_type)
.collect()
}
pub fn retain<F>(&mut self, f: F)
where
F: FnMut(&ModelSpec) -> bool,
{
self.models.retain(f);
}
pub fn extend(&mut self, other: ModelRegistry) {
self.models.extend(other.models);
}
pub fn contains(&self, name: &str) -> bool {
self.models.iter().any(|s| s.name == name)
}
}
impl Default for ModelRegistry {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::TimeSeries;
use crate::error::ForecastError;
use crate::models::baseline::{Naive, RandomWalkWithDrift, SeasonalNaive, WindowAverage};
use crate::models::exponential::{HoltLinearTrend, SimpleExponentialSmoothing, ETS};
use crate::models::intermittent::Croston;
use crate::models::theta::Theta;
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()
}
#[test]
fn test_boxed_forecaster() {
let model: BoxedForecaster = Box::new(Naive::new());
assert_eq!(model.name(), "Naive");
assert!(!model.is_fitted());
}
#[test]
fn test_boxed_forecaster_fit_predict() {
let mut model: BoxedForecaster = Box::new(Naive::new());
let ts = make_test_series(20);
assert!(model.fit(&ts).is_ok());
assert!(model.is_fitted());
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn test_boxed_forecaster_with_intervals() {
let mut model: BoxedForecaster = Box::new(Naive::new());
let ts = make_test_series(20);
model.fit(&ts).unwrap();
let forecast = model.predict_with_intervals(5, 0.95).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(forecast.has_lower());
assert!(forecast.has_upper());
}
#[test]
fn test_model_spec_simple() {
let spec = ModelSpec::new("Naive", || Box::new(Naive::new()), true);
assert_eq!(spec.name, "Naive");
assert_eq!(spec.model_type, "Naive");
assert!(spec.has_intervals);
let model = spec.create();
assert_eq!(model.name(), "Naive");
}
#[test]
fn test_model_spec_with_type() {
let spec =
ModelSpec::with_type("MFLES_additive", "MFLES", || Box::new(Naive::new()), false);
assert_eq!(spec.name, "MFLES_additive");
assert_eq!(spec.model_type, "MFLES");
assert!(!spec.has_intervals);
}
#[test]
fn test_model_spec_with_type_dynamic_name() {
let name = format!("SMA_{}", 10);
let spec = ModelSpec::new(name, || Box::new(Naive::new()), false);
assert_eq!(spec.name, "SMA_10");
assert_eq!(spec.model_type, "SMA_10");
}
#[test]
fn test_registry_by_type() {
let mut registry = ModelRegistry::new();
registry.register(ModelSpec::with_type(
"MFLES_add",
"MFLES",
|| Box::new(Naive::new()),
false,
));
registry.register(ModelSpec::with_type(
"MFLES_mul",
"MFLES",
|| Box::new(Naive::new()),
false,
));
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
let mfles = registry.by_type("MFLES");
assert_eq!(mfles.len(), 2);
assert_eq!(mfles[0].name, "MFLES_add");
assert_eq!(mfles[1].name, "MFLES_mul");
let naive = registry.by_type("Naive");
assert_eq!(naive.len(), 1);
let empty = registry.by_type("NonExistent");
assert!(empty.is_empty());
}
#[test]
fn test_model_spec_with_period() {
let spec = ModelSpec::with_period(
"SeasonalNaive",
|p| Box::new(SeasonalNaive::new(p)),
12,
true,
);
let model = spec.create();
assert_eq!(model.name(), "SeasonalNaive");
}
#[test]
fn test_model_spec_no_intervals() {
let spec = ModelSpec::new(
"SES",
|| Box::new(SimpleExponentialSmoothing::new(0.3)),
false,
);
assert!(!spec.has_intervals);
}
#[test]
fn test_model_spec_creates_independent_instances() {
let spec = ModelSpec::new("Naive", || Box::new(Naive::new()), true);
let ts = make_test_series(20);
let mut model1 = spec.create();
let model2 = spec.create();
model1.fit(&ts).unwrap();
assert!(model1.is_fitted());
assert!(!model2.is_fitted());
}
#[test]
fn test_model_registry() {
let mut registry = ModelRegistry::new();
assert!(registry.is_empty());
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
assert_eq!(registry.len(), 1);
let names: Vec<_> = registry.iter().map(|s| s.name.as_str()).collect();
assert_eq!(names, vec!["Naive"]);
}
#[test]
fn test_model_registry_default() {
let registry = ModelRegistry::default();
assert!(registry.is_empty());
assert_eq!(registry.len(), 0);
}
#[test]
fn test_registry_batch_create() {
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 models: Vec<_> = registry.iter().map(|s| s.create()).collect();
assert_eq!(models.len(), 2);
assert_eq!(models[0].name(), "Naive");
assert_eq!(models[1].name(), "SeasonalNaive");
}
#[test]
fn test_registry_multiple_models() {
let mut registry = ModelRegistry::new();
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
registry.register(ModelSpec::new(
"RandomWalk",
|| Box::new(RandomWalkWithDrift::new()),
true,
));
registry.register(ModelSpec::new(
"SES",
|| Box::new(SimpleExponentialSmoothing::new(0.3)),
false,
));
registry.register(ModelSpec::with_period(
"WindowAvg",
|p| Box::new(WindowAverage::new(p)),
5,
false,
));
assert_eq!(registry.len(), 4);
let intervals_count = registry.iter().filter(|s| s.has_intervals).count();
assert_eq!(intervals_count, 2);
}
#[test]
fn test_registry_batch_fit_predict() {
let mut registry = ModelRegistry::new();
registry.register(ModelSpec::new("Naive", || Box::new(Naive::new()), true));
registry.register(ModelSpec::new(
"RandomWalk",
|| Box::new(RandomWalkWithDrift::new()),
true,
));
let ts = make_test_series(30);
let mut results = Vec::new();
for spec in registry.iter() {
let mut model = spec.create();
if model.fit(&ts).is_ok() {
if let Ok(forecast) = model.predict(5) {
results.push((spec.name.to_string(), forecast.primary().to_vec()));
}
}
}
assert_eq!(results.len(), 2);
assert_eq!(results[0].1.len(), 5);
assert_eq!(results[1].1.len(), 5);
}
#[test]
fn test_forecaster_trait_methods() {
let mut model = Naive::new();
let ts = make_test_series(20);
assert!(!model.is_fitted());
assert!(model.fitted_values().is_none());
assert!(model.residuals().is_none());
model.fit(&ts).unwrap();
assert!(model.is_fitted());
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
assert_eq!(model.name(), "Naive");
}
#[test]
fn test_boxed_forecaster_residuals() {
let mut model: BoxedForecaster = Box::new(Naive::new());
let ts = make_test_series(20);
model.fit(&ts).unwrap();
let residuals = model.residuals().unwrap();
assert_eq!(residuals.len(), 20);
}
fn make_nan_series() -> TimeSeries {
let timestamps = make_timestamps(20);
let mut values: Vec<f64> = (1..=20).map(|i| i as f64).collect();
values[5] = f64::NAN;
TimeSeries::univariate(timestamps, values).unwrap()
}
fn make_inf_series() -> TimeSeries {
let timestamps = make_timestamps(20);
let mut values: Vec<f64> = (1..=20).map(|i| i as f64).collect();
values[10] = f64::INFINITY;
TimeSeries::univariate(timestamps, values).unwrap()
}
#[test]
fn test_validate_series_complete_ok() {
let ts = make_test_series(20);
assert!(validate_series_complete(&ts).is_ok());
}
#[test]
fn test_validate_series_complete_nan() {
let ts = make_nan_series();
let err = validate_series_complete(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_validate_series_complete_inf() {
let ts = make_inf_series();
let err = validate_series_complete(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_naive_rejects_nan() {
let ts = make_nan_series();
let mut model = Naive::new();
let err = model.fit(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_ses_rejects_nan() {
let ts = make_nan_series();
let mut model = SimpleExponentialSmoothing::new(0.3);
let err = model.fit(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_holt_rejects_nan() {
let ts = make_nan_series();
let mut model = HoltLinearTrend::auto();
let err = model.fit(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_ets_rejects_nan() {
let ts = make_nan_series();
let mut model = ETS::default();
let err = model.fit(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_theta_rejects_nan() {
let ts = make_nan_series();
let mut model = Theta::new();
let err = model.fit(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_croston_rejects_nan() {
let ts = make_nan_series();
let mut model = Croston::new();
let err = model.fit(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_random_walk_rejects_inf() {
let ts = make_inf_series();
let mut model = RandomWalkWithDrift::new();
let err = model.fit(&ts).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_fit_predict_naive() {
let mut model = Naive::new();
let ts = make_test_series(20);
let forecast = model.fit_predict(&ts, 5).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(model.is_fitted());
for &v in forecast.primary() {
assert!((v - 20.0).abs() < 1e-10);
}
}
#[test]
fn test_fit_predict_theta() {
let mut model = Theta::new();
let ts = make_test_series(30);
let forecast = model.fit_predict(&ts, 5).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(model.is_fitted());
}
#[test]
fn test_fit_predict_with_intervals_naive() {
let mut model = Naive::new();
let ts = make_test_series(20);
let forecast = model.fit_predict_with_intervals(&ts, 5, 0.95).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(model.is_fitted());
assert!(forecast.has_lower());
assert!(forecast.has_upper());
}
#[test]
fn test_fit_predict_with_intervals_theta() {
let mut model = Theta::new();
let ts = make_test_series(30);
let forecast = model.fit_predict_with_intervals(&ts, 5, 0.95).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(model.is_fitted());
}
#[test]
fn test_fit_predict_rejects_nan() {
let mut model = Naive::new();
let ts = make_nan_series();
let err = model.fit_predict(&ts, 5).unwrap_err();
assert_eq!(err, ForecastError::MissingValues);
}
#[test]
fn test_fit_predict_boxed() {
let mut model: BoxedForecaster = Box::new(Naive::new());
let ts = make_test_series(20);
let forecast = model.fit_predict(&ts, 5).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(model.is_fitted());
}
}