use crate::error::{ForecastError, Result};
use crate::seasonality::traits::SeasonalComponent;
#[derive(Debug, Clone)]
pub struct DummySeasonality {
seasonal_means: Option<Vec<f64>>,
fitted: Vec<f64>,
period: Option<usize>,
}
impl DummySeasonality {
pub fn new() -> Self {
Self {
seasonal_means: None,
fitted: Vec::new(),
period: None,
}
}
pub fn seasonal_means(&self) -> Option<&[f64]> {
self.seasonal_means.as_deref()
}
pub fn period(&self) -> Option<usize> {
self.period
}
}
impl Default for DummySeasonality {
fn default() -> Self {
Self::new()
}
}
impl SeasonalComponent for DummySeasonality {
fn fit_seasonal(&mut self, values: &[f64], period: usize) -> Result<()> {
if period < 2 {
return Err(ForecastError::InvalidParameter(format!(
"seasonal period must be >= 2, got {}",
period
)));
}
if values.is_empty() {
return Err(ForecastError::EmptyData);
}
if values.len() < period {
return Err(ForecastError::InsufficientData {
needed: period,
got: values.len(),
hint: Some(format!(
"need at least {} data points for period {}",
period, period
)),
});
}
let mut sums = vec![0.0; period];
let mut counts = vec![0usize; period];
for (i, &v) in values.iter().enumerate() {
let pos = i % period;
sums[pos] += v;
counts[pos] += 1;
}
let seasonal_means: Vec<f64> = sums
.iter()
.zip(counts.iter())
.map(|(&s, &c)| s / c as f64)
.collect();
let fitted: Vec<f64> = (0..values.len())
.map(|i| seasonal_means[i % period])
.collect();
self.seasonal_means = Some(seasonal_means);
self.fitted = fitted;
self.period = Some(period);
Ok(())
}
fn fitted_seasonal(&self) -> &[f64] {
&self.fitted
}
fn predict_seasonal(&self, n_ahead: usize) -> Vec<f64> {
let means = match &self.seasonal_means {
Some(m) => m,
None => return Vec::new(),
};
let period = means.len();
let n_train = self.fitted.len();
(0..n_ahead)
.map(|i| {
let pos = (n_train + i) % period;
means[pos]
})
.collect()
}
fn seasonal_features(&self) -> Vec<(&str, f64)> {
let means = match &self.seasonal_means {
Some(m) => m,
None => return Vec::new(),
};
let overall_mean: f64 = self.fitted.iter().copied().sum::<f64>() / self.fitted.len() as f64;
let ss_seasonal: f64 = self
.fitted
.iter()
.map(|&f| (f - overall_mean).powi(2))
.sum();
let strength = if ss_seasonal > 0.0 {
1.0 } else {
0.0
};
let max_mean = means.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let min_mean = means.iter().copied().fold(f64::INFINITY, f64::min);
let amplitude = max_mean - min_mean;
let peak_pos = means
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
.map(|(i, _)| i)
.unwrap_or(0);
let trough_pos = means
.iter()
.enumerate()
.min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
.map(|(i, _)| i)
.unwrap_or(0);
vec![
("dummy_seasonal_strength", strength),
("dummy_seasonal_amplitude", amplitude),
("dummy_seasonal_peak_position", peak_pos as f64),
("dummy_seasonal_trough_position", trough_pos as f64),
]
}
fn seasonal_name(&self) -> &str {
"DummySeasonality"
}
fn n_params(&self) -> usize {
self.period.unwrap_or(0)
}
}
pub fn dummy_seasonal_strength(values: &[f64], period: usize) -> Result<f64> {
if values.is_empty() {
return Err(ForecastError::EmptyData);
}
if values.len() < period {
return Err(ForecastError::InsufficientData {
needed: period,
got: values.len(),
hint: Some(format!(
"need at least {} data points for period {}",
period, period
)),
});
}
let n = values.len();
let overall_mean: f64 = values.iter().sum::<f64>() / n as f64;
let ss_total: f64 = values.iter().map(|&v| (v - overall_mean).powi(2)).sum();
if ss_total == 0.0 {
return Ok(0.0);
}
let mut sums = vec![0.0; period];
let mut counts = vec![0usize; period];
for (i, &v) in values.iter().enumerate() {
let pos = i % period;
sums[pos] += v;
counts[pos] += 1;
}
let means: Vec<f64> = sums
.iter()
.zip(counts.iter())
.map(|(&s, &c)| s / c as f64)
.collect();
let ss_residual: f64 = values
.iter()
.enumerate()
.map(|(i, &v)| {
let m = means[i % period];
(v - m).powi(2)
})
.sum();
Ok(1.0 - ss_residual / ss_total)
}
pub fn dummy_seasonal_amplitude(values: &[f64], period: usize) -> Result<f64> {
if values.is_empty() {
return Err(ForecastError::EmptyData);
}
if values.len() < period {
return Err(ForecastError::InsufficientData {
needed: period,
got: values.len(),
hint: Some(format!(
"need at least {} data points for period {}",
period, period
)),
});
}
let mut sums = vec![0.0; period];
let mut counts = vec![0usize; period];
for (i, &v) in values.iter().enumerate() {
let pos = i % period;
sums[pos] += v;
counts[pos] += 1;
}
let means: Vec<f64> = sums
.iter()
.zip(counts.iter())
.map(|(&s, &c)| s / c as f64)
.collect();
let max_mean = means.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let min_mean = means.iter().copied().fold(f64::INFINITY, f64::min);
Ok(max_mean - min_mean)
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_abs_diff_eq;
#[test]
fn fit_perfect_seasonal_pattern() {
let mut model = DummySeasonality::new();
let values = vec![
10.0, 20.0, 30.0, 40.0, 10.0, 20.0, 30.0, 40.0, 10.0, 20.0, 30.0, 40.0,
];
model.fit_seasonal(&values, 4).unwrap();
let means = model.seasonal_means().unwrap();
assert_eq!(means.len(), 4);
assert_abs_diff_eq!(means[0], 10.0, epsilon = 1e-10);
assert_abs_diff_eq!(means[1], 20.0, epsilon = 1e-10);
assert_abs_diff_eq!(means[2], 30.0, epsilon = 1e-10);
assert_abs_diff_eq!(means[3], 40.0, epsilon = 1e-10);
let fitted = model.fitted_seasonal();
assert_eq!(fitted.len(), 12);
for (i, &f) in fitted.iter().enumerate() {
assert_abs_diff_eq!(f, values[i], epsilon = 1e-10);
}
}
#[test]
fn fit_noisy_data() {
let mut model = DummySeasonality::new();
let values = vec![
11.0, 19.0, 31.0, 9.0, 21.0, 29.0, 10.0, 20.0, 30.0, ];
model.fit_seasonal(&values, 3).unwrap();
let means = model.seasonal_means().unwrap();
assert_eq!(means.len(), 3);
assert_abs_diff_eq!(means[0], 10.0, epsilon = 1e-10);
assert_abs_diff_eq!(means[1], 20.0, epsilon = 1e-10);
assert_abs_diff_eq!(means[2], 30.0, epsilon = 1e-10);
}
#[test]
fn predict_continues_pattern() {
let mut model = DummySeasonality::new();
let values = vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0];
model.fit_seasonal(&values, 3).unwrap();
let forecast = model.predict_seasonal(7);
assert_eq!(forecast.len(), 7);
assert_abs_diff_eq!(forecast[0], 1.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[1], 2.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[2], 3.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[3], 1.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[4], 2.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[5], 3.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[6], 1.0, epsilon = 1e-10);
}
#[test]
fn predict_continues_from_correct_position() {
let mut model = DummySeasonality::new();
let values = vec![10.0, 20.0, 30.0, 10.0, 20.0];
model.fit_seasonal(&values, 3).unwrap();
let forecast = model.predict_seasonal(4);
assert_abs_diff_eq!(forecast[0], 30.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[1], 10.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[2], 20.0, epsilon = 1e-10);
assert_abs_diff_eq!(forecast[3], 30.0, epsilon = 1e-10);
}
#[test]
fn predict_zero_ahead() {
let mut model = DummySeasonality::new();
let values = vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0];
model.fit_seasonal(&values, 3).unwrap();
let forecast = model.predict_seasonal(0);
assert!(forecast.is_empty());
}
#[test]
fn features_extraction() {
let mut model = DummySeasonality::new();
let values = vec![
10.0, 20.0, 50.0, 30.0, 10.0, 20.0, 50.0, 30.0, ];
model.fit_seasonal(&values, 4).unwrap();
let features = model.seasonal_features();
assert_eq!(features.len(), 4);
let feature_map: std::collections::HashMap<&str, f64> = features.into_iter().collect();
assert_abs_diff_eq!(
*feature_map.get("dummy_seasonal_strength").unwrap(),
1.0,
epsilon = 1e-10
);
assert_abs_diff_eq!(
*feature_map.get("dummy_seasonal_amplitude").unwrap(),
40.0,
epsilon = 1e-10
);
assert_abs_diff_eq!(
*feature_map.get("dummy_seasonal_peak_position").unwrap(),
2.0,
epsilon = 1e-10
);
assert_abs_diff_eq!(
*feature_map.get("dummy_seasonal_trough_position").unwrap(),
0.0,
epsilon = 1e-10
);
}
#[test]
fn features_before_fit_returns_empty() {
let model = DummySeasonality::new();
let features = model.seasonal_features();
assert!(features.is_empty());
}
#[test]
fn standalone_strength_perfect_pattern() {
let values = vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0];
let r2 = dummy_seasonal_strength(&values, 3).unwrap();
assert_abs_diff_eq!(r2, 1.0, epsilon = 1e-10);
}
#[test]
fn standalone_strength_no_pattern() {
let values = vec![5.0, 5.0, 5.0, 5.0, 5.0, 5.0];
let r2 = dummy_seasonal_strength(&values, 3).unwrap();
assert_abs_diff_eq!(r2, 0.0, epsilon = 1e-10);
}
#[test]
fn standalone_strength_partial_pattern() {
let values = vec![
10.0, 20.0, 30.0, 12.0, 18.0, 32.0, ];
let r2 = dummy_seasonal_strength(&values, 3).unwrap();
assert!(r2 > 0.9, "R² = {} should be > 0.9", r2);
assert!(r2 < 1.0, "R² = {} should be < 1.0", r2);
}
#[test]
fn standalone_amplitude_basic() {
let values = vec![5.0, 15.0, 10.0, 5.0, 15.0, 10.0];
let amp = dummy_seasonal_amplitude(&values, 3).unwrap();
assert_abs_diff_eq!(amp, 10.0, epsilon = 1e-10);
}
#[test]
fn standalone_amplitude_constant() {
let values = vec![7.0, 7.0, 7.0, 7.0];
let amp = dummy_seasonal_amplitude(&values, 2).unwrap();
assert_abs_diff_eq!(amp, 0.0, epsilon = 1e-10);
}
#[test]
fn fit_empty_data() {
let mut model = DummySeasonality::new();
let result = model.fit_seasonal(&[], 4);
assert!(matches!(result, Err(ForecastError::EmptyData)));
}
#[test]
fn fit_insufficient_data() {
let mut model = DummySeasonality::new();
let values = vec![1.0, 2.0, 3.0];
let result = model.fit_seasonal(&values, 5);
assert!(matches!(
result,
Err(ForecastError::InsufficientData {
needed: 5,
got: 3,
..
})
));
}
#[test]
fn predict_before_fit_returns_empty() {
let model = DummySeasonality::new();
let forecast = model.predict_seasonal(10);
assert!(forecast.is_empty());
}
#[test]
fn fitted_before_fit_returns_empty() {
let model = DummySeasonality::new();
let fitted = model.fitted_seasonal();
assert!(fitted.is_empty());
}
#[test]
fn standalone_strength_empty_data() {
let result = dummy_seasonal_strength(&[], 3);
assert!(matches!(result, Err(ForecastError::EmptyData)));
}
#[test]
fn standalone_strength_insufficient_data() {
let result = dummy_seasonal_strength(&[1.0, 2.0], 5);
assert!(matches!(
result,
Err(ForecastError::InsufficientData {
needed: 5,
got: 2,
..
})
));
}
#[test]
fn standalone_amplitude_empty_data() {
let result = dummy_seasonal_amplitude(&[], 3);
assert!(matches!(result, Err(ForecastError::EmptyData)));
}
#[test]
fn standalone_amplitude_insufficient_data() {
let result = dummy_seasonal_amplitude(&[1.0], 4);
assert!(matches!(
result,
Err(ForecastError::InsufficientData {
needed: 4,
got: 1,
..
})
));
}
#[test]
fn seasonal_name() {
let model = DummySeasonality::new();
assert_eq!(model.seasonal_name(), "DummySeasonality");
}
#[test]
fn default_is_unfitted() {
let model = DummySeasonality::default();
assert!(model.seasonal_means().is_none());
assert!(model.period().is_none());
assert!(model.fitted_seasonal().is_empty());
}
#[test]
fn refit_overwrites_previous() {
let mut model = DummySeasonality::new();
let values1 = vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0];
model.fit_seasonal(&values1, 3).unwrap();
let means1: Vec<f64> = model.seasonal_means().unwrap().to_vec();
let values2 = vec![10.0, 20.0, 10.0, 20.0];
model.fit_seasonal(&values2, 2).unwrap();
let means2 = model.seasonal_means().unwrap();
assert_eq!(means2.len(), 2);
assert_ne!(means1.len(), means2.len());
assert_abs_diff_eq!(means2[0], 10.0, epsilon = 1e-10);
assert_abs_diff_eq!(means2[1], 20.0, epsilon = 1e-10);
assert_eq!(model.period(), Some(2));
}
#[test]
fn fit_non_exact_multiple_of_period() {
let mut model = DummySeasonality::new();
let values = vec![10.0, 20.0, 30.0, 40.0, 50.0, 60.0, 70.0];
model.fit_seasonal(&values, 3).unwrap();
let means = model.seasonal_means().unwrap();
assert_abs_diff_eq!(means[0], 40.0, epsilon = 1e-10);
assert_abs_diff_eq!(means[1], 35.0, epsilon = 1e-10);
assert_abs_diff_eq!(means[2], 45.0, epsilon = 1e-10);
let fitted = model.fitted_seasonal();
assert_eq!(fitted.len(), 7);
assert_abs_diff_eq!(fitted[0], 40.0, epsilon = 1e-10);
assert_abs_diff_eq!(fitted[3], 40.0, epsilon = 1e-10);
assert_abs_diff_eq!(fitted[6], 40.0, epsilon = 1e-10);
}
#[test]
fn fit_rejects_period_zero() {
let mut model = DummySeasonality::new();
let values = vec![1.0, 2.0, 3.0, 4.0];
let result = model.fit_seasonal(&values, 0);
assert!(matches!(
result,
Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("seasonal period must be >= 2")
));
}
#[test]
fn fit_rejects_period_one() {
let mut model = DummySeasonality::new();
let values = vec![1.0, 2.0, 3.0, 4.0];
let result = model.fit_seasonal(&values, 1);
assert!(matches!(
result,
Err(ForecastError::InvalidParameter(ref msg)) if msg.contains("seasonal period must be >= 2")
));
}
}