use crate::error::{ForecastError, Result};
use std::f64::consts::PI;
const SECONDS_PER_DAY: f64 = 86_400.0;
pub fn fourier_terms(timestamps: &[f64], period: f64, order: usize) -> Result<Vec<Vec<f64>>> {
if timestamps.is_empty() {
return Err(ForecastError::EmptyData);
}
if period <= 0.0 || !period.is_finite() {
return Err(ForecastError::InvalidParameter(
"period must be a positive finite number".to_string(),
));
}
if order == 0 {
return Err(ForecastError::InvalidParameter(
"order must be at least 1".to_string(),
));
}
let n = timestamps.len();
let mut basis = Vec::with_capacity(2 * order);
for k in 1..=order {
let mut sin_vec = Vec::with_capacity(n);
let mut cos_vec = Vec::with_capacity(n);
let freq = 2.0 * PI * k as f64 / period;
for &t in timestamps {
let angle = freq * t;
sin_vec.push(angle.sin());
cos_vec.push(angle.cos());
}
basis.push(sin_vec);
basis.push(cos_vec);
}
Ok(basis)
}
#[derive(Debug, Clone)]
pub struct FourierSeasonality {
period: f64,
order: usize,
coefficients: Option<Vec<f64>>,
}
impl FourierSeasonality {
pub fn new(period: f64, order: usize) -> Result<Self> {
if period <= 0.0 || !period.is_finite() {
return Err(ForecastError::InvalidParameter(
"period must be a positive finite number".to_string(),
));
}
if order == 0 {
return Err(ForecastError::InvalidParameter(
"order must be at least 1".to_string(),
));
}
Ok(Self {
period,
order,
coefficients: None,
})
}
pub fn daily(order: usize) -> Result<Self> {
Self::new(SECONDS_PER_DAY, order)
}
pub fn weekly(order: usize) -> Result<Self> {
Self::new(7.0 * SECONDS_PER_DAY, order)
}
pub fn yearly(order: usize) -> Result<Self> {
Self::new(365.25 * SECONDS_PER_DAY, order)
}
pub fn period(&self) -> f64 {
self.period
}
pub fn order(&self) -> usize {
self.order
}
pub fn coefficients(&self) -> Option<&[f64]> {
self.coefficients.as_deref()
}
pub fn fit(&mut self, timestamps: &[f64], values: &[f64]) -> Result<()> {
if timestamps.is_empty() || values.is_empty() {
return Err(ForecastError::EmptyData);
}
if timestamps.len() != values.len() {
return Err(ForecastError::DimensionMismatch {
expected: timestamps.len(),
got: values.len(),
});
}
let n = timestamps.len();
let p = 2 * self.order;
if n < p {
return Err(ForecastError::InsufficientData {
needed: p,
got: n,
hint: Some(format!(
"need at least {} data points for {} Fourier terms",
p, p
)),
});
}
let basis = fourier_terms(timestamps, self.period, self.order)?;
let mut xtx = vec![0.0; p * p];
for i in 0..p {
for j in i..p {
let dot: f64 = basis[i]
.iter()
.zip(basis[j].iter())
.map(|(a, b)| a * b)
.sum();
xtx[i * p + j] = dot;
xtx[j * p + i] = dot;
}
}
let mut xty = vec![0.0; p];
for i in 0..p {
xty[i] = basis[i].iter().zip(values.iter()).map(|(a, b)| a * b).sum();
}
let l = cholesky_decompose(&xtx, p)?;
let coefficients = cholesky_solve(&l, &xty, p);
self.coefficients = Some(coefficients);
Ok(())
}
pub fn predict(&self, timestamps: &[f64]) -> Result<Vec<f64>> {
let coefficients = self
.coefficients
.as_ref()
.ok_or(ForecastError::FitRequired { model: None })?;
if timestamps.is_empty() {
return Err(ForecastError::EmptyData);
}
let basis = fourier_terms(timestamps, self.period, self.order)?;
let n = timestamps.len();
let mut result = vec![0.0; n];
for (j, coeff) in coefficients.iter().enumerate() {
for i in 0..n {
result[i] += coeff * basis[j][i];
}
}
Ok(result)
}
}
fn cholesky_decompose(a: &[f64], n: usize) -> Result<Vec<f64>> {
let mut l = vec![0.0; n * n];
for i in 0..n {
for j in 0..=i {
let mut sum = 0.0;
for k in 0..j {
sum += l[i * n + k] * l[j * n + k];
}
if i == j {
let diag = a[i * n + i] - sum;
if diag <= 0.0 {
return Err(ForecastError::SingularMatrix(
"Fourier basis matrix is singular or nearly singular; \
try reducing the order or ensuring varied timestamps"
.to_string(),
));
}
l[i * n + j] = diag.sqrt();
} else {
l[i * n + j] = (a[i * n + j] - sum) / l[j * n + j];
}
}
}
Ok(l)
}
fn cholesky_solve(l: &[f64], b: &[f64], n: usize) -> Vec<f64> {
let mut z = vec![0.0; n];
for i in 0..n {
let mut sum = 0.0;
for j in 0..i {
sum += l[i * n + j] * z[j];
}
z[i] = (b[i] - sum) / l[i * n + i];
}
let mut x = vec![0.0; n];
for i in (0..n).rev() {
let mut sum = 0.0;
for j in (i + 1)..n {
sum += l[j * n + i] * x[j];
}
x[i] = (z[i] - sum) / l[i * n + i];
}
x
}
#[cfg(test)]
mod tests {
use super::*;
const TOLERANCE: f64 = 1e-6;
fn assert_approx_eq(a: f64, b: f64, tol: f64) {
assert!(
(a - b).abs() < tol,
"expected {} ≈ {}, diff = {}",
a,
b,
(a - b).abs()
);
}
#[test]
fn fourier_terms_basic_shape() {
let ts: Vec<f64> = (0..10).map(|i| i as f64).collect();
let basis = fourier_terms(&ts, 7.0, 3).unwrap();
assert_eq!(basis.len(), 6); for col in &basis {
assert_eq!(col.len(), 10);
}
}
#[test]
fn fourier_terms_values_at_zero() {
let ts = vec![0.0];
let basis = fourier_terms(&ts, 10.0, 2).unwrap();
assert_approx_eq(basis[0][0], 0.0, TOLERANCE); assert_approx_eq(basis[1][0], 1.0, TOLERANCE); assert_approx_eq(basis[2][0], 0.0, TOLERANCE); assert_approx_eq(basis[3][0], 1.0, TOLERANCE); }
#[test]
fn fourier_terms_periodicity() {
let period = 7.0;
let ts = vec![1.5, 1.5 + period, 1.5 + 2.0 * period];
let basis = fourier_terms(&ts, period, 4).unwrap();
for col in &basis {
assert_approx_eq(col[0], col[1], TOLERANCE);
assert_approx_eq(col[0], col[2], TOLERANCE);
}
}
#[test]
fn fourier_terms_orthogonality() {
let period = 100.0;
let n = 1000;
let ts: Vec<f64> = (0..n).map(|i| i as f64 * period / n as f64).collect();
let basis = fourier_terms(&ts, period, 3).unwrap();
let dot: f64 = basis[0]
.iter()
.zip(basis[1].iter())
.map(|(a, b)| a * b)
.sum();
assert_approx_eq(dot / n as f64, 0.0, 0.01);
let dot: f64 = basis[0]
.iter()
.zip(basis[2].iter())
.map(|(a, b)| a * b)
.sum();
assert_approx_eq(dot / n as f64, 0.0, 0.01);
}
#[test]
fn fourier_terms_empty_timestamps() {
let result = fourier_terms(&[], 7.0, 3);
assert!(matches!(result, Err(ForecastError::EmptyData)));
}
#[test]
fn fourier_terms_invalid_period() {
let ts = vec![1.0, 2.0];
assert!(fourier_terms(&ts, 0.0, 1).is_err());
assert!(fourier_terms(&ts, -5.0, 1).is_err());
assert!(fourier_terms(&ts, f64::NAN, 1).is_err());
assert!(fourier_terms(&ts, f64::INFINITY, 1).is_err());
}
#[test]
fn fourier_terms_invalid_order() {
let ts = vec![1.0, 2.0];
assert!(fourier_terms(&ts, 7.0, 0).is_err());
}
#[test]
fn new_valid_parameters() {
let fs = FourierSeasonality::new(7.0, 3).unwrap();
assert_approx_eq(fs.period(), 7.0, TOLERANCE);
assert_eq!(fs.order(), 3);
assert!(fs.coefficients().is_none());
}
#[test]
fn new_invalid_period() {
assert!(FourierSeasonality::new(0.0, 3).is_err());
assert!(FourierSeasonality::new(-1.0, 3).is_err());
assert!(FourierSeasonality::new(f64::NAN, 3).is_err());
}
#[test]
fn new_invalid_order() {
assert!(FourierSeasonality::new(7.0, 0).is_err());
}
#[test]
fn preset_daily() {
let fs = FourierSeasonality::daily(4).unwrap();
assert_approx_eq(fs.period(), SECONDS_PER_DAY, TOLERANCE);
assert_eq!(fs.order(), 4);
}
#[test]
fn preset_weekly() {
let fs = FourierSeasonality::weekly(3).unwrap();
assert_approx_eq(fs.period(), 7.0 * SECONDS_PER_DAY, TOLERANCE);
assert_eq!(fs.order(), 3);
}
#[test]
fn preset_yearly() {
let fs = FourierSeasonality::yearly(10).unwrap();
assert_approx_eq(fs.period(), 365.25 * SECONDS_PER_DAY, TOLERANCE);
assert_eq!(fs.order(), 10);
}
#[test]
fn fit_recovers_single_sinusoid() {
let period = 7.0;
let n = 100;
let ts: Vec<f64> = (0..n).map(|i| i as f64 * period / n as f64).collect();
let values: Vec<f64> = ts
.iter()
.map(|&t| 3.0 * (2.0 * PI * t / period).sin())
.collect();
let mut model = FourierSeasonality::new(period, 3).unwrap();
model.fit(&ts, &values).unwrap();
let coeffs = model.coefficients().unwrap();
assert_eq!(coeffs.len(), 6);
assert_approx_eq(coeffs[0], 3.0, 0.01);
assert_approx_eq(coeffs[1], 0.0, 0.01);
for &c in &coeffs[2..] {
assert_approx_eq(c, 0.0, 0.01);
}
}
#[test]
fn fit_recovers_mixed_harmonics() {
let period = 10.0;
let n = 200;
let ts: Vec<f64> = (0..n).map(|i| i as f64 * period / n as f64).collect();
let values: Vec<f64> = ts
.iter()
.map(|&t| 2.0 * (2.0 * PI * t / period).sin() + 1.5 * (4.0 * PI * t / period).cos())
.collect();
let mut model = FourierSeasonality::new(period, 3).unwrap();
model.fit(&ts, &values).unwrap();
let coeffs = model.coefficients().unwrap();
assert_approx_eq(coeffs[0], 2.0, 0.01); assert_approx_eq(coeffs[1], 0.0, 0.01); assert_approx_eq(coeffs[2], 0.0, 0.01); assert_approx_eq(coeffs[3], 1.5, 0.01); assert_approx_eq(coeffs[4], 0.0, 0.01); assert_approx_eq(coeffs[5], 0.0, 0.01); }
#[test]
fn predict_matches_original() {
let period = 7.0;
let n = 50;
let ts: Vec<f64> = (0..n).map(|i| i as f64 * 0.5).collect();
let values: Vec<f64> = ts
.iter()
.map(|&t| 2.0 * (2.0 * PI * t / period).sin() - 1.0 * (2.0 * PI * t / period).cos())
.collect();
let mut model = FourierSeasonality::new(period, 3).unwrap();
model.fit(&ts, &values).unwrap();
let predicted = model.predict(&ts).unwrap();
assert_eq!(predicted.len(), n);
for i in 0..n {
assert_approx_eq(predicted[i], values[i], 0.05);
}
}
#[test]
fn predict_on_new_timestamps() {
let period = 7.0;
let n = 100;
let ts: Vec<f64> = (0..n).map(|i| i as f64 * 0.1).collect();
let values: Vec<f64> = ts
.iter()
.map(|&t| 4.0 * (2.0 * PI * t / period).sin())
.collect();
let mut model = FourierSeasonality::new(period, 2).unwrap();
model.fit(&ts, &values).unwrap();
let new_ts: Vec<f64> = (0..20).map(|i| 10.0 + i as f64 * 0.1).collect();
let predicted = model.predict(&new_ts).unwrap();
assert_eq!(predicted.len(), 20);
for (i, &t) in new_ts.iter().enumerate() {
let expected = 4.0 * (2.0 * PI * t / period).sin();
assert_approx_eq(predicted[i], expected, 0.1);
}
}
#[test]
fn predict_before_fit_errors() {
let model = FourierSeasonality::new(7.0, 3).unwrap();
let ts = vec![1.0, 2.0, 3.0];
assert!(matches!(
model.predict(&ts),
Err(ForecastError::FitRequired { .. })
));
}
#[test]
fn fit_empty_data() {
let mut model = FourierSeasonality::new(7.0, 3).unwrap();
assert!(matches!(model.fit(&[], &[]), Err(ForecastError::EmptyData)));
}
#[test]
fn fit_mismatched_lengths() {
let mut model = FourierSeasonality::new(7.0, 3).unwrap();
let ts = vec![1.0, 2.0, 3.0];
let values = vec![1.0, 2.0];
assert!(matches!(
model.fit(&ts, &values),
Err(ForecastError::DimensionMismatch { .. })
));
}
#[test]
fn fit_insufficient_data() {
let mut model = FourierSeasonality::new(7.0, 3).unwrap();
let ts = vec![1.0, 2.0, 3.0];
let values = vec![1.0, 2.0, 3.0];
assert!(matches!(
model.fit(&ts, &values),
Err(ForecastError::InsufficientData { .. })
));
}
#[test]
fn predict_empty_timestamps() {
let period = 7.0;
let ts: Vec<f64> = (0..20).map(|i| i as f64).collect();
let values: Vec<f64> = ts.iter().map(|&t| (2.0 * PI * t / period).sin()).collect();
let mut model = FourierSeasonality::new(period, 2).unwrap();
model.fit(&ts, &values).unwrap();
assert!(matches!(model.predict(&[]), Err(ForecastError::EmptyData)));
}
#[test]
fn cholesky_2x2() {
let a = vec![4.0, 2.0, 2.0, 3.0];
let l = cholesky_decompose(&a, 2).unwrap();
assert_approx_eq(l[0], 2.0, TOLERANCE);
assert_approx_eq(l[1], 0.0, TOLERANCE);
assert_approx_eq(l[2], 1.0, TOLERANCE);
assert_approx_eq(l[3], 2.0_f64.sqrt(), TOLERANCE);
}
#[test]
fn cholesky_solve_2x2() {
let a = vec![4.0, 2.0, 2.0, 3.0];
let l = cholesky_decompose(&a, 2).unwrap();
let b = vec![8.0, 7.0];
let x = cholesky_solve(&l, &b, 2);
let r0 = 4.0 * x[0] + 2.0 * x[1];
let r1 = 2.0 * x[0] + 3.0 * x[1];
assert_approx_eq(r0, 8.0, TOLERANCE);
assert_approx_eq(r1, 7.0, TOLERANCE);
}
#[test]
fn cholesky_singular_matrix() {
let a = vec![1.0, 2.0, 2.0, 1.0];
assert!(cholesky_decompose(&a, 2).is_err());
}
#[test]
fn model_is_cloneable() {
let period = 7.0;
let ts: Vec<f64> = (0..20).map(|i| i as f64).collect();
let values: Vec<f64> = ts.iter().map(|&t| (2.0 * PI * t / period).sin()).collect();
let mut model = FourierSeasonality::new(period, 2).unwrap();
model.fit(&ts, &values).unwrap();
let clone = model.clone();
let pred1 = model.predict(&ts).unwrap();
let pred2 = clone.predict(&ts).unwrap();
for (a, b) in pred1.iter().zip(pred2.iter()) {
assert_approx_eq(*a, *b, TOLERANCE);
}
}
#[test]
fn fit_can_be_called_multiple_times() {
let period = 7.0;
let ts: Vec<f64> = (0..20).map(|i| i as f64).collect();
let values1: Vec<f64> = ts
.iter()
.map(|&t| 2.0 * (2.0 * PI * t / period).sin())
.collect();
let values2: Vec<f64> = ts
.iter()
.map(|&t| 5.0 * (2.0 * PI * t / period).cos())
.collect();
let mut model = FourierSeasonality::new(period, 2).unwrap();
model.fit(&ts, &values1).unwrap();
let pred1 = model.predict(&ts).unwrap();
model.fit(&ts, &values2).unwrap();
let pred2 = model.predict(&ts).unwrap();
let diff: f64 = pred1
.iter()
.zip(pred2.iter())
.map(|(a, b)| (a - b).abs())
.sum();
assert!(diff > 1.0, "re-fitting should change predictions");
}
#[test]
fn multiple_periods_combined() {
let n = 400;
let weekly_period = 7.0;
let yearly_period = 365.25;
let ts: Vec<f64> = (0..n).map(|i| i as f64).collect();
let weekly_component: Vec<f64> = ts
.iter()
.map(|&t| 3.0 * (2.0 * PI * t / weekly_period).sin())
.collect();
let yearly_component: Vec<f64> = ts
.iter()
.map(|&t| 2.0 * (2.0 * PI * t / yearly_period).cos())
.collect();
let values: Vec<f64> = weekly_component
.iter()
.zip(yearly_component.iter())
.map(|(w, y)| w + y)
.collect();
let mut weekly_model = FourierSeasonality::new(weekly_period, 3).unwrap();
weekly_model.fit(&ts, &values).unwrap();
let weekly_pred = weekly_model.predict(&ts).unwrap();
let weekly_rmse: f64 = weekly_pred
.iter()
.zip(weekly_component.iter())
.map(|(p, a)| (p - a).powi(2))
.sum::<f64>()
/ n as f64;
let weekly_rmse = weekly_rmse.sqrt();
assert!(weekly_rmse < 3.0, "weekly RMSE {} is too high", weekly_rmse);
}
#[test]
fn high_order_captures_sharp_pattern() {
let period = 10.0;
let n = 200;
let ts: Vec<f64> = (0..n).map(|i| i as f64 * period / n as f64).collect();
let values: Vec<f64> = ts
.iter()
.map(|&t| {
if (t % period) < period / 2.0 {
1.0
} else {
-1.0
}
})
.collect();
let mut low_order = FourierSeasonality::new(period, 1).unwrap();
low_order.fit(&ts, &values).unwrap();
let pred_low = low_order.predict(&ts).unwrap();
let mut high_order = FourierSeasonality::new(period, 10).unwrap();
high_order.fit(&ts, &values).unwrap();
let pred_high = high_order.predict(&ts).unwrap();
let rmse_low: f64 = pred_low
.iter()
.zip(values.iter())
.map(|(p, a)| (p - a).powi(2))
.sum::<f64>()
/ n as f64;
let rmse_high: f64 = pred_high
.iter()
.zip(values.iter())
.map(|(p, a)| (p - a).powi(2))
.sum::<f64>()
/ n as f64;
assert!(
rmse_high < rmse_low,
"higher order should fit better: rmse_high={} vs rmse_low={}",
rmse_high,
rmse_low
);
}
}