use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::{validate_series_complete, Forecaster};
use crate::utils::optimization::{nelder_mead, NelderMeadConfig};
use crate::utils::stats::quantile_normal;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum SeasonalESErrorType {
#[default]
Additive,
Multiplicative,
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct SeasonalES {
period: usize,
alpha: f64,
optimize: bool,
error_type: SeasonalESErrorType,
seasonal_values: Option<Vec<f64>>,
fitted: Option<Vec<f64>>,
residuals: Option<Vec<f64>>,
residual_variance: Option<f64>,
n: usize,
training_values_store: Option<Vec<f64>>,
training_regressors_store: Option<std::collections::HashMap<String, Vec<f64>>>,
}
impl SeasonalES {
pub fn new(period: usize) -> Self {
Self {
period,
alpha: 0.1,
optimize: false,
error_type: SeasonalESErrorType::Additive,
seasonal_values: None,
fitted: None,
residuals: None,
residual_variance: None,
n: 0,
training_values_store: None,
training_regressors_store: None,
}
}
pub fn optimized(period: usize) -> Self {
Self {
period,
alpha: 0.1,
optimize: true,
error_type: SeasonalESErrorType::Additive,
seasonal_values: None,
fitted: None,
residuals: None,
residual_variance: None,
n: 0,
training_values_store: None,
training_regressors_store: None,
}
}
pub fn with_params(period: usize, alpha: f64, _gamma: f64) -> Self {
Self {
period,
alpha: alpha.clamp(0.001, 0.999),
optimize: false,
error_type: SeasonalESErrorType::Additive,
seasonal_values: None,
fitted: None,
residuals: None,
residual_variance: None,
n: 0,
training_values_store: None,
training_regressors_store: None,
}
}
pub fn with_error_type(mut self, error_type: SeasonalESErrorType) -> Self {
self.error_type = error_type;
self
}
pub fn alpha(&self) -> f64 {
self.alpha
}
pub fn gamma(&self) -> f64 {
self.alpha
}
pub fn period(&self) -> usize {
self.period
}
pub fn seasonal_indices(&self) -> Option<Vec<f64>> {
self.seasonal_values.as_ref().map(|sv| {
let mean: f64 = sv.iter().sum::<f64>() / sv.len() as f64;
if mean.abs() > 1e-10 {
sv.iter().map(|&v| v / mean).collect()
} else {
vec![1.0; sv.len()]
}
})
}
fn ses_forecast(values: &[f64], alpha: f64) -> (f64, Vec<f64>) {
if values.is_empty() {
return (0.0, Vec::new());
}
let mut fitted = Vec::with_capacity(values.len());
let mut level = values[0];
for &y in values.iter() {
fitted.push(level);
level = alpha * y + (1.0 - alpha) * level;
}
(level, fitted)
}
fn calculate_sse(values: &[f64], alpha: f64, period: usize, n: usize) -> f64 {
if n < period {
return f64::MAX;
}
let mut total_sse = 0.0;
for slot in 0..period {
let init_idx = slot + (n % period);
let slot_values: Vec<f64> = (init_idx..n).step_by(period).map(|i| values[i]).collect();
if slot_values.is_empty() {
continue;
}
let (_, fitted) = Self::ses_forecast(&slot_values, alpha);
for (i, &y) in slot_values.iter().enumerate() {
let error = y - fitted[i];
total_sse += error * error;
}
}
total_sse / n as f64
}
fn optimize_params(values: &[f64], period: usize) -> f64 {
let n = values.len();
let objective = |params: &[f64]| {
let alpha = params[0];
if alpha <= 0.001 || alpha >= 0.999 {
return f64::MAX;
}
Self::calculate_sse(values, alpha, period, n)
};
let starts = [[0.1], [0.3], [0.5], [0.7]];
let mut best_alpha = 0.1;
let mut best_value = f64::MAX;
let config = NelderMeadConfig {
max_iter: 200,
tolerance: 1e-6,
..Default::default()
};
for start in starts {
let result = nelder_mead(objective, &start, Some(&[(0.001, 0.999)]), config);
if result.optimal_value < best_value {
best_value = result.optimal_value;
best_alpha = result.optimal_point[0].clamp(0.001, 0.999);
}
}
best_alpha
}
}
impl Default for SeasonalES {
fn default() -> Self {
Self::new(12)
}
}
impl Forecaster for SeasonalES {
fn fit(&mut self, series: &TimeSeries) -> Result<()> {
validate_series_complete(series)?;
if self.period < 2 {
return Err(ForecastError::InvalidParameter(format!(
"seasonal period must be >= 2, got {}",
self.period
)));
}
let values = series.primary_values();
let n = values.len();
self.n = n;
if n < self.period {
return Err(ForecastError::InsufficientData {
needed: self.period,
got: n,
hint: Some(format!(
"Seasonal exponential smoothing requires at least period = {} observations",
self.period
)),
});
}
if self.optimize {
self.alpha = Self::optimize_params(values, self.period);
}
let mut seasonal_values = vec![f64::NAN; self.period];
let mut fitted = vec![f64::NAN; n];
for slot in 0..self.period {
let init_idx = slot + (n % self.period);
let slot_indices: Vec<usize> = (init_idx..n).step_by(self.period).collect();
if slot_indices.is_empty() {
seasonal_values[slot] = f64::NAN;
continue;
}
let slot_values: Vec<f64> = slot_indices.iter().map(|&i| values[i]).collect();
let (final_level, slot_fitted) = Self::ses_forecast(&slot_values, self.alpha);
seasonal_values[slot] = final_level;
for (i, &idx) in slot_indices.iter().enumerate() {
fitted[idx] = slot_fitted[i];
}
}
let residuals: Vec<f64> = values
.iter()
.zip(fitted.iter())
.map(|(y, f)| if f.is_nan() { f64::NAN } else { y - f })
.collect();
let valid_residuals: Vec<f64> = residuals
.iter()
.filter(|r| r.is_finite())
.copied()
.collect();
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.seasonal_values = Some(seasonal_values);
self.fitted = Some(fitted);
self.residuals = Some(residuals);
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 seasonal_values = self
.seasonal_values
.as_ref()
.ok_or(ForecastError::FitRequired { model: None })?;
if horizon == 0 {
return Ok(Forecast::new());
}
let mut predictions = Vec::with_capacity(horizon);
for h in 0..horizon {
let slot = h % self.period;
let forecast = seasonal_values[slot];
predictions.push(forecast);
}
Ok(Forecast::from_values(predictions))
}
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);
for h in 0..horizon {
let factor = (1.0 + 0.1 * h as f64).sqrt();
let se = (variance * factor).sqrt();
lower.push(preds[h] - z * se);
upper.push(preds[h] + 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 training_values(&self) -> Result<&[f64]> {
self.training_values_store
.as_deref()
.ok_or(ForecastError::FitRequired {
model: Some("SeasonalES".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("SeasonalES".into()),
})
}
fn name(&self) -> &str {
if self.optimize {
"SeasonalES (Optimized)"
} else {
"SeasonalES"
}
}
}
#[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_seasonal_series(n: usize, period: usize) -> TimeSeries {
let timestamps = make_timestamps(n);
let values: Vec<f64> = (0..n)
.map(|i| 50.0 + 10.0 * (2.0 * std::f64::consts::PI * i as f64 / period as f64).sin())
.collect();
TimeSeries::univariate(timestamps, values).unwrap()
}
#[test]
fn seasonal_es_basic() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::new(12);
model.fit(&ts).unwrap();
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn seasonal_es_optimized() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::optimized(12);
model.fit(&ts).unwrap();
assert!(model.alpha() > 0.0 && model.alpha() < 1.0);
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn seasonal_es_with_params() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::with_params(12, 0.2, 0.1);
model.fit(&ts).unwrap();
assert!((model.alpha() - 0.2).abs() < 1e-10);
}
#[test]
fn seasonal_es_seasonal_pattern() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::new(12);
model.fit(&ts).unwrap();
let forecast = model.predict(24).unwrap();
let preds = forecast.primary();
for h in 0..12 {
let diff = (preds[h] - preds[h + 12]).abs();
assert!(
diff < 1e-10,
"Seasonal pattern should repeat exactly: {} vs {}",
preds[h],
preds[h + 12]
);
}
}
#[test]
fn seasonal_es_confidence_intervals() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::new(12);
model.fit(&ts).unwrap();
let forecast = model.predict_with_intervals(12, 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..12 {
assert!(lower[i] < preds[i]);
assert!(upper[i] > preds[i]);
}
}
#[test]
fn seasonal_es_fitted_and_residuals() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::new(12);
model.fit(&ts).unwrap();
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
assert_eq!(model.fitted_values().unwrap().len(), 48);
}
#[test]
fn seasonal_es_insufficient_data() {
let ts = make_seasonal_series(10, 12);
let mut model = SeasonalES::new(12);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InsufficientData { .. })
));
}
#[test]
fn seasonal_es_requires_fit() {
let model = SeasonalES::new(12);
assert!(matches!(
model.predict(5),
Err(ForecastError::FitRequired { .. })
));
}
#[test]
fn seasonal_es_zero_horizon() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::new(12);
model.fit(&ts).unwrap();
let forecast = model.predict(0).unwrap();
assert_eq!(forecast.horizon(), 0);
}
#[test]
fn seasonal_es_name() {
let model = SeasonalES::new(12);
assert_eq!(model.name(), "SeasonalES");
let optimized = SeasonalES::optimized(12);
assert_eq!(optimized.name(), "SeasonalES (Optimized)");
}
#[test]
fn constant_series_produces_constant_forecast() {
let timestamps = make_timestamps(20);
let values = vec![5.0; 20];
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = SeasonalES::new(4);
model.fit(&ts).unwrap();
let forecast = model.predict(8).unwrap();
let preds = forecast.primary();
assert!(preds.iter().all(|v| v.is_finite()));
for &p in preds {
assert!((p - 5.0).abs() < 1e-6, "Expected ~5.0, got {}", p);
}
}
#[test]
fn seasonal_es_statsforecast_match() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| {
50.0 + 0.5 * i as f64 + 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 = SeasonalES::with_params(12, 0.3, 0.0);
model.fit(&ts).unwrap();
let forecast = model.predict(12).unwrap();
let preds = forecast.primary();
for &p in preds {
assert!(p.is_finite());
assert!(p > 30.0 && p < 120.0);
}
}
#[test]
fn seasonal_es_rejects_period_zero() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::new(0);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("seasonal period must be >= 2")
));
}
#[test]
fn seasonal_es_rejects_period_one() {
let ts = make_seasonal_series(48, 12);
let mut model = SeasonalES::new(1);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("seasonal period must be >= 2")
));
}
}