use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::inspect::{Explanation, Inspectable, ThetaExplanation};
use crate::models::theta::{DecompositionType, DynamicTheta, OptimizedTheta, Theta};
use crate::models::{validate_series_complete, Forecaster};
use crate::utils::ols::OLSResult;
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ThetaModelType {
STM,
OTM,
DSTM,
DOTM,
}
impl std::fmt::Display for ThetaModelType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ThetaModelType::STM => write!(f, "STM"),
ThetaModelType::OTM => write!(f, "OTM"),
ThetaModelType::DSTM => write!(f, "DSTM"),
ThetaModelType::DOTM => write!(f, "DOTM"),
}
}
}
#[derive(Debug, Clone)]
enum FittedModel {
STM(Theta),
OTM(OptimizedTheta),
DSTM(DynamicTheta),
DOTM(DynamicTheta),
}
#[derive(Debug, Clone)]
pub struct AutoTheta {
seasonal_period: usize,
decomposition_type: DecompositionType,
selected_type: Option<ThetaModelType>,
fitted_model: Option<FittedModel>,
model_scores: Option<Vec<(ThetaModelType, f64)>>,
n: usize,
training_values_store: Option<Vec<f64>>,
training_regressors_store: Option<std::collections::HashMap<String, Vec<f64>>>,
}
impl AutoTheta {
pub fn new() -> Self {
Self {
seasonal_period: 0,
decomposition_type: DecompositionType::Multiplicative,
selected_type: None,
fitted_model: None,
model_scores: None,
n: 0,
training_values_store: None,
training_regressors_store: None,
}
}
pub fn seasonal(period: usize) -> Self {
Self {
seasonal_period: period,
decomposition_type: DecompositionType::Multiplicative,
selected_type: None,
fitted_model: None,
model_scores: None,
n: 0,
training_values_store: None,
training_regressors_store: None,
}
}
pub fn seasonal_with_decomposition(period: usize, decomposition: DecompositionType) -> Self {
Self {
seasonal_period: period,
decomposition_type: decomposition,
selected_type: None,
fitted_model: None,
model_scores: None,
n: 0,
training_values_store: None,
training_regressors_store: None,
}
}
pub fn selected_model(&self) -> Option<ThetaModelType> {
self.selected_type
}
pub fn model_scores(&self) -> Option<&[(ThetaModelType, f64)]> {
self.model_scores.as_deref()
}
fn calculate_mse(residuals: &[f64]) -> f64 {
if residuals.is_empty() {
return f64::MAX;
}
let valid: Vec<f64> = residuals.iter().skip(1).copied().collect();
if valid.is_empty() {
return f64::MAX;
}
crate::simd::sum_of_squares(&valid) / valid.len() as f64
}
}
impl Default for AutoTheta {
fn default() -> Self {
Self::new()
}
}
impl Forecaster for AutoTheta {
fn fit(&mut self, series: &TimeSeries) -> Result<()> {
validate_series_complete(series)?;
let values = series.primary_values();
if values.len() < 6 {
return Err(ForecastError::InsufficientData {
needed: 6,
got: values.len(),
hint: Some(
"AutoTheta requires at least 6 observations for model comparison".into(),
),
});
}
self.n = values.len();
let mut scores: Vec<(ThetaModelType, f64, FittedModel)> = Vec::new();
{
let mut model = if self.seasonal_period > 0 {
Theta::seasonal_with_decomposition(self.seasonal_period, self.decomposition_type)
} else {
Theta::new()
};
if model.fit(series).is_ok() {
if let Some(residuals) = model.residuals() {
let mse = Self::calculate_mse(residuals);
scores.push((ThetaModelType::STM, mse, FittedModel::STM(model)));
}
}
}
{
let mut model = if self.seasonal_period > 0 {
OptimizedTheta::seasonal_with_decomposition(
self.seasonal_period,
self.decomposition_type,
)
} else {
OptimizedTheta::new()
};
if model.fit(series).is_ok() {
if let Some(residuals) = model.residuals() {
let mse = Self::calculate_mse(residuals);
scores.push((ThetaModelType::OTM, mse, FittedModel::OTM(model)));
}
}
}
{
let mut model = if self.seasonal_period > 0 {
DynamicTheta::seasonal_with_decomposition(
self.seasonal_period,
self.decomposition_type,
)
} else {
DynamicTheta::new(0.1)
};
if model.fit(series).is_ok() {
if let Some(residuals) = model.residuals() {
let mse = Self::calculate_mse(residuals);
scores.push((ThetaModelType::DSTM, mse, FittedModel::DSTM(model)));
}
}
}
{
let mut model = if self.seasonal_period > 0 {
DynamicTheta::seasonal_optimized_with_decomposition(
self.seasonal_period,
self.decomposition_type,
)
} else {
DynamicTheta::optimized()
};
if model.fit(series).is_ok() {
if let Some(residuals) = model.residuals() {
let mse = Self::calculate_mse(residuals);
scores.push((ThetaModelType::DOTM, mse, FittedModel::DOTM(model)));
}
}
}
if scores.is_empty() {
return Err(ForecastError::ConvergenceFailure(
"All Theta model variants failed to fit".to_string(),
));
}
scores.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
let (best_type, _, best_model) = scores.remove(0);
let score_summary: Vec<(ThetaModelType, f64)> =
std::iter::once((best_type, scores.first().map(|s| s.1).unwrap_or(0.0)))
.chain(scores.iter().map(|(t, s, _)| (*t, *s)))
.collect();
self.selected_type = Some(best_type);
self.fitted_model = Some(best_model);
self.model_scores = Some(score_summary);
self.training_values_store = Some(values.to_vec());
let regs = series.all_regressors();
self.training_regressors_store = if regs.is_empty() {
None
} else {
Some(regs.clone())
};
Ok(())
}
fn predict(&self, horizon: usize) -> Result<Forecast> {
let model = self
.fitted_model
.as_ref()
.ok_or(ForecastError::FitRequired { model: None })?;
match model {
FittedModel::STM(m) => m.predict(horizon),
FittedModel::OTM(m) => m.predict(horizon),
FittedModel::DSTM(m) => m.predict(horizon),
FittedModel::DOTM(m) => m.predict(horizon),
}
}
fn predict_with_intervals(&self, horizon: usize, confidence: f64) -> Result<Forecast> {
let model = self
.fitted_model
.as_ref()
.ok_or(ForecastError::FitRequired { model: None })?;
match model {
FittedModel::STM(m) => m.predict_with_intervals(horizon, confidence),
FittedModel::OTM(m) => m.predict_with_intervals(horizon, confidence),
FittedModel::DSTM(m) => m.predict_with_intervals(horizon, confidence),
FittedModel::DOTM(m) => m.predict_with_intervals(horizon, confidence),
}
}
fn fitted_values(&self) -> Option<&[f64]> {
match self.fitted_model.as_ref()? {
FittedModel::STM(m) => m.fitted_values(),
FittedModel::OTM(m) => m.fitted_values(),
FittedModel::DSTM(m) => m.fitted_values(),
FittedModel::DOTM(m) => m.fitted_values(),
}
}
fn fitted_values_with_intervals(&self, level: f64) -> Option<Forecast> {
match self.fitted_model.as_ref()? {
FittedModel::STM(m) => m.fitted_values_with_intervals(level),
FittedModel::OTM(m) => m.fitted_values_with_intervals(level),
FittedModel::DSTM(m) => m.fitted_values_with_intervals(level),
FittedModel::DOTM(m) => m.fitted_values_with_intervals(level),
}
}
fn residuals(&self) -> Option<&[f64]> {
match self.fitted_model.as_ref()? {
FittedModel::STM(m) => m.residuals(),
FittedModel::OTM(m) => m.residuals(),
FittedModel::DSTM(m) => m.residuals(),
FittedModel::DOTM(m) => m.residuals(),
}
}
fn training_values(&self) -> Result<&[f64]> {
self.training_values_store
.as_deref()
.ok_or(ForecastError::FitRequired {
model: Some("AutoTheta".into()),
})
}
fn training_regressors(&self) -> Option<&std::collections::HashMap<String, Vec<f64>>> {
self.training_regressors_store.as_ref()
}
fn trend_component(&self) -> Result<&[f64]> {
self.fitted_values().ok_or(ForecastError::FitRequired {
model: Some("AutoTheta".into()),
})
}
fn name(&self) -> &str {
"AutoTheta"
}
fn supports_exog(&self) -> bool {
true
}
fn has_exog(&self) -> bool {
match &self.fitted_model {
Some(FittedModel::STM(m)) => m.has_exog(),
Some(FittedModel::OTM(m)) => m.has_exog(),
Some(FittedModel::DSTM(m)) => m.has_exog(),
Some(FittedModel::DOTM(m)) => m.has_exog(),
None => false,
}
}
fn exog_names(&self) -> Option<&[String]> {
match self.fitted_model.as_ref()? {
FittedModel::STM(m) => m.exog_names(),
FittedModel::OTM(m) => m.exog_names(),
FittedModel::DSTM(m) => m.exog_names(),
FittedModel::DOTM(m) => m.exog_names(),
}
}
fn exog_coefficients(&self) -> Option<&OLSResult> {
match self.fitted_model.as_ref()? {
FittedModel::STM(m) => m.exog_coefficients(),
FittedModel::OTM(m) => m.exog_coefficients(),
FittedModel::DSTM(m) => m.exog_coefficients(),
FittedModel::DOTM(m) => m.exog_coefficients(),
}
}
fn predict_with_exog(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
) -> Result<Forecast> {
let model = self
.fitted_model
.as_ref()
.ok_or(ForecastError::FitRequired { model: None })?;
match model {
FittedModel::STM(m) => m.predict_with_exog(horizon, future_regressors),
FittedModel::OTM(m) => m.predict_with_exog(horizon, future_regressors),
FittedModel::DSTM(m) => m.predict_with_exog(horizon, future_regressors),
FittedModel::DOTM(m) => m.predict_with_exog(horizon, future_regressors),
}
}
fn predict_with_exog_intervals(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
level: f64,
) -> Result<Forecast> {
let model = self
.fitted_model
.as_ref()
.ok_or(ForecastError::FitRequired { model: None })?;
match model {
FittedModel::STM(m) => m.predict_with_exog_intervals(horizon, future_regressors, level),
FittedModel::OTM(m) => m.predict_with_exog_intervals(horizon, future_regressors, level),
FittedModel::DSTM(m) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
FittedModel::DOTM(m) => {
m.predict_with_exog_intervals(horizon, future_regressors, level)
}
}
}
}
impl Inspectable for AutoTheta {
fn explanation(&self) -> Result<Explanation> {
let model = self
.fitted_model
.as_ref()
.ok_or_else(|| ForecastError::FitRequired {
model: Some("AutoTheta".to_string()),
})?;
let (variant, theta, alpha, fitted_values, residuals) = match model {
FittedModel::STM(m) => (
"STM".to_string(),
m.theta(),
m.alpha(),
m.fitted_values().map(|v| v.to_vec()).unwrap_or_default(),
m.residuals().map(|v| v.to_vec()).unwrap_or_default(),
),
FittedModel::OTM(m) => (
"OTM".to_string(),
m.theta().unwrap_or(2.0),
m.alpha(),
m.fitted_values().map(|v| v.to_vec()).unwrap_or_default(),
m.residuals().map(|v| v.to_vec()).unwrap_or_default(),
),
FittedModel::DSTM(m) => (
"DSTM".to_string(),
m.theta(),
Some(m.alpha()),
m.fitted_values().map(|v| v.to_vec()).unwrap_or_default(),
m.residuals().map(|v| v.to_vec()).unwrap_or_default(),
),
FittedModel::DOTM(m) => (
"DOTM".to_string(),
m.theta(),
Some(m.alpha()),
m.fitted_values().map(|v| v.to_vec()).unwrap_or_default(),
m.residuals().map(|v| v.to_vec()).unwrap_or_default(),
),
};
Ok(Explanation::Theta(ThetaExplanation {
variant,
theta,
alpha,
seasonal_period: self.seasonal_period,
fitted_values,
residuals,
}))
}
}
#[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()
}
#[test]
fn auto_theta_basic() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50)
.map(|i| 10.0 + 0.5 * i as f64 + (i as f64 * 0.3).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoTheta::new();
model.fit(&ts).unwrap();
assert!(model.selected_model().is_some());
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn auto_theta_trending_data() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).map(|i| 10.0 + 2.0 * i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values.clone()).unwrap();
let mut model = AutoTheta::new();
model.fit(&ts).unwrap();
let selected = model.selected_model().unwrap();
println!("Selected model for trending data: {}", selected);
let forecast = model.predict(5).unwrap();
let preds = forecast.primary();
assert!(preds[0] > values.last().unwrap() - 10.0);
}
#[test]
fn auto_theta_model_scores() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).map(|i| 10.0 + 0.5 * i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoTheta::new();
model.fit(&ts).unwrap();
assert!(model.model_scores().is_some());
let scores = model.model_scores().unwrap();
assert!(!scores.is_empty());
for (_, score) in scores {
assert!(score.is_finite());
}
}
#[test]
fn auto_theta_seasonal() {
let timestamps = make_timestamps(48);
let values: Vec<f64> = (0..48)
.map(|i| 50.0 + 10.0 * (2.0 * std::f64::consts::PI * i as f64 / 12.0).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoTheta::seasonal(12);
model.fit(&ts).unwrap();
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn auto_theta_confidence_intervals() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).map(|i| 10.0 + i as f64 * 0.5).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoTheta::new();
model.fit(&ts).unwrap();
let forecast = model.predict_with_intervals(5, 0.95).unwrap();
assert!(forecast.has_lower());
assert!(forecast.has_upper());
let lower = forecast.lower_series(0).unwrap();
let upper = forecast.upper_series(0).unwrap();
let preds = forecast.primary();
for i in 0..5 {
assert!(lower[i] < preds[i]);
assert!(upper[i] > preds[i]);
}
}
#[test]
fn auto_theta_fitted_and_residuals() {
let timestamps = make_timestamps(30);
let values: Vec<f64> = (0..30).map(|i| 10.0 + i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoTheta::new();
model.fit(&ts).unwrap();
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
}
#[test]
fn auto_theta_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 = AutoTheta::new();
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InsufficientData { .. })
));
}
#[test]
fn auto_theta_requires_fit() {
let model = AutoTheta::new();
assert!(matches!(
model.predict(5),
Err(ForecastError::FitRequired { .. })
));
}
#[test]
fn auto_theta_name() {
let model = AutoTheta::new();
assert_eq!(model.name(), "AutoTheta");
}
#[test]
fn auto_theta_default() {
let model = AutoTheta::default();
assert!(model.selected_model().is_none());
}
#[test]
fn theta_model_type_display() {
assert_eq!(format!("{}", ThetaModelType::STM), "STM");
assert_eq!(format!("{}", ThetaModelType::OTM), "OTM");
assert_eq!(format!("{}", ThetaModelType::DSTM), "DSTM");
assert_eq!(format!("{}", ThetaModelType::DOTM), "DOTM");
}
}