use crate::core::TimeSeries;
use crate::error::{ForecastError, Result};
use crate::models::Forecaster;
use crate::utils::metrics::{calculate_metrics, AccuracyMetrics};
#[cfg(feature = "parallel")]
use rayon::prelude::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum CVStrategy {
Rolling,
#[default]
Expanding,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Fold {
pub train_start: usize,
pub train_end: usize,
pub test_start: usize,
pub test_end: usize,
}
impl Fold {
pub fn train_size(&self) -> usize {
self.train_end - self.train_start
}
pub fn test_size(&self) -> usize {
self.test_end - self.test_start
}
}
#[derive(Debug, Clone)]
pub struct CvFoldGenerator {
pub initial_window: usize,
pub horizon: usize,
pub step_size: usize,
pub gap: usize,
pub purge: usize,
pub strategy: CVStrategy,
}
impl Default for CvFoldGenerator {
fn default() -> Self {
Self {
initial_window: 10,
horizon: 1,
step_size: 1,
gap: 0,
purge: 0,
strategy: CVStrategy::Expanding,
}
}
}
impl CvFoldGenerator {
pub fn new() -> Self {
Self::default()
}
pub fn initial_window(mut self, size: usize) -> Self {
self.initial_window = size;
self
}
pub fn horizon(mut self, h: usize) -> Self {
self.horizon = h;
self
}
pub fn step_size(mut self, step: usize) -> Self {
self.step_size = step;
self
}
pub fn gap(mut self, g: usize) -> Self {
self.gap = g;
self
}
pub fn purge(mut self, p: usize) -> Self {
self.purge = p;
self
}
pub fn strategy(mut self, s: CVStrategy) -> Self {
self.strategy = s;
self
}
pub fn generate(&self, series_len: usize) -> Vec<Fold> {
let mut folds = Vec::new();
let mut origin = self.initial_window;
while origin + self.gap + self.horizon <= series_len {
let train_start = match self.strategy {
CVStrategy::Rolling => origin.saturating_sub(self.initial_window),
CVStrategy::Expanding => 0,
};
let train_end = origin.saturating_sub(self.purge);
if train_end <= train_start {
origin += self.step_size;
continue;
}
let test_start = origin + self.gap;
let test_end = test_start + self.horizon;
folds.push(Fold {
train_start,
train_end,
test_start,
test_end,
});
origin += self.step_size;
}
folds
}
pub fn n_folds(&self, series_len: usize) -> usize {
self.generate(series_len).len()
}
}
#[derive(Debug, Clone)]
pub struct CVConfig {
pub horizon: usize,
pub initial_window: usize,
pub step_size: usize,
pub strategy: CVStrategy,
pub seasonal_period: Option<usize>,
pub gap: usize,
pub purge: usize,
}
impl Default for CVConfig {
fn default() -> Self {
Self {
horizon: 1,
initial_window: 10,
step_size: 1,
strategy: CVStrategy::Expanding,
seasonal_period: None,
gap: 0,
purge: 0,
}
}
}
impl CVConfig {
pub fn expanding(initial_window: usize, horizon: usize) -> Self {
Self {
initial_window,
horizon,
step_size: 1,
strategy: CVStrategy::Expanding,
seasonal_period: None,
gap: 0,
purge: 0,
}
}
pub fn rolling(window_size: usize, horizon: usize) -> Self {
Self {
initial_window: window_size,
horizon,
step_size: 1,
strategy: CVStrategy::Rolling,
seasonal_period: None,
gap: 0,
purge: 0,
}
}
pub fn with_step_size(mut self, step_size: usize) -> Self {
self.step_size = step_size;
self
}
pub fn with_seasonal_period(mut self, period: usize) -> Self {
self.seasonal_period = Some(period);
self
}
pub fn with_gap(mut self, gap: usize) -> Self {
self.gap = gap;
self
}
pub fn with_purge(mut self, purge: usize) -> Self {
self.purge = purge;
self
}
pub fn to_fold_generator(&self) -> CvFoldGenerator {
CvFoldGenerator {
initial_window: self.initial_window,
horizon: self.horizon,
step_size: self.step_size,
gap: self.gap,
purge: self.purge,
strategy: self.strategy,
}
}
}
#[derive(Debug, Clone)]
pub struct CVResults {
pub n_folds: usize,
pub aggregated: AggregatedMetrics,
pub fold_metrics: Vec<AccuracyMetrics>,
pub actual_values: Vec<f64>,
pub predicted_values: Vec<f64>,
pub folds: Vec<Fold>,
}
#[derive(Debug, Clone)]
pub struct AggregatedMetrics {
pub mae: f64,
pub rmse: f64,
pub smape: f64,
pub mape: Option<f64>,
pub mae_std: f64,
pub rmse_std: f64,
}
pub trait FillStrategy: Clone {
fn fill(&self, train_values: &[f64], test_len: usize) -> Vec<f64>;
}
#[derive(Debug, Clone, Default)]
pub struct LastValueFill;
impl FillStrategy for LastValueFill {
fn fill(&self, train_values: &[f64], test_len: usize) -> Vec<f64> {
let last = train_values.last().copied().unwrap_or(0.0);
vec![last; test_len]
}
}
#[derive(Debug, Clone, Default)]
pub struct MeanFill;
impl FillStrategy for MeanFill {
fn fill(&self, train_values: &[f64], test_len: usize) -> Vec<f64> {
if train_values.is_empty() {
return vec![0.0; test_len];
}
let mean = train_values.iter().sum::<f64>() / train_values.len() as f64;
vec![mean; test_len]
}
}
#[derive(Debug, Clone, Default)]
pub struct MedianFill;
impl FillStrategy for MedianFill {
fn fill(&self, train_values: &[f64], test_len: usize) -> Vec<f64> {
if train_values.is_empty() {
return vec![0.0; test_len];
}
let mut sorted: Vec<f64> = train_values
.iter()
.filter(|x| x.is_finite())
.copied()
.collect();
if sorted.is_empty() {
return vec![0.0; test_len];
}
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
let median = if sorted.len() % 2 == 0 {
(sorted[sorted.len() / 2 - 1] + sorted[sorted.len() / 2]) / 2.0
} else {
sorted[sorted.len() / 2]
};
vec![median; test_len]
}
}
#[derive(Debug, Clone, Default)]
pub struct ZeroFill;
impl FillStrategy for ZeroFill {
fn fill(&self, _train_values: &[f64], test_len: usize) -> Vec<f64> {
vec![0.0; test_len]
}
}
#[derive(Debug, Clone)]
pub struct ConstantFill(pub f64);
impl FillStrategy for ConstantFill {
fn fill(&self, _train_values: &[f64], test_len: usize) -> Vec<f64> {
vec![self.0; test_len]
}
}
#[derive(Debug, Clone, Default)]
pub struct ModeFill;
impl FillStrategy for ModeFill {
fn fill(&self, train_values: &[f64], test_len: usize) -> Vec<f64> {
if train_values.is_empty() {
return vec![0.0; test_len];
}
use std::collections::HashMap;
let mut counts: HashMap<i64, usize> = HashMap::new();
for &v in train_values {
if v.is_finite() {
let key = (v * 1_000_000.0).round() as i64;
*counts.entry(key).or_insert(0) += 1;
}
}
let mode = counts
.into_iter()
.max_by_key(|&(_, count)| count)
.map(|(key, _)| key as f64 / 1_000_000.0)
.unwrap_or(0.0);
vec![mode; test_len]
}
}
pub fn train_test_split(series: &TimeSeries, split_point: f64) -> Result<(TimeSeries, TimeSeries)> {
let n = series.len();
if n < 2 {
return Err(ForecastError::InvalidParameter(
"Series must have at least 2 observations for train/test split".to_string(),
));
}
let split_idx = if split_point > 0.0 && split_point < 1.0 {
(n as f64 * split_point).round() as usize
} else if split_point >= 1.0 && split_point < n as f64 {
split_point as usize
} else {
return Err(ForecastError::InvalidParameter(format!(
"split_point must be a ratio (0.0-1.0) or index (1 to {}), got {}",
n - 1,
split_point
)));
};
let split_idx = split_idx.clamp(1, n - 1);
let train = series.slice(0, split_idx)?;
let test = series.slice(split_idx, n)?;
Ok((train, test))
}
pub fn train_test_split_at(series: &TimeSeries, index: usize) -> Result<(TimeSeries, TimeSeries)> {
let n = series.len();
if index == 0 || index >= n {
return Err(ForecastError::InvalidParameter(format!(
"split index must be between 1 and {}, got {}",
n - 1,
index
)));
}
let train = series.slice(0, index)?;
let test = series.slice(index, n)?;
Ok((train, test))
}
fn evaluate_fold<F: Forecaster>(
series: &TimeSeries,
fold: &Fold,
model_factory: &dyn Fn() -> F,
seasonal_period: Option<usize>,
) -> Result<(AccuracyMetrics, Vec<f64>, Vec<f64>)> {
let train_series = series.slice(fold.train_start, fold.train_end)?;
let mut model = model_factory();
model.fit(&train_series)?;
let forecast = model.predict(fold.test_size())?;
let predicted: Vec<f64> = forecast.primary().to_vec();
let actual: Vec<f64> = (fold.test_start..fold.test_end)
.map(|i| series.primary_values()[i])
.collect();
let metrics = calculate_metrics(&actual, &predicted, seasonal_period)?;
Ok((metrics, actual, predicted))
}
pub fn cross_validate<F, Factory>(
config: &CVConfig,
series: &TimeSeries,
model_factory: Factory,
) -> Result<CVResults>
where
F: Forecaster + Send,
Factory: Fn() -> F + Sync,
{
let generator = config.to_fold_generator();
let folds = generator.generate(series.len());
if folds.is_empty() {
return Ok(CVResults {
n_folds: 0,
aggregated: AggregatedMetrics {
mae: f64::NAN,
rmse: f64::NAN,
smape: f64::NAN,
mape: None,
mae_std: f64::NAN,
rmse_std: f64::NAN,
},
fold_metrics: vec![],
actual_values: vec![],
predicted_values: vec![],
folds: vec![],
});
}
let seasonal_period = config.seasonal_period;
#[cfg(feature = "parallel")]
let fold_results: Vec<Result<(AccuracyMetrics, Vec<f64>, Vec<f64>)>> = folds
.par_iter()
.map(|fold| evaluate_fold(series, fold, &model_factory, seasonal_period))
.collect();
#[cfg(not(feature = "parallel"))]
let fold_results: Vec<Result<(AccuracyMetrics, Vec<f64>, Vec<f64>)>> = folds
.iter()
.map(|fold| evaluate_fold(series, fold, &model_factory, seasonal_period))
.collect();
let mut fold_metrics = Vec::with_capacity(folds.len());
let mut all_actual = Vec::new();
let mut all_predicted = Vec::new();
for result in fold_results {
let (metrics, actual, predicted) = result?;
fold_metrics.push(metrics);
all_actual.extend_from_slice(&actual);
all_predicted.extend_from_slice(&predicted);
}
let n_folds = fold_metrics.len();
let mae_values: Vec<f64> = fold_metrics.iter().map(|m| m.mae).collect();
let rmse_values: Vec<f64> = fold_metrics.iter().map(|m| m.rmse).collect();
let smape_values: Vec<f64> = fold_metrics.iter().map(|m| m.smape).collect();
let mae_mean = mae_values.iter().sum::<f64>() / n_folds as f64;
let rmse_mean = rmse_values.iter().sum::<f64>() / n_folds as f64;
let smape_mean = smape_values.iter().sum::<f64>() / n_folds as f64;
let mae_std = std_dev(&mae_values);
let rmse_std = std_dev(&rmse_values);
let mape = if fold_metrics.iter().all(|m| m.mape.is_some()) {
let mape_values: Vec<f64> = fold_metrics.iter().filter_map(|m| m.mape).collect();
Some(mape_values.iter().sum::<f64>() / n_folds as f64)
} else {
None
};
Ok(CVResults {
n_folds,
aggregated: AggregatedMetrics {
mae: mae_mean,
rmse: rmse_mean,
smape: smape_mean,
mape,
mae_std,
rmse_std,
},
fold_metrics,
actual_values: all_actual,
predicted_values: all_predicted,
folds,
})
}
#[derive(Debug, Clone)]
pub struct GroupedCVResults {
pub group_results: Vec<(String, CVResults)>,
pub aggregated: AggregatedMetrics,
pub folds: Vec<Fold>,
}
pub fn grouped_cross_validate<F, Factory, I>(
config: &CVConfig,
series_map: I,
model_factory: Factory,
) -> Result<GroupedCVResults>
where
F: Forecaster + Send,
Factory: Fn() -> F + Sync,
I: IntoIterator<Item = (String, TimeSeries)>,
{
let series_vec: Vec<(String, TimeSeries)> = series_map.into_iter().collect();
if series_vec.is_empty() {
return Err(ForecastError::InvalidParameter(
"No series provided for grouped cross-validation".to_string(),
));
}
let min_len = series_vec.iter().map(|(_, s)| s.len()).min().unwrap_or(0);
let generator = config.to_fold_generator();
let folds = generator.generate(min_len);
if folds.is_empty() {
return Err(ForecastError::InvalidParameter(
"Not enough data for any CV folds".to_string(),
));
}
let mut group_results = Vec::with_capacity(series_vec.len());
let mut all_mae = Vec::new();
let mut all_rmse = Vec::new();
let mut all_smape = Vec::new();
let mut all_mape = Vec::new();
for (group_id, series) in series_vec {
let cv_result = cross_validate_with_folds(config, &series, &folds, &model_factory)?;
all_mae.push(cv_result.aggregated.mae);
all_rmse.push(cv_result.aggregated.rmse);
all_smape.push(cv_result.aggregated.smape);
if let Some(mape) = cv_result.aggregated.mape {
all_mape.push(mape);
}
group_results.push((group_id, cv_result));
}
let n_groups = group_results.len() as f64;
let aggregated = AggregatedMetrics {
mae: all_mae.iter().sum::<f64>() / n_groups,
rmse: all_rmse.iter().sum::<f64>() / n_groups,
smape: all_smape.iter().sum::<f64>() / n_groups,
mape: if all_mape.len() == group_results.len() {
Some(all_mape.iter().sum::<f64>() / n_groups)
} else {
None
},
mae_std: std_dev(&all_mae),
rmse_std: std_dev(&all_rmse),
};
Ok(GroupedCVResults {
group_results,
aggregated,
folds,
})
}
fn cross_validate_with_folds<F, Factory>(
config: &CVConfig,
series: &TimeSeries,
folds: &[Fold],
model_factory: &Factory,
) -> Result<CVResults>
where
F: Forecaster,
Factory: Fn() -> F,
{
if folds.is_empty() {
return Ok(CVResults {
n_folds: 0,
aggregated: AggregatedMetrics {
mae: f64::NAN,
rmse: f64::NAN,
smape: f64::NAN,
mape: None,
mae_std: f64::NAN,
rmse_std: f64::NAN,
},
fold_metrics: vec![],
actual_values: vec![],
predicted_values: vec![],
folds: vec![],
});
}
let mut fold_metrics = Vec::with_capacity(folds.len());
let mut all_actual = Vec::new();
let mut all_predicted = Vec::new();
for fold in folds {
let train_series = series.slice(fold.train_start, fold.train_end)?;
let mut model = model_factory();
model.fit(&train_series)?;
let forecast = model.predict(fold.test_size())?;
let predictions = forecast.primary();
let actual: Vec<f64> = (fold.test_start..fold.test_end)
.map(|i| series.primary_values()[i])
.collect();
let metrics = calculate_metrics(&actual, predictions, config.seasonal_period)?;
fold_metrics.push(metrics);
all_actual.extend_from_slice(&actual);
all_predicted.extend_from_slice(predictions);
}
let n_folds = fold_metrics.len();
let mae_values: Vec<f64> = fold_metrics.iter().map(|m| m.mae).collect();
let rmse_values: Vec<f64> = fold_metrics.iter().map(|m| m.rmse).collect();
let smape_values: Vec<f64> = fold_metrics.iter().map(|m| m.smape).collect();
let mae_mean = mae_values.iter().sum::<f64>() / n_folds as f64;
let rmse_mean = rmse_values.iter().sum::<f64>() / n_folds as f64;
let smape_mean = smape_values.iter().sum::<f64>() / n_folds as f64;
let mae_std = std_dev(&mae_values);
let rmse_std = std_dev(&rmse_values);
let mape = if fold_metrics.iter().all(|m| m.mape.is_some()) {
let mape_values: Vec<f64> = fold_metrics.iter().filter_map(|m| m.mape).collect();
Some(mape_values.iter().sum::<f64>() / n_folds as f64)
} else {
None
};
Ok(CVResults {
n_folds,
aggregated: AggregatedMetrics {
mae: mae_mean,
rmse: rmse_mean,
smape: smape_mean,
mape,
mae_std,
rmse_std,
},
fold_metrics,
actual_values: all_actual,
predicted_values: all_predicted,
folds: folds.to_vec(),
})
}
fn std_dev(values: &[f64]) -> f64 {
if values.len() < 2 {
return 0.0;
}
let mean = values.iter().sum::<f64>() / values.len() as f64;
let variance =
values.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / (values.len() - 1) as f64;
variance.sqrt()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::models::baseline::{Naive, SimpleMovingAverage};
use approx::assert_relative_eq;
use chrono::{TimeZone, Utc};
fn make_timestamps(n: usize) -> Vec<chrono::DateTime<Utc>> {
use chrono::Duration;
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()
}
#[test]
fn fold_generator_basic() {
let gen = CvFoldGenerator::new()
.initial_window(10)
.horizon(1)
.step_size(1);
let folds = gen.generate(20);
assert_eq!(folds.len(), 10);
assert_eq!(folds[0].train_start, 0);
assert_eq!(folds[0].train_end, 10);
assert_eq!(folds[0].test_start, 10);
assert_eq!(folds[0].test_end, 11);
}
#[test]
fn fold_generator_with_gap() {
let gen = CvFoldGenerator::new()
.initial_window(10)
.horizon(1)
.step_size(1)
.gap(2);
let folds = gen.generate(20);
assert_eq!(folds[0].train_end, 10);
assert_eq!(folds[0].test_start, 12); assert_eq!(folds[0].test_end, 13);
assert_eq!(folds.len(), 8);
}
#[test]
fn fold_generator_with_purge() {
let gen = CvFoldGenerator::new()
.initial_window(10)
.horizon(1)
.step_size(1)
.purge(2);
let folds = gen.generate(20);
assert_eq!(folds[0].train_end, 8); assert_eq!(folds[0].test_start, 10);
}
#[test]
fn fold_generator_rolling() {
let gen = CvFoldGenerator::new()
.initial_window(5)
.horizon(1)
.step_size(1)
.strategy(CVStrategy::Rolling);
let folds = gen.generate(15);
assert_eq!(folds[0].train_start, 0);
assert_eq!(folds[0].train_end, 5);
assert_eq!(folds[5].train_start, 5);
assert_eq!(folds[5].train_end, 10);
}
#[test]
fn fold_generator_multi_step_horizon() {
let gen = CvFoldGenerator::new()
.initial_window(10)
.horizon(3)
.step_size(2);
let folds = gen.generate(20);
assert_eq!(folds[0].test_start, 10);
assert_eq!(folds[0].test_end, 13);
assert_eq!(folds[0].test_size(), 3);
}
#[test]
fn fold_generator_insufficient_data() {
let gen = CvFoldGenerator::new().initial_window(10).horizon(5);
let folds = gen.generate(10); assert!(folds.is_empty());
}
#[test]
fn train_test_split_by_ratio() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let (train, test) = train_test_split(&ts, 0.8).unwrap();
assert_eq!(train.len(), 80);
assert_eq!(test.len(), 20);
assert_relative_eq!(train.primary_values()[0], 0.0);
assert_relative_eq!(train.primary_values()[79], 79.0);
assert_relative_eq!(test.primary_values()[0], 80.0);
}
#[test]
fn train_test_split_at_index() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let (train, test) = train_test_split_at(&ts, 70).unwrap();
assert_eq!(train.len(), 70);
assert_eq!(test.len(), 30);
}
#[test]
fn train_test_split_edge_cases() {
let timestamps = make_timestamps(10);
let values: Vec<f64> = (0..10).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let (train, test) = train_test_split(&ts, 0.1).unwrap();
assert!(train.len() >= 1);
assert!(test.len() >= 1);
let (train, test) = train_test_split(&ts, 0.99).unwrap();
assert!(train.len() >= 1);
assert!(test.len() >= 1);
}
#[test]
fn train_test_split_invalid() {
let timestamps = make_timestamps(10);
let values: Vec<f64> = (0..10).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
assert!(train_test_split_at(&ts, 0).is_err());
assert!(train_test_split_at(&ts, 10).is_err());
}
#[test]
fn fill_strategy_last_value() {
let train = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let fill = LastValueFill;
let result = fill.fill(&train, 3);
assert_eq!(result, vec![5.0, 5.0, 5.0]);
}
#[test]
fn fill_strategy_mean() {
let train = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let fill = MeanFill;
let result = fill.fill(&train, 3);
assert_eq!(result, vec![3.0, 3.0, 3.0]);
}
#[test]
fn fill_strategy_median() {
let train = vec![1.0, 2.0, 3.0, 4.0, 5.0];
let fill = MedianFill;
let result = fill.fill(&train, 3);
assert_eq!(result, vec![3.0, 3.0, 3.0]);
let train_even = vec![1.0, 2.0, 3.0, 4.0];
let result_even = fill.fill(&train_even, 2);
assert_eq!(result_even, vec![2.5, 2.5]);
}
#[test]
fn fill_strategy_zero() {
let train = vec![1.0, 2.0, 3.0];
let fill = ZeroFill;
let result = fill.fill(&train, 5);
assert_eq!(result, vec![0.0; 5]);
}
#[test]
fn fill_strategy_constant() {
let train = vec![1.0, 2.0, 3.0];
let fill = ConstantFill(42.0);
let result = fill.fill(&train, 3);
assert_eq!(result, vec![42.0, 42.0, 42.0]);
}
#[test]
fn fill_strategy_mode() {
let train = vec![1.0, 2.0, 2.0, 3.0, 2.0, 4.0];
let fill = ModeFill;
let result = fill.fill(&train, 3);
assert_eq!(result, vec![2.0, 2.0, 2.0]);
}
#[test]
fn fill_strategy_empty_input() {
let train: Vec<f64> = vec![];
assert_eq!(LastValueFill.fill(&train, 2), vec![0.0, 0.0]);
assert_eq!(MeanFill.fill(&train, 2), vec![0.0, 0.0]);
assert_eq!(MedianFill.fill(&train, 2), vec![0.0, 0.0]);
assert_eq!(ModeFill.fill(&train, 2), vec![0.0, 0.0]);
}
#[test]
fn cv_config_with_gap() {
let config = CVConfig::expanding(10, 1).with_gap(3);
assert_eq!(config.gap, 3);
let gen = config.to_fold_generator();
assert_eq!(gen.gap, 3);
}
#[test]
fn cv_config_with_purge() {
let config = CVConfig::expanding(10, 1).with_purge(2);
assert_eq!(config.purge, 2);
let gen = config.to_fold_generator();
assert_eq!(gen.purge, 2);
}
#[test]
fn cv_expanding_window_basic() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
assert_eq!(results.n_folds, 10);
assert!(results.aggregated.mae.is_finite());
assert_eq!(results.folds.len(), 10);
}
#[test]
fn cv_rolling_window_basic() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::rolling(10, 1);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
assert_eq!(results.n_folds, 10);
assert!(results.aggregated.mae.is_finite());
}
#[test]
fn cv_with_gap_reduces_folds() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config_no_gap = CVConfig::expanding(10, 1);
let config_gap = CVConfig::expanding(10, 1).with_gap(3);
let results_no_gap = cross_validate(&config_no_gap, &ts, Naive::new).unwrap();
let results_gap = cross_validate(&config_gap, &ts, Naive::new).unwrap();
assert!(results_gap.n_folds < results_no_gap.n_folds);
}
#[test]
fn cv_with_purge() {
let timestamps = make_timestamps(25);
let values: Vec<f64> = (0..25).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1).with_purge(2);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
let first_fold = &results.folds[0];
assert_eq!(first_fold.train_end, 8); }
#[test]
fn cv_with_step_size() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1).with_step_size(2);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
assert_eq!(results.n_folds, 5);
}
#[test]
fn cv_multi_step_horizon() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 3);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
assert_eq!(results.n_folds, 8);
assert_eq!(results.actual_values.len(), 8 * 3);
assert_eq!(results.predicted_values.len(), 8 * 3);
}
#[test]
fn cv_insufficient_data_returns_zero_folds() {
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 config = CVConfig::expanding(10, 1);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
assert_eq!(results.n_folds, 0);
}
#[test]
fn cv_naive_perfect_on_constant() {
let timestamps = make_timestamps(20);
let values = vec![5.0; 20]; let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
assert_relative_eq!(results.aggregated.mae, 0.0, epsilon = 1e-10);
}
#[test]
fn cv_sma_on_linear_trend() {
let timestamps = make_timestamps(30);
let values: Vec<f64> = (0..30).map(|i| 10.0 + i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(15, 1);
let results = cross_validate(&config, &ts, || SimpleMovingAverage::new(5)).unwrap();
assert!(results.aggregated.mae > 0.0);
assert!(results.aggregated.rmse >= results.aggregated.mae);
}
#[test]
fn cv_metrics_are_consistent() {
let timestamps = make_timestamps(25);
let values: Vec<f64> = (0..25).map(|i| (i as f64).sin() * 10.0 + 50.0).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(15, 1);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
assert!(results.aggregated.rmse >= results.aggregated.mae);
assert!(results.aggregated.smape >= 0.0);
assert!(results.aggregated.smape <= 200.0);
assert!(results.aggregated.mae_std >= 0.0);
}
#[test]
fn cv_fold_metrics_match_aggregated() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20).map(|i| i as f64 + 0.1 * (i as f64).sin()).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
let manual_mae_mean: f64 =
results.fold_metrics.iter().map(|m| m.mae).sum::<f64>() / results.n_folds as f64;
assert_relative_eq!(results.aggregated.mae, manual_mae_mean, epsilon = 1e-10);
}
#[test]
fn cv_with_seasonal_period() {
let timestamps = make_timestamps(30);
let values: Vec<f64> = (0..30)
.map(|i| ((i % 4) as f64) * 10.0 + 5.0 + 0.5 * (i as f64))
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(12, 5)
.with_seasonal_period(4)
.with_step_size(3);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
let mase_count = results
.fold_metrics
.iter()
.filter(|m| m.mase.is_some())
.count();
assert!(mase_count > 0);
}
#[test]
fn cv_config_builders() {
let expanding = CVConfig::expanding(10, 3);
assert_eq!(expanding.initial_window, 10);
assert_eq!(expanding.horizon, 3);
assert_eq!(expanding.strategy, CVStrategy::Expanding);
let rolling = CVConfig::rolling(15, 2);
assert_eq!(rolling.initial_window, 15);
assert_eq!(rolling.horizon, 2);
assert_eq!(rolling.strategy, CVStrategy::Rolling);
let with_step = CVConfig::expanding(10, 1).with_step_size(5);
assert_eq!(with_step.step_size, 5);
let with_seasonal = CVConfig::expanding(10, 1).with_seasonal_period(12);
assert_eq!(with_seasonal.seasonal_period, Some(12));
}
#[test]
fn cv_default_config() {
let config = CVConfig::default();
assert_eq!(config.horizon, 1);
assert_eq!(config.initial_window, 10);
assert_eq!(config.step_size, 1);
assert_eq!(config.strategy, CVStrategy::Expanding);
assert_eq!(config.seasonal_period, None);
assert_eq!(config.gap, 0);
assert_eq!(config.purge, 0);
}
#[test]
fn cv_values_stored_correctly() {
let timestamps = make_timestamps(15);
let values: Vec<f64> = (0..15).map(|i| i as f64 * 2.0).collect();
let ts = TimeSeries::univariate(timestamps, values.clone()).unwrap();
let config = CVConfig::expanding(10, 2).with_step_size(2);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
for &actual in &results.actual_values {
assert!(values.iter().any(|&v| (v - actual).abs() < 1e-10));
}
}
#[test]
fn grouped_cv_basic() {
let timestamps = make_timestamps(30);
let series_a =
TimeSeries::univariate(timestamps.clone(), (0..30).map(|i| i as f64).collect())
.unwrap();
let series_b = TimeSeries::univariate(
timestamps.clone(),
(0..30).map(|i| (i as f64) * 2.0).collect(),
)
.unwrap();
let series_map = vec![
("product_a".to_string(), series_a),
("product_b".to_string(), series_b),
];
let config = CVConfig::expanding(15, 3).with_step_size(3);
let results = grouped_cross_validate(&config, series_map.into_iter(), Naive::new).unwrap();
assert_eq!(results.group_results.len(), 2);
assert!(results.aggregated.mae.is_finite());
let n_folds_a = results.group_results[0].1.n_folds;
let n_folds_b = results.group_results[1].1.n_folds;
assert_eq!(n_folds_a, n_folds_b);
}
#[test]
fn grouped_cv_uses_min_length() {
let timestamps_short = make_timestamps(20);
let timestamps_long = make_timestamps(30);
let series_short =
TimeSeries::univariate(timestamps_short, (0..20).map(|i| i as f64).collect()).unwrap();
let series_long =
TimeSeries::univariate(timestamps_long, (0..30).map(|i| i as f64).collect()).unwrap();
let series_map = vec![
("short".to_string(), series_short),
("long".to_string(), series_long),
];
let config = CVConfig::expanding(10, 1);
let results = grouped_cross_validate(&config, series_map.into_iter(), Naive::new).unwrap();
for (_, cv_result) in &results.group_results {
assert_eq!(cv_result.n_folds, 10); }
}
#[test]
fn grouped_cv_empty_input() {
let series_map: Vec<(String, TimeSeries)> = vec![];
let config = CVConfig::expanding(10, 1);
let result = grouped_cross_validate(&config, series_map.into_iter(), Naive::new);
assert!(result.is_err());
}
}