use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::{validate_series_complete, Forecaster};
use crate::utils::optimization::{nelder_mead, NelderMeadConfig};
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct TBATS {
seasonal_periods: Vec<usize>,
fourier_k: Vec<usize>,
lambda: Option<f64>,
use_trend: bool,
use_damped_trend: bool,
phi: f64,
arma_p: usize,
arma_q: usize,
alpha: f64,
beta: f64,
gamma_one: Vec<f64>,
gamma_two: Vec<f64>,
#[allow(dead_code)]
ar_coeffs: Vec<f64>,
#[allow(dead_code)]
ma_coeffs: Vec<f64>,
state: Vec<f64>,
fitted: Option<Vec<f64>>,
residuals: Option<Vec<f64>>,
n: usize,
sigma2: f64,
aic: Option<f64>,
max_nm_iter: usize,
}
impl TBATS {
pub fn new(seasonal_periods: Vec<usize>) -> Self {
let n_periods = seasonal_periods.len();
let fourier_k = vec![1; n_periods];
Self {
seasonal_periods,
fourier_k,
lambda: None,
use_trend: true,
use_damped_trend: false,
phi: 1.0,
arma_p: 0,
arma_q: 0,
alpha: 0.09, beta: 0.05, gamma_one: vec![0.0; n_periods], gamma_two: vec![0.0; n_periods], ar_coeffs: Vec::new(),
ma_coeffs: Vec::new(),
state: Vec::new(),
fitted: None,
residuals: None,
n: 0,
sigma2: 1.0,
aic: None,
max_nm_iter: 200,
}
}
pub(crate) fn fit_with_max_iter(&mut self, series: &TimeSeries, max_iter: usize) -> Result<()> {
self.max_nm_iter = max_iter;
self.fit(series)
}
fn max_harmonics(period: usize) -> usize {
if period % 2 == 0 {
period / 2
} else {
(period - 1) / 2
}
}
pub fn default_k(period: usize) -> usize {
if period <= 2 {
1
} else if period <= 12 {
period / 2
} else if period <= 24 {
6
} else if period <= 52 {
10
} else {
15
}
}
fn find_harmonics(values: &[f64], period: usize) -> (usize, Vec<f64>) {
let n = values.len();
let window = 2 * period;
let mut trend = vec![0.0; n];
for i in 0..n {
let start = i.saturating_sub(window / 2);
let end = (i + window / 2 + 1).min(n);
let count = end - start;
trend[i] = values[start..end].iter().sum::<f64>() / count as f64;
}
let z: Vec<f64> = values
.iter()
.zip(trend.iter())
.map(|(y, t)| y - t)
.collect();
let max_k = Self::max_harmonics(period).min(n).min(6);
if max_k == 0 {
return (1, values.to_vec());
}
let max_cols = 2 * max_k;
let mut fourier = vec![0.0; max_cols * n];
for j in 0..max_k {
let freq = 2.0 * std::f64::consts::PI * (j + 1) as f64 / period as f64;
let cos_col = 2 * j;
let sin_col = 2 * j + 1;
for t in 0..n {
let angle = freq * t as f64;
fourier[cos_col * n + t] = angle.cos();
fourier[sin_col * n + t] = angle.sin();
}
}
let mut xtx = vec![vec![0.0; max_cols]; max_cols];
let mut xty = vec![0.0; max_cols];
let mut best_k = 1;
let mut best_aic = f64::INFINITY;
let mut best_residuals = z.clone();
for h in 1..=max_k {
let new_cos = 2 * (h - 1);
let new_sin = 2 * (h - 1) + 1;
let k_params = 2 * h;
for new_col in [new_cos, new_sin] {
for prev_col in 0..k_params {
let mut dot = 0.0;
for t in 0..n {
dot += fourier[new_col * n + t] * fourier[prev_col * n + t];
}
xtx[new_col][prev_col] = dot;
xtx[prev_col][new_col] = dot;
}
let mut dot = 0.0;
for t in 0..n {
dot += fourier[new_col * n + t] * z[t];
}
xty[new_col] = dot;
}
let sub_xtx: Vec<Vec<f64>> = xtx[..k_params]
.iter()
.map(|row| row[..k_params].to_vec())
.collect();
let sub_xty = xty[..k_params].to_vec();
let coeffs = match Self::solve_linear_system(&sub_xtx, &sub_xty) {
Some(c) => c,
None => continue,
};
let mut sse = 0.0;
let mut residuals = vec![0.0; n];
for t in 0..n {
let mut fitted = 0.0;
for j in 0..k_params {
fitted += fourier[j * n + t] * coeffs[j];
}
residuals[t] = z[t] - fitted;
sse += residuals[t] * residuals[t];
}
if sse > 0.0 {
let aic = n as f64 * (sse / n as f64).ln() + 2.0 * k_params as f64;
if aic < best_aic {
best_aic = aic;
best_k = h;
best_residuals = residuals;
}
}
}
(best_k, best_residuals)
}
fn solve_linear_system(a: &[Vec<f64>], b: &[f64]) -> Option<Vec<f64>> {
let n = b.len();
let mut aug = vec![vec![0.0; n + 1]; n];
for i in 0..n {
for j in 0..n {
aug[i][j] = a[i][j];
}
aug[i][n] = b[i];
}
for i in 0..n {
let mut max_row = i;
for k in (i + 1)..n {
if aug[k][i].abs() > aug[max_row][i].abs() {
max_row = k;
}
}
aug.swap(i, max_row);
if aug[i][i].abs() < 1e-12 {
return None;
}
for k in (i + 1)..n {
let factor = aug[k][i] / aug[i][i];
for j in i..=n {
aug[k][j] -= factor * aug[i][j];
}
}
}
let mut x = vec![0.0; n];
for i in (0..n).rev() {
x[i] = aug[i][n];
for j in (i + 1)..n {
x[i] -= aug[i][j] * x[j];
}
x[i] /= aug[i][i];
}
Some(x)
}
pub fn with_box_cox(mut self, lambda: f64) -> Self {
self.lambda = Some(lambda.clamp(0.0, 1.0));
self
}
pub fn without_trend(mut self) -> Self {
self.use_trend = false;
self
}
pub fn with_damped_trend(mut self, phi: f64) -> Self {
self.use_damped_trend = true;
self.phi = phi.clamp(0.8, 0.99);
self
}
pub fn with_arma(mut self, p: usize, q: usize) -> Self {
self.arma_p = p.min(2);
self.arma_q = q.min(2);
self
}
pub fn with_fourier_k(mut self, k: Vec<usize>) -> Self {
for (i, &ki) in k.iter().enumerate() {
if i < self.fourier_k.len() {
let max_k = self.seasonal_periods[i] / 2;
self.fourier_k[i] = ki.min(max_k).max(1);
}
}
self
}
pub fn aic(&self) -> Option<f64> {
self.aic
}
pub fn lambda(&self) -> Option<f64> {
self.lambda
}
fn box_cox_transform(value: f64, lambda: f64) -> f64 {
if lambda.abs() < 1e-10 {
value.ln()
} else {
(value.powf(lambda) - 1.0) / lambda
}
}
fn inverse_box_cox(value: f64, lambda: f64) -> f64 {
if lambda.abs() < 1e-10 {
value.exp()
} else {
let inner = lambda * value + 1.0;
if inner > 0.0 {
inner.powf(1.0 / lambda)
} else {
0.0
}
}
}
fn estimate_lambda(values: &[f64]) -> f64 {
if values.iter().any(|&v| v <= 0.0) {
return 1.0; }
let objective = |params: &[f64]| {
let lambda = params[0];
let transformed: Vec<f64> = values
.iter()
.map(|&v| Self::box_cox_transform(v, lambda))
.collect();
let mean = crate::simd::mean(&transformed);
let variance = crate::simd::variance(&transformed);
if mean.abs() < 1e-10 {
f64::MAX
} else {
variance / (mean * mean)
}
};
let config = NelderMeadConfig {
max_iter: 50,
tolerance: 1e-4,
..Default::default()
};
let result = nelder_mead(objective, &[0.5], Some(&[(0.0, 1.0)]), config);
result.optimal_point[0].clamp(0.0, 1.0)
}
fn tau(&self) -> usize {
self.fourier_k.iter().map(|&k| 2 * k).sum()
}
fn state_dim(&self) -> usize {
let base = if self.use_trend { 2 } else { 1 };
base + self.tau()
}
fn initialize_state(&self, values: &[f64]) -> Vec<f64> {
let n = values.len();
let dim = self.state_dim();
let mut state = vec![0.0; dim];
let mean = values.iter().sum::<f64>() / n as f64;
state[0] = mean;
if self.use_trend {
state[1] = 0.0;
}
state
}
fn precompute_trig(&self) -> Vec<(f64, f64)> {
let mut table = Vec::new();
for (period_idx, &k) in self.fourier_k.iter().enumerate() {
let period = self.seasonal_periods[period_idx];
for j in 0..k {
let freq = 2.0 * std::f64::consts::PI * (j + 1) as f64 / period as f64;
table.push((freq.cos(), freq.sin()));
}
}
table
}
fn run_filter(
&self,
values: &[f64],
initial_state: &[f64],
alpha: f64,
beta: f64,
phi: f64,
gamma_one: &[f64],
gamma_two: &[f64],
) -> (f64, Vec<f64>, Vec<f64>, Vec<f64>) {
let n = values.len();
let base = if self.use_trend { 2 } else { 1 };
let trig = self.precompute_trig();
let mut state = initial_state.to_vec();
let mut fitted = Vec::with_capacity(n);
let mut residuals = Vec::with_capacity(n);
let mut sse = 0.0;
for t in 0..n {
let level = state[0];
let trend = if self.use_trend { state[1] } else { 0.0 };
let mut seasonal = 0.0;
let mut pos = base;
for &k in &self.fourier_k {
for j in 0..k {
seasonal += state[pos + 2 * j];
}
pos += 2 * k;
}
let predicted = level + phi * trend + seasonal;
let error = values[t] - predicted;
fitted.push(predicted);
residuals.push(error);
sse += error * error;
state[0] = level + phi * trend + alpha * error;
if self.use_trend {
state[1] = phi * trend + beta * error;
}
let mut pos = base;
let mut trig_idx = 0;
for (period_idx, &k) in self.fourier_k.iter().enumerate() {
let g1 = gamma_one.get(period_idx).copied().unwrap_or(0.0);
let g2 = gamma_two.get(period_idx).copied().unwrap_or(0.0);
for j in 0..k {
let (cos_freq, sin_freq) = trig[trig_idx];
trig_idx += 1;
let idx_cos = pos + 2 * j;
let idx_sin = pos + 2 * j + 1;
let old_cos = state[idx_cos];
let old_sin = state[idx_sin];
state[idx_cos] = cos_freq * old_cos + sin_freq * old_sin + g1 * error;
state[idx_sin] = -sin_freq * old_cos + cos_freq * old_sin + g2 * error;
}
pos += 2 * k;
}
}
(sse, state, fitted, residuals)
}
fn optimize_parameters(&mut self, values: &[f64]) {
let n = values.len();
let n_periods = self.seasonal_periods.len();
let mut initial = vec![0.09]; let mut bounds = vec![(0.001, 0.999)];
if self.use_trend {
initial.push(0.05); bounds.push((-0.5, 0.5));
if self.use_damped_trend {
initial.push(0.98); bounds.push((0.8, 0.999));
}
}
for _ in 0..n_periods {
initial.push(0.0);
bounds.push((-0.1, 0.1));
}
for _ in 0..n_periods {
initial.push(0.0);
bounds.push((-0.1, 0.1));
}
let fourier_k = self.fourier_k.clone();
let use_trend = self.use_trend;
let use_damped = self.use_damped_trend;
let initial_state = self.initialize_state(values);
let base = if use_trend { 2 } else { 1 };
let trig = self.precompute_trig();
let state_len = initial_state.len();
let scratch = std::cell::RefCell::new(vec![0.0f64; state_len]);
let objective = move |params: &[f64]| {
let alpha = params[0];
let mut idx = 1;
let beta = if use_trend {
let b = params[idx];
idx += 1;
b
} else {
0.0
};
let phi = if use_trend && use_damped {
let p = params[idx];
idx += 1;
p
} else if use_trend {
1.0
} else {
0.0
};
let gamma_one_start = idx;
idx += n_periods;
let gamma_two_start = idx;
let mut state = scratch.borrow_mut();
state.copy_from_slice(&initial_state);
let mut sse = 0.0;
for t in 0..n {
let level = state[0];
let trend = if use_trend { state[1] } else { 0.0 };
let mut seasonal = 0.0;
let mut pos = base;
for &k in &fourier_k {
for j in 0..k {
seasonal += state[pos + 2 * j];
}
pos += 2 * k;
}
let predicted = level + phi * trend + seasonal;
let error = values[t] - predicted;
sse += error * error;
state[0] = level + phi * trend + alpha * error;
if use_trend {
state[1] = phi * trend + beta * error;
}
let mut pos = base;
let mut trig_idx = 0;
for (period_idx, &k) in fourier_k.iter().enumerate() {
let g1 = params[gamma_one_start + period_idx];
let g2 = params[gamma_two_start + period_idx];
for j in 0..k {
let (cos_freq, sin_freq) = trig[trig_idx];
trig_idx += 1;
let idx_cos = pos + 2 * j;
let idx_sin = pos + 2 * j + 1;
let old_cos = state[idx_cos];
let old_sin = state[idx_sin];
state[idx_cos] = cos_freq * old_cos + sin_freq * old_sin + g1 * error;
state[idx_sin] = -sin_freq * old_cos + cos_freq * old_sin + g2 * error;
}
pos += 2 * k;
}
}
sse / n as f64
};
let config = NelderMeadConfig {
max_iter: self.max_nm_iter,
tolerance: 1e-7,
..Default::default()
};
let result = nelder_mead(objective, &initial, Some(&bounds), config);
self.alpha = result.optimal_point[0];
let mut idx = 1;
if self.use_trend {
self.beta = result.optimal_point[idx];
idx += 1;
if self.use_damped_trend {
self.phi = result.optimal_point[idx];
idx += 1;
}
}
for i in 0..n_periods {
self.gamma_one[i] = result.optimal_point[idx + i];
}
idx += n_periods;
for i in 0..n_periods {
self.gamma_two[i] = result.optimal_point[idx + i];
}
}
fn n_parameters(&self) -> usize {
let mut k = 2;
if self.lambda.is_some() {
k += 1;
}
if self.use_trend {
k += 1; if self.use_damped_trend {
k += 1; }
}
k += 2 * self.gamma_one.len();
k += self.tau();
k += self.arma_p + self.arma_q;
k
}
}
impl Default for TBATS {
fn default() -> Self {
Self::new(vec![12])
}
}
impl Forecaster for TBATS {
fn fit(&mut self, series: &TimeSeries) -> Result<()> {
validate_series_complete(series)?;
for &period in &self.seasonal_periods {
if period < 2 {
return Err(ForecastError::InvalidParameter(format!(
"seasonal period must be >= 2, got {}",
period
)));
}
}
let values = series.primary_values();
self.n = values.len();
let min_required = self
.seasonal_periods
.iter()
.max()
.copied()
.unwrap_or(4)
.max(10);
if values.len() < min_required {
return Err(ForecastError::InsufficientData {
needed: min_required,
got: values.len(),
hint: Some(format!(
"TBATS requires at least max(max_period, 10) = {} observations for seasonal estimation",
min_required
)),
});
}
let transformed: Vec<f64> = if self.lambda.is_some() || values.iter().all(|&v| v > 0.0) {
let lambda = self.lambda.unwrap_or_else(|| Self::estimate_lambda(values));
self.lambda = Some(lambda);
values
.iter()
.map(|&v| Self::box_cox_transform(v, lambda))
.collect()
} else {
values.to_vec()
};
let mut residuals_for_next = transformed.clone();
for (i, &period) in self.seasonal_periods.iter().enumerate() {
let (k, residuals) = Self::find_harmonics(&residuals_for_next, period);
self.fourier_k[i] = k;
residuals_for_next = residuals;
}
let n = transformed.len();
self.state = self.initialize_state(&transformed);
self.optimize_parameters(&transformed);
let phi = if self.use_trend { self.phi } else { 0.0 };
let (sse, final_state, fitted, _residuals) = self.run_filter(
&transformed,
&self.initialize_state(&transformed),
self.alpha,
self.beta,
phi,
&self.gamma_one,
&self.gamma_two,
);
self.state = final_state;
self.sigma2 = sse / n as f64;
let lambda = self.lambda.unwrap_or(1.0);
let fitted_original: Vec<f64> = fitted
.iter()
.map(|&f| Self::inverse_box_cox(f, lambda))
.collect();
let residuals_original: Vec<f64> = values
.iter()
.zip(fitted_original.iter())
.map(|(y, f)| y - f)
.collect();
let log_likelihood =
-0.5 * n as f64 * (1.0 + (2.0 * std::f64::consts::PI * self.sigma2).ln());
let k = self.n_parameters();
self.aic = Some(-2.0 * log_likelihood + 2.0 * k as f64);
self.fitted = Some(fitted_original);
self.residuals = Some(residuals_original);
Ok(())
}
fn predict(&self, horizon: usize) -> Result<Forecast> {
if self.fitted.is_none() {
return Err(ForecastError::FitRequired { model: None });
}
if horizon == 0 {
return Ok(Forecast::from_values(Vec::new()));
}
let lambda = self.lambda.unwrap_or(1.0);
let mut predictions = Vec::with_capacity(horizon);
let trig = self.precompute_trig();
let mut state = self.state.clone();
let base = if self.use_trend { 2 } else { 1 };
let phi = if self.use_trend { self.phi } else { 0.0 };
for _h in 0..horizon {
let level = state[0];
let trend = if self.use_trend { state[1] } else { 0.0 };
let mut seasonal = 0.0;
let mut pos = base;
for &k in &self.fourier_k {
for j in 0..k {
seasonal += state[pos + 2 * j]; }
pos += 2 * k;
}
let pred_transformed = level + phi * trend + seasonal;
let pred = Self::inverse_box_cox(pred_transformed, lambda);
predictions.push(pred);
state[0] = level + phi * trend;
if self.use_trend {
state[1] = phi * trend;
}
let mut pos = base;
let mut trig_idx = 0;
for &k in &self.fourier_k {
for j in 0..k {
let (cos_freq, sin_freq) = trig[trig_idx];
trig_idx += 1;
let idx_cos = pos + 2 * j;
let idx_sin = pos + 2 * j + 1;
let old_cos = state[idx_cos];
let old_sin = state[idx_sin];
state[idx_cos] = cos_freq * old_cos + sin_freq * old_sin;
state[idx_sin] = -sin_freq * old_cos + cos_freq * old_sin;
}
pos += 2 * k;
}
}
Ok(Forecast::from_values(predictions))
}
fn predict_with_intervals(&self, horizon: usize, level: f64) -> Result<Forecast> {
let point_forecast = self.predict(horizon)?;
if horizon == 0 {
return Ok(point_forecast);
}
let z = crate::utils::quantile_normal(0.5 + level / 2.0);
let std_dev = self.sigma2.sqrt();
let lower: Vec<f64> = point_forecast
.primary()
.iter()
.enumerate()
.map(|(h, &f)| f - z * std_dev * ((h + 1) as f64).sqrt())
.collect();
let upper: Vec<f64> = point_forecast
.primary()
.iter()
.enumerate()
.map(|(h, &f)| f + z * std_dev * ((h + 1) as f64).sqrt())
.collect();
Ok(Forecast::from_values_with_intervals(
point_forecast.primary().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 residuals = self.residuals.as_ref()?;
let valid_residuals: Vec<f64> = residuals.iter().copied().filter(|r| !r.is_nan()).collect();
if valid_residuals.is_empty() {
return Some(Forecast::from_values(fitted.clone()));
}
let n = valid_residuals.len() as f64;
let variance = crate::simd::sum_of_squares(&valid_residuals) / n;
if variance <= 0.0 {
return Some(Forecast::from_values(fitted.clone()));
}
let z = crate::utils::quantile_normal(0.5 + 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 {
"TBATS"
}
}
#[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_complex_seasonal_series(n: usize) -> TimeSeries {
let timestamps = make_timestamps(n);
let values: Vec<f64> = (0..n)
.map(|i| {
let trend = 50.0 + 0.1 * i as f64;
let daily = 10.0 * (2.0 * std::f64::consts::PI * (i % 24) as f64 / 24.0).sin();
let weekly = 5.0 * (2.0 * std::f64::consts::PI * (i % 168) as f64 / 168.0).sin();
let noise = ((i * 17) % 7) as f64 * 0.3 - 1.0;
(trend + daily + weekly + noise).max(1.0) })
.collect();
TimeSeries::univariate(timestamps, values).unwrap()
}
#[test]
fn tbats_basic() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]);
model.fit(&ts).unwrap();
let forecast = model.predict(24).unwrap();
assert_eq!(forecast.horizon(), 24);
}
#[test]
fn tbats_multiple_seasonality() {
let ts = make_complex_seasonal_series(500);
let mut model = TBATS::new(vec![24, 168]);
model.fit(&ts).unwrap();
let forecast = model.predict(48).unwrap();
assert_eq!(forecast.horizon(), 48);
}
#[test]
fn tbats_without_trend() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]).without_trend();
model.fit(&ts).unwrap();
let forecast = model.predict(24).unwrap();
assert_eq!(forecast.horizon(), 24);
}
#[test]
fn tbats_damped_trend() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]).with_damped_trend(0.95);
model.fit(&ts).unwrap();
let forecast = model.predict(24).unwrap();
assert_eq!(forecast.horizon(), 24);
}
#[test]
fn tbats_box_cox() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]).with_box_cox(0.5);
model.fit(&ts).unwrap();
assert!(model.lambda().is_some());
let forecast = model.predict(24).unwrap();
assert_eq!(forecast.horizon(), 24);
}
#[test]
fn tbats_confidence_intervals() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]);
model.fit(&ts).unwrap();
let forecast = model.predict_with_intervals(24, 0.95).unwrap();
assert!(forecast.has_lower());
assert!(forecast.has_upper());
}
#[test]
fn tbats_fitted_and_residuals() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]);
model.fit(&ts).unwrap();
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
assert_eq!(model.fitted_values().unwrap().len(), 200);
}
#[test]
fn tbats_aic() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]);
model.fit(&ts).unwrap();
assert!(model.aic().is_some());
assert!(model.aic().unwrap().is_finite());
}
#[test]
fn tbats_requires_fit() {
let model = TBATS::new(vec![24]);
assert!(matches!(
model.predict(5),
Err(ForecastError::FitRequired { .. })
));
}
#[test]
fn tbats_zero_horizon() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]);
model.fit(&ts).unwrap();
let forecast = model.predict(0).unwrap();
assert_eq!(forecast.horizon(), 0);
}
#[test]
fn tbats_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 = TBATS::new(vec![24]);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InsufficientData { .. })
));
}
#[test]
fn tbats_name() {
let model = TBATS::new(vec![24]);
assert_eq!(model.name(), "TBATS");
}
#[test]
fn tbats_state_space_observation_uses_cosine_only() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![24]);
model.fit(&ts).unwrap();
assert!(!model.state.is_empty());
for &val in &model.state {
assert!(val.is_finite(), "State value should be finite");
}
}
#[test]
fn tbats_level_accumulates_trend() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 50.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 = TBATS::new(vec![12]);
model.fit(&ts).unwrap();
let forecast = model.predict(24).unwrap();
let preds = forecast.primary();
for &p in preds {
assert!(
p > 0.0 && p < 200.0,
"Forecast {} is out of reasonable bounds for stationary data",
p
);
}
let mean_forecast: f64 = preds.iter().sum::<f64>() / preds.len() as f64;
assert!(
(mean_forecast - 50.0).abs() < 30.0,
"Mean forecast {} should be near 50 for stationary data",
mean_forecast
);
}
#[test]
fn tbats_non_damped_trend_constant() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100).map(|i| 10.0 + 0.5 * i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = TBATS::new(vec![12]);
model.fit(&ts).unwrap();
let forecast = model.predict(12).unwrap();
let preds = forecast.primary();
let first = preds[0];
let last = preds[preds.len() - 1];
assert!(
last >= first - 10.0,
"For trending data, forecasts should not decrease dramatically: first={}, last={}",
first,
last
);
}
#[test]
fn tbats_damped_trend_converges() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100).map(|i| 10.0 + 0.5 * i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = TBATS::new(vec![12]).with_damped_trend(0.9);
model.fit(&ts).unwrap();
let forecast = model.predict(50).unwrap();
let preds = forecast.primary();
let early_diff = preds[5] - preds[0];
let late_diff = preds[49] - preds[44];
assert!(
late_diff.abs() <= early_diff.abs() * 2.0 + 10.0,
"Damped trend should reduce growth rate: early_diff={}, late_diff={}",
early_diff,
late_diff
);
}
#[test]
fn tbats_fourier_rotation() {
let timestamps = make_timestamps(200);
let values: Vec<f64> = (0..200)
.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 = TBATS::new(vec![12]);
model.fit(&ts).unwrap();
let forecast = model.predict(24).unwrap();
let preds = forecast.primary();
for i in 0..12 {
let diff = (preds[i] - preds[i + 12]).abs();
assert!(
diff < 20.0,
"Seasonal pattern should repeat: step {} = {}, step {} = {}, diff = {}",
i,
preds[i],
i + 12,
preds[i + 12],
diff
);
}
}
#[test]
fn tbats_box_cox_auto_estimation() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 50.0 + 10.0 * (i as f64 / 10.0).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = TBATS::new(vec![12]);
model.fit(&ts).unwrap();
assert!(model.lambda().is_some());
let lambda = model.lambda().unwrap();
assert!((0.0..=1.0).contains(&lambda));
}
#[test]
fn tbats_deterministic_forecasts() {
let ts = make_complex_seasonal_series(200);
let mut model1 = TBATS::new(vec![24]);
model1.fit(&ts).unwrap();
let pred1 = model1.predict(24).unwrap();
let mut model2 = TBATS::new(vec![24]);
model2.fit(&ts).unwrap();
let pred2 = model2.predict(24).unwrap();
for (p1, p2) in pred1.primary().iter().zip(pred2.primary().iter()) {
assert!(
(p1 - p2).abs() < 1e-6,
"Forecasts should be deterministic: {} vs {}",
p1,
p2
);
}
}
#[test]
fn constant_series_produces_constant_forecast() {
let timestamps = make_timestamps(40);
let values = vec![5.0; 40];
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = TBATS::new(vec![4]);
model.fit(&ts).unwrap();
let forecast = model.predict(8).unwrap();
let preds = forecast.primary();
assert!(
preds.iter().all(|v| v.is_finite()),
"All predictions must be finite, got: {:?}",
preds
);
for &p in preds {
assert!((p - 5.0).abs() < 0.5, "Expected ~5.0, got {}", p);
}
}
#[test]
fn tbats_forecasts_reasonable_magnitude() {
let ts = make_complex_seasonal_series(200);
let values = ts.primary_values();
let data_min = values.iter().cloned().fold(f64::INFINITY, f64::min);
let data_max = values.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let data_range = data_max - data_min;
let mut model = TBATS::new(vec![24]);
model.fit(&ts).unwrap();
let forecast = model.predict(48).unwrap();
let preds = forecast.primary();
for &p in preds {
assert!(
p > data_min - 2.0 * data_range && p < data_max + 2.0 * data_range,
"Forecast {} is outside reasonable range [{}, {}]",
p,
data_min - 2.0 * data_range,
data_max + 2.0 * data_range
);
}
}
#[test]
fn tbats_rejects_period_zero() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![0]);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("seasonal period must be >= 2")
));
}
#[test]
fn tbats_rejects_period_one() {
let ts = make_complex_seasonal_series(200);
let mut model = TBATS::new(vec![1]);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("seasonal period must be >= 2")
));
}
#[test]
fn tbats_rejects_any_period_below_two() {
let ts = make_complex_seasonal_series(500);
let mut model = TBATS::new(vec![24, 0]);
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("seasonal period must be >= 2")
));
}
}