use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::explain::{Explainable, ForecastExplanation};
use crate::models::{validate_series_complete, FittedParams, Forecaster};
use crate::utils::ols::{ols_fit, ols_residuals, OLSResult};
use crate::utils::optimization::{nelder_mead, NelderMeadConfig};
use crate::utils::stats::quantile_normal;
use statrs::distribution::{ContinuousCDF, Normal};
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum DecompositionType {
Additive,
#[default]
Multiplicative,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct Theta {
theta: f64,
alpha: Option<f64>,
optimize: bool,
seasonal_period: usize,
decomposition_type: DecompositionType,
decomposition_fallback: bool,
b: Option<f64>,
level: Option<f64>,
seasonals: Option<Vec<f64>>,
seasonal_forecast: Option<Vec<f64>>,
#[cfg_attr(feature = "serde", serde(with = "crate::utils::persistence::nan_vec"))]
fitted: Option<Vec<f64>>,
#[cfg_attr(feature = "serde", serde(with = "crate::utils::persistence::nan_vec"))]
residuals: Option<Vec<f64>>,
residual_variance: Option<f64>,
n: usize,
#[cfg_attr(feature = "serde", serde(skip))]
exog_ols: Option<OLSResult>,
skip_optimization: bool,
}
impl Theta {
pub fn new() -> Self {
Self {
theta: 2.0,
alpha: Some(0.1), optimize: false, seasonal_period: 0,
decomposition_type: DecompositionType::Multiplicative,
decomposition_fallback: false,
b: None,
level: None,
seasonals: None,
seasonal_forecast: None,
fitted: None,
residuals: None,
residual_variance: None,
n: 0,
exog_ols: None,
skip_optimization: false,
}
}
pub fn with_theta(theta: f64) -> Self {
Self {
theta,
..Self::new()
}
}
pub fn with_theta_value(theta: f64, alpha: f64, level: f64, b: f64) -> Self {
Self {
theta,
alpha: Some(alpha.clamp(0.0001, 0.9999)),
optimize: false,
seasonal_period: 0,
decomposition_type: DecompositionType::Multiplicative,
decomposition_fallback: false,
b: Some(b),
level: Some(level),
seasonals: None,
seasonal_forecast: None,
fitted: None,
residuals: None,
residual_variance: None,
n: 0,
exog_ols: None,
skip_optimization: true,
}
}
pub fn seasonal(period: usize) -> Self {
Self {
seasonal_period: period,
decomposition_type: DecompositionType::Multiplicative,
..Self::new()
}
}
pub fn seasonal_with_decomposition(period: usize, decomposition: DecompositionType) -> Self {
Self {
seasonal_period: period,
decomposition_type: decomposition,
..Self::new()
}
}
pub fn with_alpha(alpha: f64) -> Self {
Self {
alpha: Some(alpha.clamp(0.0001, 0.9999)),
optimize: false,
..Self::new()
}
}
pub fn optimized() -> Self {
Self {
alpha: None,
optimize: true,
..Self::new()
}
}
pub fn seasonal_optimized(period: usize) -> Self {
Self {
seasonal_period: period,
decomposition_type: DecompositionType::Multiplicative,
alpha: None,
optimize: true,
..Self::new()
}
}
pub fn theta(&self) -> f64 {
self.theta
}
pub fn alpha(&self) -> Option<f64> {
self.alpha
}
pub fn slope(&self) -> Option<f64> {
self.b
}
#[deprecated(note = "Use slope() instead")]
pub fn drift(&self) -> Option<f64> {
self.b
}
pub fn decomposition_type(&self) -> DecompositionType {
self.decomposition_type
}
pub fn used_fallback(&self) -> bool {
self.decomposition_fallback
}
fn deseasonalize(&self, series: &[f64], seasonals: &[f64]) -> Vec<f64> {
if seasonals.is_empty() || self.seasonal_period == 0 {
return series.to_vec();
}
match self.decomposition_type {
DecompositionType::Additive => {
series
.iter()
.enumerate()
.map(|(i, &y)| y - seasonals[i % self.seasonal_period])
.collect()
}
DecompositionType::Multiplicative => {
series
.iter()
.enumerate()
.map(|(i, &y)| {
let s = seasonals[i % self.seasonal_period];
if s.abs() < 1e-10 {
y
} else {
y / s
}
})
.collect()
}
}
}
fn reseasonalize(&self, forecasts: &[f64], start_idx: usize, seasonals: &[f64]) -> Vec<f64> {
if seasonals.is_empty() || self.seasonal_period == 0 {
return forecasts.to_vec();
}
match self.decomposition_type {
DecompositionType::Additive => {
forecasts
.iter()
.enumerate()
.map(|(i, &y)| y + seasonals[(start_idx + i) % self.seasonal_period])
.collect()
}
DecompositionType::Multiplicative => {
forecasts
.iter()
.enumerate()
.map(|(i, &y)| y * seasonals[(start_idx + i) % self.seasonal_period])
.collect()
}
}
}
fn calculate_seasonal_component(
&self,
series: &[f64],
decomposition: DecompositionType,
) -> (Vec<f64>, Vec<f64>) {
let period = self.seasonal_period;
if period == 0 || series.len() < 2 * period {
return (vec![], vec![]);
}
let half = period / 2;
let mut trend = vec![f64::NAN; series.len()];
for i in half..(series.len() - half) {
let sum: f64 = if period % 2 == 0 {
let mut s = 0.5 * series[i - half] + 0.5 * series[i + half];
for &val in series.iter().take(i + half).skip(i - half + 1) {
s += val;
}
s / period as f64
} else {
let start = i - period / 2;
let end = i + period / 2 + 1;
series[start..end].iter().sum::<f64>() / period as f64
};
trend[i] = sum;
}
let detrended: Vec<f64> = match decomposition {
DecompositionType::Additive => series
.iter()
.zip(trend.iter())
.map(|(&y, &t)| if t.is_nan() { f64::NAN } else { y - t })
.collect(),
DecompositionType::Multiplicative => series
.iter()
.zip(trend.iter())
.map(|(&y, &t)| {
if t.is_nan() || t.abs() < 1e-10 {
f64::NAN
} else {
y / t
}
})
.collect(),
};
let mut seasonal_indices = vec![0.0; period];
let mut counts = vec![0usize; period];
for (i, &d) in detrended.iter().enumerate() {
if !d.is_nan() {
seasonal_indices[i % period] += d;
counts[i % period] += 1;
}
}
for i in 0..period {
if counts[i] > 0 {
seasonal_indices[i] /= counts[i] as f64;
}
}
match decomposition {
DecompositionType::Additive => {
let mean = seasonal_indices.iter().sum::<f64>() / period as f64;
for s in &mut seasonal_indices {
*s -= mean;
}
}
DecompositionType::Multiplicative => {
let mean = seasonal_indices.iter().sum::<f64>() / period as f64;
if mean.abs() > 1e-10 {
for s in &mut seasonal_indices {
*s /= mean;
}
}
}
}
let full_seasonal: Vec<f64> = (0..series.len())
.map(|i| seasonal_indices[i % period])
.collect();
let last_cycle: Vec<f64> = full_seasonal[(series.len() - period)..].to_vec();
(full_seasonal, last_cycle)
}
fn calculate_seasonals(&self, series: &[f64], decomposition: DecompositionType) -> Vec<f64> {
let (_, last_cycle) = self.calculate_seasonal_component(series, decomposition);
if last_cycle.is_empty() {
return vec![];
}
let period = self.seasonal_period;
let mut seasonals = vec![0.0; period];
for (i, &s) in last_cycle.iter().enumerate() {
seasonals[i] = s;
}
seasonals
}
fn determine_decomposition(&mut self, series: &[f64]) -> DecompositionType {
if self.decomposition_type == DecompositionType::Additive {
return DecompositionType::Additive;
}
if series.iter().any(|&y| y <= 0.0) {
self.decomposition_fallback = true;
return DecompositionType::Additive;
}
let seasonals = self.calculate_seasonals(series, DecompositionType::Multiplicative);
if !seasonals.is_empty() {
if seasonals.iter().any(|&s| s < 0.01) {
self.decomposition_fallback = true;
return DecompositionType::Additive;
}
}
DecompositionType::Multiplicative
}
fn calculate_sse(series: &[f64], alpha: f64) -> f64 {
if series.is_empty() {
return f64::MAX;
}
let mut level = series[0];
let mut sse = 0.0;
for &y in &series[1..] {
let error = y - level;
sse += error * error;
level = alpha * y + (1.0 - alpha) * level;
}
sse
}
fn optimize_alpha(series: &[f64]) -> f64 {
let config = NelderMeadConfig {
max_iter: 500,
tolerance: 1e-8,
..Default::default()
};
let result = nelder_mead(
|params| Self::calculate_sse(series, params[0]),
&[0.5],
Some(&[(0.0001, 0.9999)]),
config,
);
result.optimal_point[0].clamp(0.0001, 0.9999)
}
fn acf(series: &[f64], nlags: usize) -> Vec<f64> {
let n = series.len();
if n < 2 || nlags == 0 {
return vec![1.0];
}
let mean = crate::simd::mean(series);
let var = crate::simd::variance(series);
if var < 1e-10 {
return vec![1.0; nlags + 1];
}
let mut acf_values = Vec::with_capacity(nlags + 1);
acf_values.push(1.0);
for lag in 1..=nlags {
if lag >= n {
acf_values.push(0.0);
continue;
}
let mut sum = 0.0;
for i in 0..(n - lag) {
sum += (series[i] - mean) * (series[i + lag] - mean);
}
acf_values.push(sum / (n as f64 * var));
}
acf_values
}
fn seasonal_test(series: &[f64], period: usize) -> bool {
if period < 4 || series.len() < 2 * period {
return false;
}
let acf_vals = Self::acf(series, period);
let r: Vec<f64> = acf_vals[1..].to_vec();
let r_sq_sum = crate::simd::sum_of_squares(&r[..r.len() - 1]);
let stat = ((1.0 + 2.0 * r_sq_sum) / series.len() as f64).sqrt();
let r_m = r[r.len() - 1];
let normal = Normal::new(0.0, 1.0).unwrap();
let z_90 = normal.inverse_cdf(0.90);
(r_m.abs() / stat) > z_90
}
fn predict_internal(
&self,
horizon: usize,
future_regressors: Option<&HashMap<String, Vec<f64>>>,
) -> Result<Forecast> {
let smoothed = self
.level
.ok_or(ForecastError::FitRequired { model: None })?;
let alpha = self
.alpha
.ok_or(ForecastError::FitRequired { model: None })?;
let b = self.b.ok_or(ForecastError::FitRequired { model: None })?;
if horizon == 0 {
return Ok(Forecast::new());
}
let exog_contribution = if let Some(ols) = &self.exog_ols {
let future = future_regressors.ok_or_else(|| {
ForecastError::InvalidParameter(
"Model was fit with exogenous regressors. Future regressor values required."
.into(),
)
})?;
for name in &ols.regressor_names {
let values = future.get(name).ok_or_else(|| {
ForecastError::InvalidParameter(format!(
"Missing future values for regressor '{}'",
name
))
})?;
if values.len() != horizon {
return Err(ForecastError::DimensionMismatch {
expected: horizon,
got: values.len(),
});
}
}
Some(ols.predict(future)?)
} else {
if future_regressors.is_some_and(|r| !r.is_empty()) {
return Err(ForecastError::InvalidParameter(
"Model was not fit with exogenous regressors".into(),
));
}
None
};
let mut predictions = Vec::with_capacity(horizon);
for h in 1..=horizon {
let mut forecast =
smoothed + (1.0 - 1.0 / self.theta) * b * (1.0 / alpha + (h as f64 - 1.0));
if let Some(ref exog) = exog_contribution {
forecast += exog[h - 1];
}
predictions.push(forecast);
}
let predictions = if let Some(ref seasonal_forecast) = self.seasonal_forecast {
self.reseasonalize(&predictions, 0, seasonal_forecast)
} else {
predictions
};
Ok(Forecast::from_values(predictions))
}
}
impl Explainable for Theta {
fn explain(&self, horizon: usize) -> Result<ForecastExplanation> {
let smoothed = self
.level
.ok_or(ForecastError::FitRequired { model: None })?;
let alpha = self
.alpha
.ok_or(ForecastError::FitRequired { model: None })?;
let b = self.b.ok_or(ForecastError::FitRequired { model: None })?;
if horizon == 0 {
return Ok(ForecastExplanation {
level: vec![],
trend: None,
seasonal: None,
residual: None,
named_components: vec![],
});
}
let level_component = vec![smoothed; horizon];
let trend_component: Vec<f64> = (1..=horizon)
.map(|h| (1.0 - 1.0 / self.theta) * b * (1.0 / alpha + (h as f64 - 1.0)))
.collect();
let base_forecasts: Vec<f64> = level_component
.iter()
.zip(trend_component.iter())
.map(|(&l, &t)| l + t)
.collect();
let seasonal_component = if let Some(ref sf) = self.seasonal_forecast {
let reseasonalized = self.reseasonalize(&base_forecasts, 0, sf);
Some(
reseasonalized
.iter()
.zip(base_forecasts.iter())
.map(|(&r, &b)| r - b)
.collect(),
)
} else {
None
};
Ok(ForecastExplanation {
level: level_component,
trend: Some(trend_component),
seasonal: seasonal_component,
residual: None,
named_components: vec![],
})
}
}
impl Default for Theta {
fn default() -> Self {
Self::new()
}
}
impl Forecaster for Theta {
fn fit(&mut self, series: &TimeSeries) -> Result<()> {
validate_series_complete(series)?;
let raw_values = series.primary_values();
if raw_values.len() < 4 {
return Err(ForecastError::InsufficientData {
needed: 4,
got: raw_values.len(),
hint: Some("Theta requires at least 4 observations for trend estimation".into()),
});
}
let values: Vec<f64> = if series.has_regressors() {
let regressors = series.all_regressors();
let ols_result = ols_fit(raw_values, ®ressors)?;
let adjusted = ols_residuals(raw_values, &ols_result, ®ressors)?;
self.exog_ols = Some(ols_result);
adjusted
} else {
self.exog_ols = None;
raw_values.to_vec()
};
self.n = values.len();
self.decomposition_fallback = false;
let should_decompose = self.seasonal_period >= 4
&& values.len() >= 2 * self.seasonal_period
&& Self::seasonal_test(&values, self.seasonal_period);
let effective_decomposition = if should_decompose {
self.determine_decomposition(&values)
} else {
self.decomposition_type
};
self.decomposition_type = effective_decomposition;
let (full_seasonal, seasonal_forecast) = if should_decompose {
self.calculate_seasonal_component(&values, effective_decomposition)
} else {
(vec![], vec![])
};
let deseasonalized = self.deseasonalize(&values, &full_seasonal);
let n = deseasonalized.len();
let x_mean = (n - 1) as f64 / 2.0;
let y_mean = deseasonalized.iter().sum::<f64>() / n as f64;
let mut ss_xx = 0.0;
let mut ss_xy = 0.0;
for (i, &y) in deseasonalized.iter().enumerate() {
let x = i as f64;
ss_xx += (x - x_mean).powi(2);
ss_xy += (x - x_mean) * (y - y_mean);
}
let b = if ss_xx > 0.0 { ss_xy / ss_xx } else { 0.0 };
self.b = Some(b);
if self.optimize && !self.skip_optimization {
self.alpha = Some(Self::optimize_alpha(&deseasonalized));
}
let alpha = self
.alpha
.ok_or(ForecastError::FitRequired { model: None })?;
let mut level = deseasonalized[0];
let mut fitted = Vec::with_capacity(self.n);
let mut residuals = Vec::with_capacity(self.n);
let first_fitted = if full_seasonal.is_empty() {
deseasonalized[0]
} else {
match self.decomposition_type {
DecompositionType::Additive => deseasonalized[0] + full_seasonal[0],
DecompositionType::Multiplicative => deseasonalized[0] * full_seasonal[0],
}
};
fitted.push(first_fitted);
residuals.push(0.0);
for i in 1..self.n {
let forecast = level;
let seasonalized_forecast = if full_seasonal.is_empty() {
forecast
} else {
match self.decomposition_type {
DecompositionType::Additive => forecast + full_seasonal[i],
DecompositionType::Multiplicative => forecast * full_seasonal[i],
}
};
fitted.push(seasonalized_forecast);
residuals.push(values[i] - seasonalized_forecast);
level = alpha * deseasonalized[i] + (1.0 - alpha) * level;
}
self.level = Some(level);
self.seasonals = if seasonal_forecast.is_empty() {
None
} else {
let period = self.seasonal_period;
let mut averaged_indices = vec![0.0; period];
let start_pos = (self.n - period) % period;
for i in 0..period {
averaged_indices[(start_pos + i) % period] = seasonal_forecast[i];
}
Some(averaged_indices)
};
self.seasonal_forecast = if seasonal_forecast.is_empty() {
None
} else {
Some(seasonal_forecast)
};
self.fitted = Some(fitted);
let valid_residuals: Vec<f64> = residuals[1..].to_vec();
if !valid_residuals.is_empty() {
let variance =
crate::simd::sum_of_squares(&valid_residuals) / valid_residuals.len() as f64;
self.residual_variance = Some(variance);
}
self.residuals = Some(residuals);
Ok(())
}
fn predict(&self, horizon: usize) -> Result<Forecast> {
if self.exog_ols.is_some() {
return Err(ForecastError::InvalidParameter(
"Model was fit with exogenous regressors. Use predict_with_exog() and provide future regressor values.".into()
));
}
self.predict_internal(horizon, None)
}
fn supports_exog(&self) -> bool {
true
}
fn has_exog(&self) -> bool {
self.exog_ols.is_some()
}
fn exog_names(&self) -> Option<&[String]> {
self.exog_ols
.as_ref()
.map(|ols| ols.regressor_names.as_slice())
}
fn exog_coefficients(&self) -> Option<&OLSResult> {
self.exog_ols.as_ref()
}
fn predict_with_exog(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
) -> Result<Forecast> {
self.predict_internal(horizon, Some(future_regressors))
}
fn predict_with_exog_intervals(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
level: f64,
) -> Result<Forecast> {
let forecast = self.predict_internal(horizon, Some(future_regressors))?;
let variance = self.residual_variance.unwrap_or(0.0);
if horizon == 0 {
return Ok(forecast);
}
let z = quantile_normal((1.0 + level) / 2.0);
let preds = forecast.primary();
let mut lower = Vec::with_capacity(horizon);
let mut upper = Vec::with_capacity(horizon);
let alpha = self.alpha.unwrap_or(0.3);
for h in 1..=horizon {
let factor = if h == 1 {
1.0
} else {
let beta = 1.0 - alpha;
1.0 + beta.powi(2) * (1.0 - beta.powi(2 * (h as i32 - 1))) / (1.0 - beta.powi(2))
};
let se = (variance * factor).sqrt();
lower.push(preds[h - 1] - z * se);
upper.push(preds[h - 1] + z * se);
}
Ok(Forecast::from_values_with_intervals(
preds.to_vec(),
lower,
upper,
))
}
fn predict_with_intervals(&self, horizon: usize, confidence: f64) -> Result<Forecast> {
let forecast = self.predict(horizon)?;
let variance = self.residual_variance.unwrap_or(0.0);
if horizon == 0 {
return Ok(forecast);
}
let z = quantile_normal((1.0 + confidence) / 2.0);
let preds = forecast.primary();
let mut lower = Vec::with_capacity(horizon);
let mut upper = Vec::with_capacity(horizon);
let alpha = self.alpha.unwrap_or(0.3);
for h in 1..=horizon {
let factor = if h == 1 {
1.0
} else {
let beta = 1.0 - alpha;
1.0 + beta.powi(2) * (1.0 - beta.powi(2 * (h as i32 - 1))) / (1.0 - beta.powi(2))
};
let se = (variance * factor).sqrt();
lower.push(preds[h - 1] - z * se);
upper.push(preds[h - 1] + z * se);
}
Ok(Forecast::from_values_with_intervals(
preds.to_vec(),
lower,
upper,
))
}
fn fitted_values(&self) -> Option<&[f64]> {
self.fitted.as_deref()
}
fn fitted_values_with_intervals(&self, level: f64) -> Option<Forecast> {
let fitted = self.fitted.as_ref()?;
let variance = self.residual_variance?;
if variance <= 0.0 {
return Some(Forecast::from_values(fitted.clone()));
}
let z = quantile_normal((1.0 + level) / 2.0);
let sigma = variance.sqrt();
let lower: Vec<f64> = fitted.iter().map(|&f| f - z * sigma).collect();
let upper: Vec<f64> = fitted.iter().map(|&f| f + z * sigma).collect();
Some(Forecast::from_values_with_intervals(
fitted.clone(),
lower,
upper,
))
}
fn residuals(&self) -> Option<&[f64]> {
self.residuals.as_deref()
}
fn name(&self) -> &str {
"Theta"
}
fn fitted_params(&self) -> Option<FittedParams> {
let level = self.level?;
let b = self.b?;
let alpha = self.alpha?;
let mut params = HashMap::new();
params.insert("theta".to_string(), self.theta);
params.insert("alpha".to_string(), alpha);
params.insert("level".to_string(), level);
params.insert("b".to_string(), b);
Some(FittedParams {
params,
seasonal: self.seasonals.clone(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
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 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 = Theta::new();
model.fit(&ts).unwrap();
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn theta_with_trend() {
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 = Theta::new();
model.fit(&ts).unwrap();
let forecast = model.predict(5).unwrap();
let preds = forecast.primary();
assert!(preds[0] > values.last().unwrap() - 10.0);
}
#[test]
fn theta_seasonal() {
let timestamps = make_timestamps(48);
let values: Vec<f64> = (0..48)
.map(|i| 10.0 + 5.0 * (2.0 * std::f64::consts::PI * i as f64 / 12.0).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::seasonal(12);
model.fit(&ts).unwrap();
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn theta_with_alpha() {
let timestamps = make_timestamps(30);
let values: Vec<f64> = (0..30).map(|i| 10.0 + i as f64 * 0.5).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::with_alpha(0.5);
model.fit(&ts).unwrap();
assert_relative_eq!(model.alpha().unwrap(), 0.5, epsilon = 1e-10);
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn theta_custom_theta() {
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 = Theta::with_theta(1.5);
model.fit(&ts).unwrap();
assert_relative_eq!(model.theta(), 1.5, epsilon = 1e-10);
}
#[test]
fn theta_confidence_intervals() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50)
.map(|i| 10.0 + i as f64 * 0.5 + (i as f64 * 0.2).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::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 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.clone()).unwrap();
let mut model = Theta::new();
model.fit(&ts).unwrap();
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
let fitted = model.fitted_values().unwrap();
assert_eq!(fitted.len(), 30);
}
#[test]
fn theta_insufficient_data() {
let timestamps = make_timestamps(3);
let values = vec![1.0, 2.0, 3.0];
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::new();
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InsufficientData { .. })
));
}
#[test]
fn theta_requires_fit() {
let model = Theta::new();
assert!(matches!(
model.predict(5),
Err(ForecastError::FitRequired { .. })
));
}
#[test]
fn theta_zero_horizon() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::new();
model.fit(&ts).unwrap();
let forecast = model.predict(0).unwrap();
assert_eq!(forecast.horizon(), 0);
}
#[test]
fn theta_name() {
let model = Theta::new();
assert_eq!(model.name(), "Theta");
}
#[test]
fn theta_default() {
let model = Theta::default();
assert_relative_eq!(model.theta(), 2.0, epsilon = 1e-10);
}
#[test]
fn theta_slope_positive_for_trend() {
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).unwrap();
let mut model = Theta::new();
model.fit(&ts).unwrap();
assert!(model.slope().unwrap() > 0.0);
}
#[test]
fn theta_multiplicative_default() {
let model = Theta::seasonal(12);
assert_eq!(
model.decomposition_type(),
DecompositionType::Multiplicative
);
assert!(!model.used_fallback());
}
#[test]
fn theta_additive_explicit() {
let model = Theta::seasonal_with_decomposition(12, DecompositionType::Additive);
assert_eq!(model.decomposition_type(), DecompositionType::Additive);
}
#[test]
fn theta_multiplicative_seasonal_positive_data() {
let timestamps = make_timestamps(48);
let values: Vec<f64> = (0..48)
.map(|i| {
let base = 100.0;
let seasonal_factor =
1.0 + 0.2 * (2.0 * std::f64::consts::PI * i as f64 / 12.0).sin();
base * seasonal_factor
})
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::seasonal(12);
model.fit(&ts).unwrap();
assert_eq!(
model.decomposition_type(),
DecompositionType::Multiplicative
);
assert!(!model.used_fallback());
if let Some(seasonals) = &model.seasonals {
let mean = seasonals.iter().sum::<f64>() / seasonals.len() as f64;
assert_relative_eq!(mean, 1.0, epsilon = 0.05);
}
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn theta_fallback_for_negative_values() {
let timestamps = make_timestamps(48);
let values: Vec<f64> = (0..48)
.map(|i| 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 = Theta::seasonal(12);
model.fit(&ts).unwrap();
assert!(model.used_fallback());
assert_eq!(model.decomposition_type(), DecompositionType::Additive);
}
#[test]
fn theta_additive_seasonals_sum_to_zero() {
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 = Theta::seasonal_with_decomposition(12, DecompositionType::Additive);
model.fit(&ts).unwrap();
assert_eq!(model.decomposition_type(), DecompositionType::Additive);
if let Some(seasonals) = &model.seasonals {
let sum: f64 = seasonals.iter().sum();
assert_relative_eq!(sum, 0.0, epsilon = 1e-10);
}
}
#[test]
fn theta_multiplicative_seasonals_average_to_one() {
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 = Theta::seasonal(12);
model.fit(&ts).unwrap();
assert_eq!(
model.decomposition_type(),
DecompositionType::Multiplicative
);
if let Some(seasonals) = &model.seasonals {
let mean = seasonals.iter().sum::<f64>() / seasonals.len() as f64;
assert_relative_eq!(mean, 1.0, epsilon = 0.01);
}
}
#[test]
fn theta_decomposition_type_enum_default() {
let default_type: DecompositionType = Default::default();
assert_eq!(default_type, DecompositionType::Multiplicative);
}
#[test]
fn theta_stm_uses_fixed_alpha() {
let model = Theta::new();
assert_relative_eq!(model.alpha().unwrap(), 0.1, epsilon = 1e-10);
}
#[test]
fn theta_optimized_finds_alpha() {
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 = Theta::optimized();
assert!(model.alpha().is_none());
model.fit(&ts).unwrap();
let alpha = model.alpha().unwrap();
assert!(alpha > 0.0 && alpha < 1.0);
}
#[test]
fn theta_seasonal_optimized() {
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 = Theta::seasonal_optimized(12);
model.fit(&ts).unwrap();
assert!(model.seasonal_forecast.is_some());
let alpha = model.alpha().unwrap();
assert!(alpha > 0.0 && alpha < 1.0);
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn theta_seasonal_forecast_stored() {
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 = Theta::seasonal(12);
model.fit(&ts).unwrap();
assert!(model.seasonal_forecast.is_some());
assert_eq!(model.seasonal_forecast.as_ref().unwrap().len(), 12);
}
#[test]
fn theta_seasonal_forecast_matches_last_cycle() {
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 = Theta::seasonal_with_decomposition(12, DecompositionType::Additive);
model.fit(&ts).unwrap();
let seasonal_forecast = model.seasonal_forecast.as_ref().unwrap();
assert_eq!(seasonal_forecast.len(), 12);
let sum: f64 = seasonal_forecast.iter().sum();
assert_relative_eq!(sum, 0.0, epsilon = 0.1);
}
#[test]
fn theta_additive_seasonal_forecast_accuracy() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.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 = Theta::seasonal_with_decomposition(12, DecompositionType::Additive);
model.fit(&ts).unwrap();
let forecast = model.predict(12).unwrap();
let preds = forecast.primary();
assert_eq!(preds.len(), 12);
for &p in preds {
assert!(p > 30.0 && p < 70.0, "Forecast {} out of expected range", p);
}
}
#[test]
fn theta_multiplicative_seasonal_forecast_accuracy() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| {
let level = 100.0 + 0.5 * i as f64; let seasonal_factor =
1.0 + 0.3 * (2.0 * std::f64::consts::PI * i as f64 / 12.0).sin();
level * seasonal_factor
})
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::seasonal(12);
model.fit(&ts).unwrap();
assert_eq!(
model.decomposition_type(),
DecompositionType::Multiplicative
);
let forecast = model.predict(12).unwrap();
let preds = forecast.primary();
assert_eq!(preds.len(), 12);
for &p in preds {
assert!(
p > 50.0 && p < 250.0,
"Forecast {} out of expected range",
p
);
}
}
#[test]
fn theta_fallback_preserves_forecast_quality() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 5.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 = Theta::seasonal(12); model.fit(&ts).unwrap();
assert!(model.used_fallback());
assert_eq!(model.decomposition_type(), DecompositionType::Additive);
let forecast = model.predict(12).unwrap();
let preds = forecast.primary();
for &p in preds {
assert!(
p > -10.0 && p < 20.0,
"Forecast {} out of expected range after fallback",
p
);
}
}
#[test]
fn theta_warm_start_predict_without_fit() {
let model = Theta::with_theta_value(2.0, 0.1, 50.0, 0.5);
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
for &v in forecast.primary() {
assert!(v.is_finite());
}
}
#[test]
fn theta_extract_params_then_warm_start() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).map(|i| 10.0 + (i as f64) * 0.3).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::new();
model.fit(&ts).unwrap();
let forecast1 = model.predict(5).unwrap();
let fp = model.fitted_params().unwrap();
let theta = fp.params["theta"];
let alpha = fp.params["alpha"];
let level = fp.params["level"];
let b = fp.params["b"];
let warm = Theta::with_theta_value(theta, alpha, level, b);
let forecast2 = warm.predict(5).unwrap();
for (a, b_val) in forecast1.primary().iter().zip(forecast2.primary().iter()) {
assert_relative_eq!(a, b_val, epsilon = 1e-10);
}
}
#[test]
fn theta_warm_start_fit_refines() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).map(|i| 10.0 + (i as f64) * 0.3).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::with_theta_value(2.0, 0.1, 5.0, 0.1);
model.fit(&ts).unwrap();
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
for &v in forecast.primary() {
assert!(v.is_finite());
assert!(v > 5.0); }
}
#[test]
fn theta_fitted_params_returns_none_before_fit() {
let model = Theta::new();
assert!(model.fitted_params().is_none());
}
#[test]
fn theta_fitted_params_contains_expected_keys() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).map(|i| 10.0 + (i as f64) * 0.3).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = Theta::new();
model.fit(&ts).unwrap();
let fp = model.fitted_params().unwrap();
assert!(fp.params.contains_key("theta"));
assert!(fp.params.contains_key("alpha"));
assert!(fp.params.contains_key("level"));
assert!(fp.params.contains_key("b"));
}
}