use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::arima::AutoARIMA;
use crate::models::exponential::AutoETS;
use crate::models::theta::AutoTheta;
use crate::models::{validate_series_complete, AutoTBATS, Forecaster, MSTLForecaster, MFLES};
use std::collections::HashMap;
use std::fmt;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
#[derive(Debug, Clone)]
pub struct AutoForecastConfig {
pub seasonal_period: Option<usize>,
pub include_arima: bool,
pub include_ets: bool,
pub include_theta: bool,
pub include_tbats: bool,
pub include_mfles: bool,
pub include_mstl: bool,
}
impl Default for AutoForecastConfig {
fn default() -> Self {
Self {
seasonal_period: None,
include_arima: true,
include_ets: true,
include_theta: true,
include_tbats: false,
include_mfles: false,
include_mstl: false,
}
}
}
impl AutoForecastConfig {
pub fn with_period(period: usize) -> Self {
Self {
seasonal_period: Some(period),
..Default::default()
}
}
pub fn without_arima(mut self) -> Self {
self.include_arima = false;
self
}
pub fn without_ets(mut self) -> Self {
self.include_ets = false;
self
}
pub fn without_theta(mut self) -> Self {
self.include_theta = false;
self
}
pub fn with_tbats(mut self) -> Self {
self.include_tbats = true;
self
}
pub fn with_mfles(mut self) -> Self {
self.include_mfles = true;
self
}
pub fn with_mstl(mut self) -> Self {
self.include_mstl = true;
self
}
}
#[derive(Debug, Clone)]
#[allow(clippy::large_enum_variant)]
enum SelectedAutoModel {
ARIMA(AutoARIMA),
ETS(AutoETS),
Theta(AutoTheta),
TBATS(AutoTBATS),
Mfles(MFLES),
Mstl(MSTLForecaster),
}
#[derive(Debug, Clone)]
pub struct AutoForecast {
config: AutoForecastConfig,
selected: Option<SelectedAutoModel>,
scores: Vec<(String, f64)>,
}
#[derive(Debug, Clone, Default)]
pub struct AutoForecastBuilder {
seasonal_period: Option<usize>,
include_arima: Option<bool>,
include_ets: Option<bool>,
include_theta: Option<bool>,
include_tbats: Option<bool>,
include_mfles: Option<bool>,
include_mstl: Option<bool>,
}
impl AutoForecastBuilder {
fn new() -> Self {
Self::default()
}
pub fn seasonal_period(mut self, period: usize) -> Self {
self.seasonal_period = Some(period);
self
}
pub fn include_arima(mut self, include: bool) -> Self {
self.include_arima = Some(include);
self
}
pub fn include_ets(mut self, include: bool) -> Self {
self.include_ets = Some(include);
self
}
pub fn include_theta(mut self, include: bool) -> Self {
self.include_theta = Some(include);
self
}
pub fn include_tbats(mut self, include: bool) -> Self {
self.include_tbats = Some(include);
self
}
pub fn include_mfles(mut self, include: bool) -> Self {
self.include_mfles = Some(include);
self
}
pub fn include_mstl(mut self, include: bool) -> Self {
self.include_mstl = Some(include);
self
}
pub fn build(self) -> AutoForecast {
let config = AutoForecastConfig {
seasonal_period: self.seasonal_period,
include_arima: self.include_arima.unwrap_or(true),
include_ets: self.include_ets.unwrap_or(true),
include_theta: self.include_theta.unwrap_or(true),
include_tbats: self.include_tbats.unwrap_or(false),
include_mfles: self.include_mfles.unwrap_or(false),
include_mstl: self.include_mstl.unwrap_or(false),
};
AutoForecast::with_config(config)
}
}
impl AutoForecast {
pub fn builder() -> AutoForecastBuilder {
AutoForecastBuilder::new()
}
pub fn new() -> Self {
Self {
config: AutoForecastConfig::default(),
selected: None,
scores: Vec::new(),
}
}
pub fn with_config(config: AutoForecastConfig) -> Self {
Self {
config,
selected: None,
scores: Vec::new(),
}
}
pub fn seasonal(period: usize) -> Self {
Self::with_config(AutoForecastConfig::with_period(period))
}
pub fn selected_model_name(&self) -> Option<&str> {
self.selected.as_ref().map(|m| match m {
SelectedAutoModel::ARIMA(model) => model.name(),
SelectedAutoModel::ETS(model) => model.name(),
SelectedAutoModel::Theta(model) => model.name(),
SelectedAutoModel::TBATS(model) => model.name(),
SelectedAutoModel::Mfles(model) => model.name(),
SelectedAutoModel::Mstl(model) => model.name(),
})
}
pub fn all_scores(&self) -> &[(String, f64)] {
&self.scores
}
fn fit_cross_validation(&mut self, series: &TimeSeries) -> Result<()> {
use crate::utils::cross_validation::{cross_validate, CVConfig};
let n = series.len();
let horizon = self
.config
.seasonal_period
.filter(|&p| p > 1)
.unwrap_or(5)
.min(n / 4)
.max(1);
let initial_window = (n / 2).max(10).min(n - horizon);
let step_size = horizon.max(1);
let cv_config = CVConfig::expanding(initial_window, horizon).with_step_size(step_size);
let seasonal_period = self.config.seasonal_period;
let mut factories: Vec<
Box<
dyn Fn(&CVConfig, &TimeSeries) -> Option<(SelectedAutoModel, String, f64)>
+ Send
+ Sync,
>,
> = Vec::new();
if self.config.include_arima {
factories.push(Box::new(move |cv_cfg: &CVConfig, ts: &TimeSeries| {
let period = seasonal_period;
let cv_result = cross_validate(cv_cfg, ts, move || match period {
Some(p) if p > 1 => AutoARIMA::seasonal(p),
_ => AutoARIMA::new(),
});
if let Ok(results) = cv_result {
if results.n_folds > 0 && results.aggregated.rmse.is_finite() {
let mut model = match period {
Some(p) if p > 1 => AutoARIMA::seasonal(p),
_ => AutoARIMA::new(),
};
if model.fit(ts).is_ok() {
let name = model.name().to_string();
return Some((
SelectedAutoModel::ARIMA(model),
name,
results.aggregated.rmse,
));
}
}
}
None
}));
}
if self.config.include_ets {
factories.push(Box::new(move |cv_cfg: &CVConfig, ts: &TimeSeries| {
let period = seasonal_period;
let cv_result = cross_validate(cv_cfg, ts, move || match period {
Some(p) if p > 1 => AutoETS::with_period(p),
_ => AutoETS::new(),
});
if let Ok(results) = cv_result {
if results.n_folds > 0 && results.aggregated.rmse.is_finite() {
let mut model = match period {
Some(p) if p > 1 => AutoETS::with_period(p),
_ => AutoETS::new(),
};
if model.fit(ts).is_ok() {
let name = model.name().to_string();
return Some((
SelectedAutoModel::ETS(model),
name,
results.aggregated.rmse,
));
}
}
}
None
}));
}
if self.config.include_theta {
factories.push(Box::new(move |cv_cfg: &CVConfig, ts: &TimeSeries| {
let period = seasonal_period;
let cv_result = cross_validate(cv_cfg, ts, move || match period {
Some(p) if p > 1 => AutoTheta::seasonal(p),
_ => AutoTheta::new(),
});
if let Ok(results) = cv_result {
if results.n_folds > 0 && results.aggregated.rmse.is_finite() {
let mut model = match period {
Some(p) if p > 1 => AutoTheta::seasonal(p),
_ => AutoTheta::new(),
};
if model.fit(ts).is_ok() {
let name = model.name().to_string();
return Some((
SelectedAutoModel::Theta(model),
name,
results.aggregated.rmse,
));
}
}
}
None
}));
}
if self.config.include_tbats {
if let Some(p) = seasonal_period.filter(|&p| p > 1) {
factories.push(Box::new(move |cv_cfg: &CVConfig, ts: &TimeSeries| {
let cv_result = cross_validate(cv_cfg, ts, move || AutoTBATS::new(vec![p]));
if let Ok(results) = cv_result {
if results.n_folds > 0 && results.aggregated.rmse.is_finite() {
let mut model = AutoTBATS::new(vec![p]);
if model.fit(ts).is_ok() {
let name = model.name().to_string();
return Some((
SelectedAutoModel::TBATS(model),
name,
results.aggregated.rmse,
));
}
}
}
None
}));
}
}
if self.config.include_mfles {
let period = seasonal_period.unwrap_or(1);
factories.push(Box::new(move |cv_cfg: &CVConfig, ts: &TimeSeries| {
let cv_result = cross_validate(cv_cfg, ts, move || MFLES::new(vec![period]));
if let Ok(results) = cv_result {
if results.n_folds > 0 && results.aggregated.rmse.is_finite() {
let mut model = MFLES::new(vec![period]);
if model.fit(ts).is_ok() {
let name = model.name().to_string();
return Some((
SelectedAutoModel::Mfles(model),
name,
results.aggregated.rmse,
));
}
}
}
None
}));
}
if self.config.include_mstl {
if let Some(p) = seasonal_period.filter(|&p| p > 1) {
factories.push(Box::new(move |cv_cfg: &CVConfig, ts: &TimeSeries| {
let cv_result =
cross_validate(cv_cfg, ts, move || MSTLForecaster::new(vec![p]));
if let Ok(results) = cv_result {
if results.n_folds > 0 && results.aggregated.rmse.is_finite() {
let mut model = MSTLForecaster::new(vec![p]);
if model.fit(ts).is_ok() {
let name = model.name().to_string();
return Some((
SelectedAutoModel::Mstl(model),
name,
results.aggregated.rmse,
));
}
}
}
None
}));
}
}
#[cfg(feature = "parallel")]
let mut cv_scores: Vec<(SelectedAutoModel, String, f64)> = factories
.par_iter()
.filter_map(|f| f(&cv_config, series))
.collect();
#[cfg(not(feature = "parallel"))]
let mut cv_scores: Vec<(SelectedAutoModel, String, f64)> = factories
.iter()
.filter_map(|f| f(&cv_config, series))
.collect();
if cv_scores.is_empty() {
return Err(ForecastError::ConvergenceFailure(
"No candidate model produced valid cross-validation results".to_string(),
));
}
cv_scores.sort_by(|a, b| a.2.partial_cmp(&b.2).unwrap_or(std::cmp::Ordering::Equal));
self.scores = cv_scores.iter().map(|(_, n, s)| (n.clone(), *s)).collect();
let (best_model, _, _) = cv_scores.into_iter().next().unwrap();
self.selected = Some(best_model);
Ok(())
}
}
impl Default for AutoForecast {
fn default() -> Self {
Self::new()
}
}
impl Forecaster for AutoForecast {
fn fit(&mut self, series: &TimeSeries) -> Result<()> {
validate_series_complete(series)?;
if series.len() < 10 {
return Err(ForecastError::InsufficientData {
needed: 10,
got: series.len(),
hint: Some(
"AutoForecast requires at least 10 observations for model comparison".into(),
),
});
}
self.selected = None;
self.scores.clear();
self.fit_cross_validation(series)
}
fn predict(&self, horizon: usize) -> Result<Forecast> {
match self.selected.as_ref() {
Some(SelectedAutoModel::ARIMA(m)) => m.predict(horizon),
Some(SelectedAutoModel::ETS(m)) => m.predict(horizon),
Some(SelectedAutoModel::Theta(m)) => m.predict(horizon),
Some(SelectedAutoModel::TBATS(m)) => m.predict(horizon),
Some(SelectedAutoModel::Mfles(m)) => m.predict(horizon),
Some(SelectedAutoModel::Mstl(m)) => m.predict(horizon),
None => Err(ForecastError::FitRequired { model: None }),
}
}
fn predict_with_intervals(&self, horizon: usize, level: f64) -> Result<Forecast> {
match self.selected.as_ref() {
Some(SelectedAutoModel::ARIMA(m)) => m.predict_with_intervals(horizon, level),
Some(SelectedAutoModel::ETS(m)) => m.predict_with_intervals(horizon, level),
Some(SelectedAutoModel::Theta(m)) => m.predict_with_intervals(horizon, level),
Some(SelectedAutoModel::TBATS(m)) => m.predict_with_intervals(horizon, level),
Some(SelectedAutoModel::Mfles(m)) => m.predict_with_intervals(horizon, level),
Some(SelectedAutoModel::Mstl(m)) => m.predict_with_intervals(horizon, level),
None => Err(ForecastError::FitRequired { model: None }),
}
}
fn fitted_values(&self) -> Option<&[f64]> {
match self.selected.as_ref()? {
SelectedAutoModel::ARIMA(m) => m.fitted_values(),
SelectedAutoModel::ETS(m) => m.fitted_values(),
SelectedAutoModel::Theta(m) => m.fitted_values(),
SelectedAutoModel::TBATS(m) => m.fitted_values(),
SelectedAutoModel::Mfles(m) => m.fitted_values(),
SelectedAutoModel::Mstl(m) => m.fitted_values(),
}
}
fn fitted_values_with_intervals(&self, level: f64) -> Option<Forecast> {
match self.selected.as_ref()? {
SelectedAutoModel::ARIMA(m) => m.fitted_values_with_intervals(level),
SelectedAutoModel::ETS(m) => m.fitted_values_with_intervals(level),
SelectedAutoModel::Theta(m) => m.fitted_values_with_intervals(level),
SelectedAutoModel::TBATS(m) => m.fitted_values_with_intervals(level),
SelectedAutoModel::Mfles(m) => m.fitted_values_with_intervals(level),
SelectedAutoModel::Mstl(m) => m.fitted_values_with_intervals(level),
}
}
fn residuals(&self) -> Option<&[f64]> {
match self.selected.as_ref()? {
SelectedAutoModel::ARIMA(m) => m.residuals(),
SelectedAutoModel::ETS(m) => m.residuals(),
SelectedAutoModel::Theta(m) => m.residuals(),
SelectedAutoModel::TBATS(m) => m.residuals(),
SelectedAutoModel::Mfles(m) => m.residuals(),
SelectedAutoModel::Mstl(m) => m.residuals(),
}
}
fn name(&self) -> &str {
"AutoForecast"
}
fn supports_exog(&self) -> bool {
true
}
fn has_exog(&self) -> bool {
match self.selected.as_ref() {
Some(SelectedAutoModel::ARIMA(m)) => m.has_exog(),
Some(SelectedAutoModel::ETS(m)) => m.has_exog(),
Some(SelectedAutoModel::Theta(m)) => m.has_exog(),
Some(SelectedAutoModel::TBATS(m)) => m.has_exog(),
Some(SelectedAutoModel::Mfles(m)) => m.has_exog(),
Some(SelectedAutoModel::Mstl(m)) => m.has_exog(),
None => false,
}
}
fn exog_names(&self) -> Option<&[String]> {
match self.selected.as_ref()? {
SelectedAutoModel::ARIMA(m) => m.exog_names(),
SelectedAutoModel::ETS(m) => m.exog_names(),
SelectedAutoModel::Theta(m) => m.exog_names(),
SelectedAutoModel::TBATS(m) => m.exog_names(),
SelectedAutoModel::Mfles(m) => m.exog_names(),
SelectedAutoModel::Mstl(m) => m.exog_names(),
}
}
fn predict_with_exog(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
) -> Result<Forecast> {
match self.selected.as_ref() {
Some(SelectedAutoModel::ARIMA(m)) => m.predict_with_exog(horizon, future_regressors),
Some(SelectedAutoModel::ETS(m)) => m.predict_with_exog(horizon, future_regressors),
Some(SelectedAutoModel::Theta(m)) => m.predict_with_exog(horizon, future_regressors),
Some(SelectedAutoModel::TBATS(m)) => m.predict_with_exog(horizon, future_regressors),
Some(SelectedAutoModel::Mfles(m)) => m.predict_with_exog(horizon, future_regressors),
Some(SelectedAutoModel::Mstl(m)) => m.predict_with_exog(horizon, future_regressors),
None => Err(ForecastError::FitRequired { model: None }),
}
}
fn predict_with_exog_intervals(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
level: f64,
) -> Result<Forecast> {
match self.selected.as_ref() {
Some(SelectedAutoModel::ARIMA(m)) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
Some(SelectedAutoModel::ETS(m)) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
Some(SelectedAutoModel::Theta(m)) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
Some(SelectedAutoModel::TBATS(m)) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
Some(SelectedAutoModel::Mfles(m)) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
Some(SelectedAutoModel::Mstl(m)) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
None => Err(ForecastError::FitRequired { model: None }),
}
}
}
impl fmt::Display for AutoForecast {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self.selected_model_name() {
Some(name) => {
writeln!(f, "AutoForecast (selected: {})", name)?;
writeln!(f, "Candidate scores:")?;
for (model_name, score) in &self.scores {
writeln!(f, " {}: {:.4}", model_name, score)?;
}
Ok(())
}
None => write!(f, "AutoForecast (not fitted)"),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, TimeZone, Utc};
fn make_timestamps(n: usize) -> Vec<chrono::DateTime<Utc>> {
let base = Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap();
(0..n).map(|i| base + Duration::hours(i as i64)).collect()
}
fn make_trend_series(n: usize) -> TimeSeries {
let timestamps = make_timestamps(n);
let values: Vec<f64> = (0..n)
.map(|i| 10.0 + 0.5 * i as f64 + (i as f64 * 0.3).sin())
.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| {
50.0 + 0.3 * i as f64
+ 10.0 * (2.0 * std::f64::consts::PI * i as f64 / period as f64).sin()
})
.collect();
TimeSeries::univariate(timestamps, values).unwrap()
}
#[test]
fn auto_forecast_basic_fit_predict() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
assert!(model.selected_model_name().is_some());
assert!(!model.all_scores().is_empty());
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn auto_forecast_opt_in_models_extend_candidate_set() {
let ts = make_seasonal_series(120, 12);
let mut baseline = AutoForecast::seasonal(12);
baseline.fit(&ts).unwrap();
let baseline_n = baseline.all_scores().len();
let mut extended = AutoForecast::builder()
.seasonal_period(12)
.include_tbats(true)
.include_mfles(true)
.include_mstl(true)
.build();
extended.fit(&ts).unwrap();
let extended_n = extended.all_scores().len();
assert!(
extended_n > baseline_n,
"extended fit should include more candidates; baseline={}, extended={}",
baseline_n,
extended_n,
);
let extended_names: Vec<&str> = extended
.all_scores()
.iter()
.map(|(n, _)| n.as_str())
.collect();
assert!(
extended_names.iter().any(|n| n.contains("TBATS")),
"extended scores missing TBATS: {:?}",
extended_names
);
assert!(
extended_names.contains(&"MFLES"),
"extended scores missing MFLES: {:?}",
extended_names
);
assert!(
extended_names.iter().any(|n| n.contains("MSTL")),
"extended scores missing MSTL: {:?}",
extended_names
);
let forecast = extended.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn auto_forecast_tbats_mstl_require_seasonal_period() {
let ts = make_trend_series(80);
let mut model = AutoForecast::builder()
.include_tbats(true)
.include_mstl(true)
.build();
model.fit(&ts).unwrap();
let names: Vec<&str> = model.all_scores().iter().map(|(n, _)| n.as_str()).collect();
assert!(
!names.iter().any(|n| n.contains("TBATS")),
"TBATS should be skipped without seasonal_period"
);
assert!(
!names.iter().any(|n| n.contains("MSTL")),
"MSTL should be skipped without seasonal_period"
);
}
#[test]
fn auto_forecast_selects_across_families() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
let scores = model.all_scores();
assert!(
scores.len() >= 2,
"Expected at least 2 candidates, got {}",
scores.len()
);
}
#[test]
fn auto_forecast_scores_sorted() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
let scores = model.all_scores();
for i in 1..scores.len() {
assert!(
scores[i].1 >= scores[i - 1].1,
"Scores not sorted: {} > {}",
scores[i - 1].1,
scores[i].1
);
}
}
#[test]
fn auto_forecast_confidence_intervals() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
let forecast = model.predict_with_intervals(5, 0.95).unwrap();
assert!(forecast.has_lower());
assert!(forecast.has_upper());
}
#[test]
fn auto_forecast_fitted_and_residuals() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
}
#[test]
fn auto_forecast_requires_fit() {
let model = AutoForecast::new();
assert!(matches!(
model.predict(5),
Err(ForecastError::FitRequired { .. })
));
}
#[test]
fn auto_forecast_insufficient_data() {
let timestamps = make_timestamps(5);
let values = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoForecast::new();
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InsufficientData { .. })
));
}
#[test]
fn auto_forecast_name() {
let model = AutoForecast::new();
assert_eq!(model.name(), "AutoForecast");
}
#[test]
fn auto_forecast_default() {
let model = AutoForecast::default();
assert!(model.selected_model_name().is_none());
assert!(model.all_scores().is_empty());
}
#[test]
fn auto_forecast_seasonal() {
let ts = make_seasonal_series(100, 12);
let mut model = AutoForecast::seasonal(12);
model.fit(&ts).unwrap();
assert!(model.selected_model_name().is_some());
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn auto_forecast_arima_only() {
let ts = make_trend_series(100);
let config = AutoForecastConfig::default().without_ets().without_theta();
let mut model = AutoForecast::with_config(config);
model.fit(&ts).unwrap();
let scores = model.all_scores();
assert_eq!(scores.len(), 1);
assert!(scores[0].0.contains("AutoARIMA"));
}
#[test]
fn auto_forecast_ets_only() {
let ts = make_trend_series(100);
let config = AutoForecastConfig::default()
.without_arima()
.without_theta();
let mut model = AutoForecast::with_config(config);
model.fit(&ts).unwrap();
let scores = model.all_scores();
assert_eq!(scores.len(), 1);
assert!(scores[0].0.contains("AutoETS"));
}
#[test]
fn auto_forecast_theta_only() {
let ts = make_trend_series(100);
let config = AutoForecastConfig::default().without_arima().without_ets();
let mut model = AutoForecast::with_config(config);
model.fit(&ts).unwrap();
let scores = model.all_scores();
assert_eq!(scores.len(), 1);
assert!(scores[0].0.contains("AutoTheta"));
}
#[test]
fn auto_forecast_display_fitted() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
let display = format!("{}", model);
assert!(display.contains("AutoForecast"));
assert!(display.contains("selected:"));
assert!(display.contains("Candidate scores:"));
}
#[test]
fn auto_forecast_display_not_fitted() {
let model = AutoForecast::new();
let display = format!("{}", model);
assert_eq!(display, "AutoForecast (not fitted)");
}
#[test]
fn auto_forecast_refit_clears_state() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
let first_name = model.selected_model_name().unwrap().to_string();
let first_scores_len = model.all_scores().len();
model.fit(&ts).unwrap();
assert!(model.selected_model_name().is_some());
assert_eq!(model.all_scores().len(), first_scores_len);
assert_eq!(model.selected_model_name().unwrap(), first_name);
}
#[test]
fn auto_forecast_clone() {
let ts = make_trend_series(100);
let mut model = AutoForecast::new();
model.fit(&ts).unwrap();
let cloned = model.clone();
assert_eq!(cloned.selected_model_name(), model.selected_model_name());
assert_eq!(cloned.all_scores().len(), model.all_scores().len());
let forecast = cloned.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn auto_forecast_no_candidates_enabled() {
let ts = make_trend_series(100);
let config = AutoForecastConfig::default()
.without_arima()
.without_ets()
.without_theta();
let mut model = AutoForecast::with_config(config);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::ConvergenceFailure(_))
));
}
#[test]
fn auto_forecast_builder_defaults() {
let model = AutoForecast::builder().build();
assert!(model.selected_model_name().is_none());
assert_eq!(model.name(), "AutoForecast");
}
#[test]
fn auto_forecast_builder_custom() {
let model = AutoForecast::builder()
.seasonal_period(12)
.include_arima(true)
.include_ets(true)
.include_theta(false)
.build();
let ts = make_seasonal_series(100, 12);
let mut model = model;
model.fit(&ts).unwrap();
assert!(model.selected_model_name().is_some());
let scores = model.all_scores();
assert!(scores.len() <= 2);
for (name, _) in scores {
assert!(!name.contains("AutoTheta"));
}
}
#[test]
fn auto_forecast_builder_fit_predict() {
let ts = make_trend_series(100);
let mut model = AutoForecast::builder()
.include_arima(true)
.include_ets(true)
.build();
model.fit(&ts).unwrap();
assert!(model.selected_model_name().is_some());
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn auto_forecast_builder_arima_only() {
let ts = make_trend_series(100);
let mut model = AutoForecast::builder()
.include_arima(true)
.include_ets(false)
.include_theta(false)
.build();
model.fit(&ts).unwrap();
let scores = model.all_scores();
assert_eq!(scores.len(), 1);
assert!(scores[0].0.contains("AutoARIMA"));
}
#[test]
fn auto_forecast_builder_seasonal() {
let ts = make_seasonal_series(100, 12);
let mut model = AutoForecast::builder().seasonal_period(12).build();
model.fit(&ts).unwrap();
assert!(model.selected_model_name().is_some());
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
}