#[derive(Debug, Clone)]
pub struct STLResult {
pub trend: Vec<f64>,
pub seasonal: Vec<f64>,
pub remainder: Vec<f64>,
}
impl STLResult {
pub fn seasonal_strength(&self) -> f64 {
let var_remainder = variance(&self.remainder);
let seasonal_plus_remainder: Vec<f64> = self
.seasonal
.iter()
.zip(self.remainder.iter())
.map(|(s, r)| s + r)
.collect();
let var_sr = variance(&seasonal_plus_remainder);
if var_sr < 1e-10 {
return 0.0;
}
(1.0 - var_remainder / var_sr).max(0.0)
}
pub fn trend_strength(&self) -> f64 {
let var_remainder = variance(&self.remainder);
let trend_plus_remainder: Vec<f64> = self
.trend
.iter()
.zip(self.remainder.iter())
.map(|(t, r)| t + r)
.collect();
let var_tr = variance(&trend_plus_remainder);
if var_tr < 1e-10 {
return 0.0;
}
(1.0 - var_remainder / var_tr).max(0.0)
}
}
#[derive(Debug, Clone)]
pub struct STL {
seasonal_period: usize,
seasonal_smoothness: usize,
trend_smoothness: usize,
low_pass_smoothness: usize,
inner_iterations: usize,
outer_iterations: usize,
robust: bool,
}
impl STL {
pub fn new(seasonal_period: usize) -> Self {
let ns = seasonal_period;
let nt = (1.5 * seasonal_period as f64 / (1.0 - 1.5 / ns as f64)).ceil() as usize;
let nt = if nt % 2 == 0 { nt + 1 } else { nt }; let nl = seasonal_period;
let nl = if nl % 2 == 0 { nl + 1 } else { nl };
Self {
seasonal_period,
seasonal_smoothness: ns | 1, trend_smoothness: nt,
low_pass_smoothness: nl,
inner_iterations: 2,
outer_iterations: 0,
robust: false,
}
}
pub fn with_seasonal_smoothness(mut self, ns: usize) -> Self {
self.seasonal_smoothness = if ns % 2 == 0 { ns + 1 } else { ns };
self
}
pub fn with_trend_smoothness(mut self, nt: usize) -> Self {
self.trend_smoothness = if nt % 2 == 0 { nt + 1 } else { nt };
self
}
pub fn robust(mut self) -> Self {
self.robust = true;
self.outer_iterations = 6;
self
}
pub fn with_outer_iterations(mut self, n: usize) -> Self {
self.outer_iterations = n;
if n > 0 {
self.robust = true;
}
self
}
pub fn with_inner_iterations(mut self, n: usize) -> Self {
self.inner_iterations = n;
self
}
pub fn decompose(&self, series: &[f64]) -> Option<STLResult> {
let n = series.len();
if n < 2 * self.seasonal_period {
return None;
}
let mut seasonal = vec![0.0; n];
let mut trend = vec![0.0; n];
let mut weights = vec![1.0; n];
let mut detrended = vec![0.0; n];
let mut deseasonalized = vec![0.0; n];
let mut remainder = vec![0.0; n];
let mut lp_buf_a = vec![0.0; n];
let mut lp_buf_b = vec![0.0; n];
let mut low_pass = vec![0.0; n];
let unit_weights = vec![1.0; n];
let mut cycle_subseries = vec![0.0; n];
let outer_iters = if self.robust {
self.outer_iterations.max(1)
} else {
1
};
for _ in 0..outer_iters {
for _ in 0..self.inner_iterations {
for (d, (y, t)) in detrended.iter_mut().zip(series.iter().zip(trend.iter())) {
*d = y - t;
}
self.smooth_cycle_subseries_into(&detrended, &weights, &mut cycle_subseries);
self.low_pass_filter_into(
&cycle_subseries,
&mut lp_buf_a,
&mut lp_buf_b,
&mut low_pass,
&unit_weights,
);
for i in 0..n {
seasonal[i] = cycle_subseries[i] - low_pass[i];
}
for (d, (y, s)) in deseasonalized
.iter_mut()
.zip(series.iter().zip(seasonal.iter()))
{
*d = y - s;
}
Self::loess_smooth_into(
&deseasonalized,
self.trend_smoothness,
&weights,
&mut trend,
);
}
if self.robust {
for ((r, (y, s)), t) in remainder
.iter_mut()
.zip(series.iter().zip(seasonal.iter()))
.zip(trend.iter())
{
*r = y - s - t;
}
weights = self.compute_robustness_weights(&remainder);
}
}
for ((r, (y, s)), t) in remainder
.iter_mut()
.zip(series.iter().zip(seasonal.iter()))
.zip(trend.iter())
{
*r = y - s - t;
}
Some(STLResult {
trend,
seasonal,
remainder,
})
}
fn smooth_cycle_subseries_into(&self, detrended: &[f64], weights: &[f64], result: &mut [f64]) {
let n = detrended.len();
let period = self.seasonal_period;
let max_subseries_len = n.div_ceil(period);
let mut subseries_values = Vec::with_capacity(max_subseries_len);
let mut subseries_weights = Vec::with_capacity(max_subseries_len);
let mut subseries_indices = Vec::with_capacity(max_subseries_len);
let mut smoothed = Vec::with_capacity(max_subseries_len);
for cycle_pos in 0..period {
subseries_values.clear();
subseries_weights.clear();
subseries_indices.clear();
for (i, (&val, &w)) in detrended.iter().zip(weights.iter()).enumerate() {
if i % period == cycle_pos {
subseries_values.push(val);
subseries_weights.push(w);
subseries_indices.push(i);
}
}
Self::loess_smooth_into(
&subseries_values,
self.seasonal_smoothness,
&subseries_weights,
&mut smoothed,
);
for (&idx, &smooth_val) in subseries_indices.iter().zip(smoothed.iter()) {
result[idx] = smooth_val;
}
}
}
fn low_pass_filter_into(
&self,
series: &[f64],
buf_a: &mut Vec<f64>,
buf_b: &mut Vec<f64>,
result: &mut Vec<f64>,
unit_weights: &[f64],
) {
let period = self.seasonal_period;
Self::moving_average_into(series, period, buf_a);
Self::moving_average_into(buf_a, period, buf_b);
Self::moving_average_into(buf_b, 3, buf_a);
Self::loess_smooth_into(buf_a, self.low_pass_smoothness, unit_weights, result);
}
fn moving_average_into(series: &[f64], window: usize, result: &mut Vec<f64>) {
let n = series.len();
let half = window / 2;
result.resize(n, 0.0);
for (i, res) in result.iter_mut().enumerate() {
let start = i.saturating_sub(half);
let end = (i + half + 1).min(n);
let sum: f64 = series[start..end].iter().sum();
*res = sum / (end - start) as f64;
}
}
fn loess_smooth_into(values: &[f64], span: usize, weights: &[f64], result: &mut Vec<f64>) {
let n = values.len();
result.resize(n, 0.0);
if n == 0 {
return;
}
let half_span = span / 2;
for i in 0..n {
let start = i.saturating_sub(half_span);
let end = (i + half_span + 1).min(n);
let mut sum_weights = 0.0;
let mut sum_values = 0.0;
for j in start..end {
let dist = (i as f64 - j as f64).abs();
let max_dist = half_span as f64 + 1.0;
let u = dist / max_dist;
let tricube = if u < 1.0 {
(1.0 - u.powi(3)).powi(3)
} else {
0.0
};
let w = tricube * weights[j];
sum_weights += w;
sum_values += w * values[j];
}
result[i] = if sum_weights > 0.0 {
sum_values / sum_weights
} else {
values[i]
};
}
}
fn compute_robustness_weights(&self, remainder: &[f64]) -> Vec<f64> {
let n = remainder.len();
let mut sorted: Vec<f64> = remainder.iter().map(|r| r.abs()).collect();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let median = if n % 2 == 0 {
(sorted[n / 2 - 1] + sorted[n / 2]) / 2.0
} else {
sorted[n / 2]
};
let h = 6.0 * median;
remainder
.iter()
.map(|r| {
if h < 1e-10 {
return 1.0;
}
let u = r.abs() / h;
if u < 1.0 {
(1.0 - u * u).powi(2)
} else {
0.0
}
})
.collect()
}
}
impl Default for STL {
fn default() -> Self {
Self::new(12) }
}
fn variance(values: &[f64]) -> f64 {
let n = values.len();
if n < 2 {
return 0.0;
}
let mean: f64 = values.iter().sum::<f64>() / n as f64;
values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / (n - 1) as f64
}
#[cfg(test)]
mod tests {
use super::*;
fn generate_seasonal_series(n: usize, period: usize) -> Vec<f64> {
(0..n)
.map(|i| {
let trend = 0.1 * i as f64;
let seasonal =
10.0 * ((2.0 * std::f64::consts::PI * i as f64 / period as f64).sin());
trend + seasonal
})
.collect()
}
#[test]
fn stl_basic_decomposition() {
let period = 12;
let series = generate_seasonal_series(120, period);
let stl = STL::new(period);
let result = stl.decompose(&series).unwrap();
assert_eq!(result.trend.len(), series.len());
assert_eq!(result.seasonal.len(), series.len());
assert_eq!(result.remainder.len(), series.len());
for i in 0..series.len() {
let reconstructed = result.trend[i] + result.seasonal[i] + result.remainder[i];
assert!(
(series[i] - reconstructed).abs() < 1e-10,
"Reconstruction failed at index {}: {} vs {}",
i,
series[i],
reconstructed
);
}
}
#[test]
fn stl_detects_seasonality() {
let period = 12;
let series = generate_seasonal_series(120, period);
let stl = STL::new(period);
let result = stl.decompose(&series).unwrap();
let strength = result.seasonal_strength();
assert!(
strength > 0.5,
"Expected strong seasonality, got {}",
strength
);
}
#[test]
fn stl_detects_trend() {
let n = 120;
let period = 12;
let series: Vec<f64> = (0..n)
.map(|i| {
let trend = 2.0 * i as f64;
let seasonal =
0.1 * ((2.0 * std::f64::consts::PI * i as f64 / period as f64).sin());
trend + seasonal
})
.collect();
let stl = STL::new(period);
let result = stl.decompose(&series).unwrap();
let strength = result.trend_strength();
assert!(strength > 0.9, "Expected strong trend, got {}", strength);
}
#[test]
fn stl_trend_only() {
let n = 100;
let period = 10;
let series: Vec<f64> = (0..n).map(|i| 5.0 + 0.5 * i as f64).collect();
let stl = STL::new(period);
let result = stl.decompose(&series).unwrap();
let seasonal_var = variance(&result.seasonal);
let series_var = variance(&series);
assert!(
seasonal_var < series_var * 0.1,
"Seasonal variance {} should be small compared to series variance {}",
seasonal_var,
series_var
);
}
#[test]
fn stl_constant_series() {
let n = 100;
let period = 10;
let series = vec![5.0; n];
let stl = STL::new(period);
let result = stl.decompose(&series).unwrap();
for &s in &result.seasonal {
assert!(s.abs() < 1e-6, "Seasonal should be near zero");
}
for &r in &result.remainder {
assert!(r.abs() < 1e-6, "Remainder should be near zero");
}
}
#[test]
fn stl_insufficient_data() {
let period = 12;
let series = vec![1.0; 10];
let stl = STL::new(period);
assert!(stl.decompose(&series).is_none());
}
#[test]
fn stl_robust_decomposition() {
let period = 12;
let mut series = generate_seasonal_series(120, period);
series[30] = 100.0;
series[60] = -100.0;
let stl = STL::new(period).robust();
let result = stl.decompose(&series).unwrap();
let strength = result.seasonal_strength();
assert!(
strength > 0.1,
"Robust STL should still detect seasonality: {}",
strength
);
}
#[test]
fn stl_custom_smoothness() {
let period = 12;
let series = generate_seasonal_series(120, period);
let stl = STL::new(period)
.with_seasonal_smoothness(7)
.with_trend_smoothness(21)
.with_inner_iterations(3);
let result = stl.decompose(&series).unwrap();
assert_eq!(result.trend.len(), series.len());
}
#[test]
fn stl_different_periods() {
let series_weekly = generate_seasonal_series(70, 7);
let stl_weekly = STL::new(7);
assert!(stl_weekly.decompose(&series_weekly).is_some());
let series_quarterly = generate_seasonal_series(40, 4);
let stl_quarterly = STL::new(4);
assert!(stl_quarterly.decompose(&series_quarterly).is_some());
}
#[test]
fn stl_result_seasonal_strength_range() {
let period = 12;
let series = generate_seasonal_series(120, period);
let stl = STL::new(period);
let result = stl.decompose(&series).unwrap();
let strength = result.seasonal_strength();
assert!(
(0.0..=1.0).contains(&strength),
"Seasonal strength should be in [0, 1]: {}",
strength
);
}
#[test]
fn stl_result_trend_strength_range() {
let period = 12;
let series = generate_seasonal_series(120, period);
let stl = STL::new(period);
let result = stl.decompose(&series).unwrap();
let strength = result.trend_strength();
assert!(
(0.0..=1.0).contains(&strength),
"Trend strength should be in [0, 1]: {}",
strength
);
}
}