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)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
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 target_n_folds: usize,
pub min_initial_window: usize,
pub horizon: usize,
pub step_size: Option<usize>,
pub gap: usize,
pub purge: usize,
pub embargo: usize,
pub strategy: CVStrategy,
pub on_constraint_violation: ConstraintViolation,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConstraintViolation {
Error,
ReduceFolds,
}
impl Default for CvFoldGenerator {
fn default() -> Self {
Self {
target_n_folds: 5,
min_initial_window: 10,
horizon: 1,
step_size: None,
gap: 0,
purge: 0,
embargo: 0,
strategy: CVStrategy::Expanding,
on_constraint_violation: ConstraintViolation::Error,
}
}
}
impl CvFoldGenerator {
pub fn new() -> Self {
Self::default()
}
pub fn n_folds(mut self, n: usize) -> Self {
self.target_n_folds = n;
self
}
pub fn min_initial_window(mut self, size: usize) -> Self {
self.min_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 = Some(step.max(1));
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 embargo(mut self, e: usize) -> Self {
self.embargo = e;
self
}
pub fn strategy(mut self, s: CVStrategy) -> Self {
self.strategy = s;
self
}
pub fn on_constraint_violation(mut self, behavior: ConstraintViolation) -> Self {
self.on_constraint_violation = behavior;
self
}
pub fn ensure_end_coverage(self, _enable: bool) -> Self {
self
}
pub fn generate(&self, series_len: usize) -> crate::error::Result<Vec<Fold>> {
use crate::error::ForecastError;
let horizon = self.horizon.max(1);
let gap = self.gap;
let purge = self.purge;
let min_train = self.min_initial_window.max(1);
let step = self.step_size.unwrap_or(horizon).max(1);
let min_series = min_train + purge + gap + horizon;
if series_len < min_series {
return Err(ForecastError::InsufficientData {
needed: min_series,
got: series_len,
hint: Some(format!(
"need at least min_initial_window({}) + purge({}) + gap({}) + horizon({})",
min_train, purge, gap, horizon
)),
});
}
let last_origin = series_len - gap - horizon;
let mut candidates = Vec::new();
let mut origin = last_origin;
loop {
let test_start = origin + gap;
let test_end = (test_start + horizon).min(series_len);
let train_end = origin.saturating_sub(purge);
let base_train_start = match self.strategy {
CVStrategy::Rolling => train_end.saturating_sub(min_train),
CVStrategy::Expanding => 0,
};
if train_end <= base_train_start || train_end - base_train_start < min_train {
break; }
candidates.push((base_train_start, train_end, test_start, test_end));
if origin < step {
break;
}
origin -= step;
}
candidates.reverse();
let n_folds = self.target_n_folds.max(1);
if candidates.len() > n_folds {
candidates = candidates[candidates.len() - n_folds..].to_vec();
}
if candidates.is_empty() {
return match self.on_constraint_violation {
ConstraintViolation::Error => Err(ForecastError::InvalidParameter(format!(
"no valid folds: min_initial_window={}, horizon={}, gap={}, purge={}, series_len={}",
min_train, horizon, gap, purge, series_len
))),
ConstraintViolation::ReduceFolds => Ok(Vec::new()),
};
}
let mut folds = Vec::with_capacity(candidates.len());
let mut max_embargo_end: usize = 0;
for (base_train_start, train_end, test_start, test_end) in candidates {
let train_start = if self.embargo > 0 && max_embargo_end > base_train_start {
max_embargo_end.min(train_end)
} else {
base_train_start
};
if train_end <= train_start {
let embargo_end = (test_end + self.embargo).min(series_len);
if embargo_end > max_embargo_end {
max_embargo_end = embargo_end;
}
continue;
}
folds.push(Fold {
train_start,
train_end,
test_start,
test_end,
});
if self.embargo > 0 {
let embargo_end = (test_end + self.embargo).min(series_len);
if embargo_end > max_embargo_end {
max_embargo_end = embargo_end;
}
}
}
if folds.is_empty() {
return match self.on_constraint_violation {
ConstraintViolation::Error => Err(ForecastError::InvalidParameter(
"all folds eliminated by embargo".to_string(),
)),
ConstraintViolation::ReduceFolds => Ok(Vec::new()),
};
}
Ok(folds)
}
}
#[derive(Debug, Clone)]
pub struct CVConfig {
pub horizon: usize,
pub min_initial_window: usize,
pub step_size: usize,
pub strategy: CVStrategy,
pub seasonal_period: Option<usize>,
pub gap: usize,
pub purge: usize,
pub embargo: usize,
}
impl Default for CVConfig {
fn default() -> Self {
Self {
horizon: 1,
min_initial_window: 10,
step_size: 1,
strategy: CVStrategy::Expanding,
seasonal_period: None,
gap: 0,
purge: 0,
embargo: 0,
}
}
}
impl CVConfig {
pub fn expanding(min_initial_window: usize, horizon: usize) -> Self {
Self {
min_initial_window,
horizon,
step_size: 1,
strategy: CVStrategy::Expanding,
seasonal_period: None,
gap: 0,
purge: 0,
embargo: 0,
}
}
pub fn rolling(window_size: usize, horizon: usize) -> Self {
Self {
min_initial_window: window_size,
horizon,
step_size: 1,
strategy: CVStrategy::Rolling,
seasonal_period: None,
gap: 0,
purge: 0,
embargo: 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 with_embargo(mut self, embargo: usize) -> Self {
self.embargo = embargo;
self
}
pub fn to_fold_generator(&self) -> CvFoldGenerator {
CvFoldGenerator {
target_n_folds: 5, min_initial_window: self.min_initial_window,
horizon: self.horizon,
step_size: if self.step_size > 0 {
Some(self.step_size)
} else {
None
},
gap: self.gap,
purge: self.purge,
embargo: self.embargo,
strategy: self.strategy,
on_constraint_violation: ConstraintViolation::ReduceFolds,
}
}
}
#[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 horizon = fold.test_size();
let forecast = if model.has_exog() {
let test_series = series.slice(fold.test_start, fold.test_end)?;
let future_regs = test_series.all_regressors();
model.predict_with_exog(horizon, &future_regs)?
} else {
model.predict(horizon)?
};
let predicted = forecast.primary();
let actual_slice = &series.primary_values()[fold.test_start..fold.test_end];
let metrics = calculate_metrics(actual_slice, predicted, seasonal_period)?;
Ok((metrics, actual_slice.to_vec(), predicted.to_vec()))
}
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(),
));
}
#[cfg(feature = "parallel")]
let group_results: Vec<(String, CVResults)> = {
use rayon::prelude::*;
let results: Vec<std::result::Result<(String, CVResults), ForecastError>> = series_vec
.into_par_iter()
.map(|(group_id, series)| {
let cv_result = cross_validate_with_folds(config, &series, &folds, &model_factory)?;
Ok((group_id, cv_result))
})
.collect();
let mut group_results = Vec::with_capacity(results.len());
for r in results {
group_results.push(r?);
}
group_results
};
#[cfg(not(feature = "parallel"))]
let group_results: Vec<(String, CVResults)> = {
let mut group_results = Vec::with_capacity(series_vec.len());
for (group_id, series) in series_vec {
let cv_result = cross_validate_with_folds(config, &series, &folds, &model_factory)?;
group_results.push((group_id, cv_result));
}
group_results
};
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 (_, cv_result) in &group_results {
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);
}
}
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 + Send,
Factory: Fn() -> F + Sync,
{
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: folds.to_vec(),
})
}
#[derive(Debug, Clone)]
pub struct StreamingCVAggregator {
count: usize,
mae_mean: f64,
mae_m2: f64,
rmse_mean: f64,
rmse_m2: f64,
smape_mean: f64,
smape_m2: f64,
mape_mean: f64,
mape_m2: f64,
mape_count: usize,
prev_mae_mean: f64,
}
impl StreamingCVAggregator {
pub fn new() -> Self {
Self {
count: 0,
mae_mean: 0.0,
mae_m2: 0.0,
rmse_mean: 0.0,
rmse_m2: 0.0,
smape_mean: 0.0,
smape_m2: 0.0,
mape_mean: 0.0,
mape_m2: 0.0,
mape_count: 0,
prev_mae_mean: f64::NAN,
}
}
pub fn update(&mut self, metrics: &AccuracyMetrics) {
self.prev_mae_mean = self.mae_mean;
self.count += 1;
let n = self.count as f64;
let delta = metrics.mae - self.mae_mean;
self.mae_mean += delta / n;
let delta2 = metrics.mae - self.mae_mean;
self.mae_m2 += delta * delta2;
let delta = metrics.rmse - self.rmse_mean;
self.rmse_mean += delta / n;
let delta2 = metrics.rmse - self.rmse_mean;
self.rmse_m2 += delta * delta2;
let delta = metrics.smape - self.smape_mean;
self.smape_mean += delta / n;
let delta2 = metrics.smape - self.smape_mean;
self.smape_m2 += delta * delta2;
if let Some(mape) = metrics.mape {
self.mape_count += 1;
let mn = self.mape_count as f64;
let delta = mape - self.mape_mean;
self.mape_mean += delta / mn;
let delta2 = mape - self.mape_mean;
self.mape_m2 += delta * delta2;
}
}
pub fn n_folds(&self) -> usize {
self.count
}
pub fn mean_mae(&self) -> f64 {
self.mae_mean
}
pub fn mean_rmse(&self) -> f64 {
self.rmse_mean
}
pub fn mean_smape(&self) -> f64 {
self.smape_mean
}
pub fn mean_mape(&self) -> Option<f64> {
if self.mape_count > 0 {
Some(self.mape_mean)
} else {
None
}
}
pub fn std_mae(&self) -> f64 {
if self.count < 2 {
return 0.0;
}
(self.mae_m2 / (self.count - 1) as f64).sqrt()
}
pub fn std_rmse(&self) -> f64 {
if self.count < 2 {
return 0.0;
}
(self.rmse_m2 / (self.count - 1) as f64).sqrt()
}
pub fn has_converged(&self, tolerance: f64) -> bool {
if self.count < 3 || self.prev_mae_mean.is_nan() {
return false;
}
let change = (self.mae_mean - self.prev_mae_mean).abs();
let scale = self.mae_mean.abs().max(1e-10);
change / scale < tolerance
}
pub fn finalize(&self) -> AggregatedMetrics {
AggregatedMetrics {
mae: self.mae_mean,
rmse: self.rmse_mean,
smape: self.smape_mean,
mape: self.mean_mape(),
mae_std: self.std_mae(),
rmse_std: self.std_rmse(),
}
}
}
impl Default for StreamingCVAggregator {
fn default() -> Self {
Self::new()
}
}
pub fn cross_validate_early_stop<F, Factory>(
config: &CVConfig,
series: &TimeSeries,
model_factory: Factory,
tolerance: f64,
) -> Result<CVResults>
where
F: Forecaster,
Factory: Fn() -> F,
{
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 mut aggregator = StreamingCVAggregator::new();
let mut fold_metrics = Vec::new();
let mut all_actual = Vec::new();
let mut all_predicted = Vec::new();
let mut used_folds = Vec::new();
for fold in &folds {
let (metrics, actual, predicted) =
evaluate_fold(series, fold, &model_factory, config.seasonal_period)?;
aggregator.update(&metrics);
fold_metrics.push(metrics);
all_actual.extend_from_slice(&actual);
all_predicted.extend_from_slice(&predicted);
used_folds.push(fold.clone());
if aggregator.has_converged(tolerance) {
break;
}
}
Ok(CVResults {
n_folds: fold_metrics.len(),
aggregated: aggregator.finalize(),
fold_metrics,
actual_values: all_actual,
predicted_values: all_predicted,
folds: used_folds,
})
}
#[derive(Debug, Clone)]
pub struct RollingForecastConfig {
pub initial_train_size: usize,
pub horizon: usize,
pub step_size: usize,
pub expanding: bool,
}
impl RollingForecastConfig {
pub fn new(initial_train_size: usize, horizon: usize) -> Self {
Self {
initial_train_size,
horizon,
step_size: horizon,
expanding: true,
}
}
pub fn step_size(mut self, step: usize) -> Self {
self.step_size = step;
self
}
pub fn expanding(mut self, expanding: bool) -> Self {
self.expanding = expanding;
self
}
}
#[derive(Debug, Clone)]
pub struct RollingForecastWindow {
pub train_start: usize,
pub train_end: usize,
pub predictions: Vec<f64>,
pub actuals: Vec<f64>,
}
#[derive(Debug, Clone)]
pub struct RollingForecastResult {
pub windows: Vec<RollingForecastWindow>,
pub all_predictions: Vec<f64>,
pub all_actuals: Vec<f64>,
pub window_metrics: Vec<AccuracyMetrics>,
pub aggregated: AggregatedMetrics,
}
pub fn rolling_forecast<F, Factory>(
series: &TimeSeries,
config: &RollingForecastConfig,
model_factory: Factory,
) -> Result<RollingForecastResult>
where
F: Forecaster + Send,
Factory: Fn() -> F + Sync,
{
let n = series.len();
if config.initial_train_size == 0 {
return Err(ForecastError::InvalidParameter(
"initial_train_size must be at least 1".to_string(),
));
}
if config.horizon == 0 {
return Err(ForecastError::InvalidParameter(
"horizon must be at least 1".to_string(),
));
}
if config.step_size == 0 {
return Err(ForecastError::InvalidParameter(
"step_size must be at least 1".to_string(),
));
}
if config.initial_train_size + config.horizon > n {
return Err(ForecastError::InvalidParameter(format!(
"Series length ({}) is too short for initial_train_size ({}) + horizon ({})",
n, config.initial_train_size, config.horizon
)));
}
let mut window_specs: Vec<(usize, usize)> = Vec::new();
let mut origin = config.initial_train_size;
while origin + config.horizon <= n {
let train_start = if config.expanding {
0
} else {
origin.saturating_sub(config.initial_train_size)
};
let train_end = origin;
window_specs.push((train_start, train_end));
origin += config.step_size;
}
if window_specs.is_empty() {
return Err(ForecastError::InvalidParameter(
"Not enough data for any forecast window".to_string(),
));
}
let horizon = config.horizon;
let values = series.primary_values();
let evaluate_window =
|&(train_start, train_end): &(usize, usize)| -> Result<(RollingForecastWindow, AccuracyMetrics)> {
let train_series = series.slice(train_start, train_end)?;
let mut model = model_factory();
model.fit(&train_series)?;
let forecast = model.predict(horizon)?;
let predictions: Vec<f64> = forecast.primary().to_vec();
let test_end = train_end + horizon;
let actuals: Vec<f64> = (train_end..test_end).map(|i| values[i]).collect();
let metrics = calculate_metrics(&actuals, &predictions, None)?;
let window = RollingForecastWindow {
train_start,
train_end,
predictions,
actuals,
};
Ok((window, metrics))
};
#[cfg(feature = "parallel")]
let results: Vec<Result<(RollingForecastWindow, AccuracyMetrics)>> =
window_specs.par_iter().map(evaluate_window).collect();
#[cfg(not(feature = "parallel"))]
let results: Vec<Result<(RollingForecastWindow, AccuracyMetrics)>> =
window_specs.iter().map(evaluate_window).collect();
let mut windows = Vec::with_capacity(results.len());
let mut window_metrics = Vec::with_capacity(results.len());
let mut all_predictions = Vec::new();
let mut all_actuals = Vec::new();
for result in results {
let (window, metrics) = result?;
all_predictions.extend_from_slice(&window.predictions);
all_actuals.extend_from_slice(&window.actuals);
windows.push(window);
window_metrics.push(metrics);
}
let n_windows = window_metrics.len();
let mae_values: Vec<f64> = window_metrics.iter().map(|m| m.mae).collect();
let rmse_values: Vec<f64> = window_metrics.iter().map(|m| m.rmse).collect();
let smape_values: Vec<f64> = window_metrics.iter().map(|m| m.smape).collect();
let mae_mean = mae_values.iter().sum::<f64>() / n_windows as f64;
let rmse_mean = rmse_values.iter().sum::<f64>() / n_windows as f64;
let smape_mean = smape_values.iter().sum::<f64>() / n_windows as f64;
let mae_std = std_dev(&mae_values);
let rmse_std = std_dev(&rmse_values);
let mape = if window_metrics.iter().all(|m| m.mape.is_some()) {
let mape_values: Vec<f64> = window_metrics.iter().filter_map(|m| m.mape).collect();
Some(mape_values.iter().sum::<f64>() / n_windows as f64)
} else {
None
};
let aggregated = AggregatedMetrics {
mae: mae_mean,
rmse: rmse_mean,
smape: smape_mean,
mape,
mae_std,
rmse_std,
};
Ok(RollingForecastResult {
windows,
all_predictions,
all_actuals,
window_metrics,
aggregated,
})
}
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 folds = CvFoldGenerator::new()
.n_folds(5)
.min_initial_window(10)
.horizon(1)
.generate(50)
.unwrap();
assert_eq!(folds.len(), 5);
assert!(folds[0].train_size() >= 10);
assert_eq!(folds.last().unwrap().test_end, 50);
}
#[test]
fn fold_generator_with_gap() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.min_initial_window(10)
.horizon(1)
.gap(2)
.generate(50)
.unwrap();
for fold in &folds {
assert!(fold.test_start >= fold.train_end + 2);
}
}
#[test]
fn fold_generator_with_purge() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.min_initial_window(10)
.horizon(1)
.purge(2)
.generate(50)
.unwrap();
for fold in &folds {
assert!(fold.train_end < fold.test_start);
}
}
#[test]
fn fold_generator_rolling() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.min_initial_window(10)
.horizon(1)
.strategy(CVStrategy::Rolling)
.generate(50)
.unwrap();
for fold in &folds {
assert_eq!(fold.train_size(), 10);
}
}
#[test]
fn fold_generator_multi_step_horizon() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.min_initial_window(10)
.horizon(3)
.generate(50)
.unwrap();
for fold in &folds {
assert_eq!(fold.test_size(), 3);
}
}
#[test]
fn fold_generator_insufficient_data() {
let result = CvFoldGenerator::new()
.min_initial_window(10)
.horizon(5)
.generate(10);
assert!(result.is_err());
}
#[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.is_empty());
assert!(!test.is_empty());
let (train, test) = train_test_split(&ts, 0.99).unwrap();
assert!(!train.is_empty());
assert!(!test.is_empty());
}
#[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!(results.n_folds > 0);
assert!(results.aggregated.mae.is_finite());
assert_eq!(results.folds.len(), results.n_folds);
}
#[test]
fn cv_rolling_window_basic() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).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!(results.n_folds > 0);
assert!(results.aggregated.mae.is_finite());
}
#[test]
fn cv_with_gap() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1).with_gap(3);
let results = cross_validate(&config, &ts, Naive::new).unwrap();
for fold in &results.folds {
assert!(fold.test_start >= fold.train_end + 3);
}
}
#[test]
fn cv_with_purge() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).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();
for fold in &results.folds {
assert!(fold.test_start > fold.train_end);
}
}
#[test]
fn cv_multi_step_horizon() {
let timestamps = make_timestamps(50);
let values: Vec<f64> = (0..50).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!(results.n_folds > 0);
assert_eq!(results.actual_values.len(), results.n_folds * 3);
assert_eq!(results.predicted_values.len(), results.n_folds * 3);
}
#[test]
fn cv_insufficient_data_returns_error_or_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);
match cross_validate(&config, &ts, Naive::new) {
Ok(results) => assert_eq!(results.n_folds, 0),
Err(_) => {} }
}
#[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.min_initial_window, 10);
assert_eq!(expanding.horizon, 3);
assert_eq!(expanding.strategy, CVStrategy::Expanding);
let rolling = CVConfig::rolling(15, 2);
assert_eq!(rolling.min_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.min_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, 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, Naive::new).unwrap();
let fold_counts: Vec<_> = results
.group_results
.iter()
.map(|(_, r)| r.n_folds)
.collect();
assert!(fold_counts.iter().all(|&n| n == fold_counts[0]));
assert!(fold_counts[0] > 0);
}
#[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, Naive::new);
assert!(result.is_err());
}
#[test]
fn streaming_aggregator_single_fold() {
let mut agg = StreamingCVAggregator::new();
let metrics = AccuracyMetrics {
mae: 2.0,
mse: 0.0,
rmse: 3.0,
smape: 15.0,
mape: Some(10.0),
mase: None,
r_squared: 0.0,
};
agg.update(&metrics);
assert_eq!(agg.n_folds(), 1);
assert_relative_eq!(agg.mean_mae(), 2.0);
assert_relative_eq!(agg.mean_rmse(), 3.0);
assert_relative_eq!(agg.mean_smape(), 15.0);
assert_relative_eq!(agg.mean_mape().unwrap(), 10.0);
assert_relative_eq!(agg.std_mae(), 0.0);
}
#[test]
fn streaming_aggregator_matches_batch() {
let fold_metrics = vec![
AccuracyMetrics {
mae: 1.0,
mse: 0.0,
rmse: 1.5,
smape: 10.0,
mape: Some(8.0),
mase: None,
r_squared: 0.0,
},
AccuracyMetrics {
mae: 2.0,
mse: 0.0,
rmse: 2.5,
smape: 12.0,
mape: Some(9.0),
mase: None,
r_squared: 0.0,
},
AccuracyMetrics {
mae: 3.0,
mse: 0.0,
rmse: 3.5,
smape: 14.0,
mape: Some(11.0),
mase: None,
r_squared: 0.0,
},
];
let mut agg = StreamingCVAggregator::new();
for m in &fold_metrics {
agg.update(m);
}
let mae_vals: Vec<f64> = fold_metrics.iter().map(|m| m.mae).collect();
let batch_mean = mae_vals.iter().sum::<f64>() / mae_vals.len() as f64;
let batch_std = std_dev(&mae_vals);
assert_relative_eq!(agg.mean_mae(), batch_mean, epsilon = 1e-10);
assert_relative_eq!(agg.std_mae(), batch_std, epsilon = 1e-10);
assert_eq!(agg.n_folds(), 3);
}
#[test]
fn streaming_aggregator_convergence() {
let mut agg = StreamingCVAggregator::new();
agg.update(&AccuracyMetrics {
mae: 1.0,
mse: 0.0,
rmse: 1.0,
smape: 5.0,
mape: None,
mase: None,
r_squared: 0.0,
});
assert!(!agg.has_converged(0.01));
agg.update(&AccuracyMetrics {
mae: 1.0,
mse: 0.0,
rmse: 1.0,
smape: 5.0,
mape: None,
mase: None,
r_squared: 0.0,
});
assert!(!agg.has_converged(0.01));
agg.update(&AccuracyMetrics {
mae: 1.0,
mse: 0.0,
rmse: 1.0,
smape: 5.0,
mape: None,
mase: None,
r_squared: 0.0,
});
assert!(agg.has_converged(0.01));
}
#[test]
fn streaming_aggregator_no_mape() {
let mut agg = StreamingCVAggregator::new();
agg.update(&AccuracyMetrics {
mae: 1.0,
mse: 0.0,
rmse: 1.0,
smape: 5.0,
mape: None,
mase: None,
r_squared: 0.0,
});
assert!(agg.mean_mape().is_none());
}
#[test]
fn streaming_aggregator_finalize() {
let mut agg = StreamingCVAggregator::new();
agg.update(&AccuracyMetrics {
mae: 2.0,
mse: 0.0,
rmse: 3.0,
smape: 10.0,
mape: Some(5.0),
mase: None,
r_squared: 0.0,
});
agg.update(&AccuracyMetrics {
mae: 4.0,
mse: 0.0,
rmse: 5.0,
smape: 20.0,
mape: Some(15.0),
mase: None,
r_squared: 0.0,
});
let result = agg.finalize();
assert_relative_eq!(result.mae, 3.0);
assert_relative_eq!(result.rmse, 4.0);
assert_relative_eq!(result.smape, 15.0);
assert_relative_eq!(result.mape.unwrap(), 10.0);
}
#[test]
fn cv_early_stop_constant_series() {
let timestamps = make_timestamps(50);
let values = vec![5.0; 50];
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1);
let results = cross_validate_early_stop(&config, &ts, Naive::new, 0.01).unwrap();
assert!(results.n_folds >= 3);
assert!(results.n_folds < 40); assert_relative_eq!(results.aggregated.mae, 0.0, epsilon = 1e-10);
}
#[test]
fn cv_early_stop_runs_all_if_needed() {
let timestamps = make_timestamps(20);
let values: Vec<f64> = (0..20)
.map(|i| if i % 2 == 0 { 100.0 } else { 0.0 })
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = CVConfig::expanding(10, 1);
let results = cross_validate_early_stop(&config, &ts, Naive::new, 1e-15).unwrap();
assert!(results.n_folds > 0);
}
#[test]
fn rolling_forecast_expanding_basic() {
let timestamps = make_timestamps(30);
let values: Vec<f64> = (0..30).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = RollingForecastConfig::new(20, 3).step_size(3);
let result = rolling_forecast(&ts, &config, Naive::new).unwrap();
assert_eq!(result.windows.len(), 3);
assert_eq!(result.all_predictions.len(), 9);
assert_eq!(result.all_actuals.len(), 9);
for w in &result.windows {
assert_eq!(w.train_start, 0);
}
assert_eq!(result.windows[0].train_end, 20);
assert_eq!(result.windows[1].train_end, 23);
assert_eq!(result.windows[2].train_end, 26);
}
#[test]
fn rolling_forecast_fixed_window() {
let timestamps = make_timestamps(30);
let values: Vec<f64> = (0..30).map(|i| i as f64).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = RollingForecastConfig::new(15, 3)
.step_size(3)
.expanding(false);
let result = rolling_forecast(&ts, &config, Naive::new).unwrap();
for w in &result.windows {
assert_eq!(w.train_end - w.train_start, 15);
}
assert_eq!(result.windows[0].train_start, 0);
assert_eq!(result.windows[1].train_start, 3);
}
#[test]
fn rolling_forecast_step_size_one() {
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 = RollingForecastConfig::new(20, 3).step_size(1);
let result = rolling_forecast(&ts, &config, Naive::new).unwrap();
assert_eq!(result.windows.len(), 3);
}
#[test]
fn rolling_forecast_constant_series() {
let timestamps = make_timestamps(30);
let values = vec![5.0; 30];
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = RollingForecastConfig::new(20, 3).step_size(3);
let result = rolling_forecast(&ts, &config, Naive::new).unwrap();
assert!(result.aggregated.mae.abs() < 1e-10);
assert!(result.aggregated.rmse.abs() < 1e-10);
for (p, a) in result.all_predictions.iter().zip(result.all_actuals.iter()) {
assert!((p - a).abs() < 1e-10);
}
}
#[test]
fn rolling_forecast_actuals_match_series() {
let timestamps = make_timestamps(30);
let values: Vec<f64> = (0..30).map(|i| i as f64 * 2.0 + 1.0).collect();
let ts = TimeSeries::univariate(timestamps, values.clone()).unwrap();
let config = RollingForecastConfig::new(20, 5).step_size(5);
let result = rolling_forecast(&ts, &config, Naive::new).unwrap();
for w in &result.windows {
for (j, &actual) in w.actuals.iter().enumerate() {
let idx = w.train_end + j;
assert!((actual - values[idx]).abs() < 1e-10);
}
}
}
#[test]
fn rolling_forecast_insufficient_data() {
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 config = RollingForecastConfig::new(10, 5);
let result = rolling_forecast(&ts, &config, Naive::new);
assert!(result.is_err());
}
#[test]
fn rolling_forecast_invalid_params() {
let timestamps = make_timestamps(30);
let values = vec![1.0; 30];
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = RollingForecastConfig {
initial_train_size: 20,
horizon: 0,
step_size: 1,
expanding: true,
};
assert!(rolling_forecast(&ts, &config, Naive::new).is_err());
let config = RollingForecastConfig {
initial_train_size: 20,
horizon: 3,
step_size: 0,
expanding: true,
};
assert!(rolling_forecast(&ts, &config, Naive::new).is_err());
let config = RollingForecastConfig {
initial_train_size: 0,
horizon: 3,
step_size: 1,
expanding: true,
};
assert!(rolling_forecast(&ts, &config, Naive::new).is_err());
}
#[test]
fn rolling_forecast_config_builder() {
let config = RollingForecastConfig::new(50, 7)
.step_size(3)
.expanding(false);
assert_eq!(config.initial_train_size, 50);
assert_eq!(config.horizon, 7);
assert_eq!(config.step_size, 3);
assert!(!config.expanding);
}
#[test]
fn rolling_forecast_config_defaults() {
let config = RollingForecastConfig::new(100, 12);
assert_eq!(config.initial_train_size, 100);
assert_eq!(config.horizon, 12);
assert_eq!(config.step_size, 12); assert!(config.expanding); }
#[test]
fn rolling_forecast_metrics_consistent() {
let timestamps = make_timestamps(40);
let values: Vec<f64> = (0..40).map(|i| (i as f64).sin() * 10.0 + 50.0).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = RollingForecastConfig::new(25, 3).step_size(3);
let result = rolling_forecast(&ts, &config, Naive::new).unwrap();
assert!(result.aggregated.rmse >= result.aggregated.mae);
assert!(result.aggregated.smape >= 0.0);
assert!(result.aggregated.smape <= 200.0);
assert!(result.aggregated.mae_std >= 0.0);
assert!(result.aggregated.rmse_std >= 0.0);
assert_eq!(result.window_metrics.len(), result.windows.len());
}
#[test]
fn grouped_cv_parallel_matches_sequential_results() {
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![("a".to_string(), series_a), ("b".to_string(), series_b)];
let config = CVConfig::expanding(15, 3).with_step_size(3);
let results = grouped_cross_validate(&config, series_map, Naive::new).unwrap();
assert_eq!(results.group_results.len(), 2);
assert!(results.aggregated.mae.is_finite());
assert!(results.aggregated.rmse.is_finite());
let n_a = results.group_results[0].1.n_folds;
let n_b = results.group_results[1].1.n_folds;
assert_eq!(n_a, n_b);
assert!(n_a > 0);
}
#[test]
fn fold_generator_embargo_zero_matches_no_embargo() {
let g1 = CvFoldGenerator::new()
.min_initial_window(10)
.horizon(3)
.step_size(3);
let g2 = CvFoldGenerator::new()
.min_initial_window(10)
.horizon(3)
.step_size(3)
.embargo(0);
assert_eq!(g1.generate(50), g2.generate(50));
}
#[test]
fn fold_generator_embargo_shrinks_training() {
let folds = CvFoldGenerator::new()
.min_initial_window(10)
.horizon(3)
.step_size(3)
.embargo(5)
.generate(50)
.unwrap();
assert!(!folds.is_empty());
if folds.len() > 1 {
assert!(folds[1].train_start > 0);
}
}
#[test]
fn fold_generator_embargo_with_gap_and_purge() {
let folds = CvFoldGenerator::new()
.min_initial_window(10)
.horizon(3)
.step_size(3)
.gap(1)
.purge(1)
.embargo(3)
.generate(50)
.unwrap();
assert!(!folds.is_empty());
for fold in &folds {
assert!(fold.train_end > fold.train_start);
}
}
#[test]
fn fold_generator_embargo_beyond_series_clamps() {
let folds = CvFoldGenerator::new()
.min_initial_window(10)
.horizon(3)
.step_size(3)
.embargo(1000)
.generate(30)
.unwrap();
assert!(folds.len() <= 2, "got {}", folds.len());
}
#[test]
fn fold_generator_embargo_expanding_vs_rolling() {
assert!(!CvFoldGenerator::new()
.min_initial_window(10)
.horizon(2)
.step_size(2)
.strategy(CVStrategy::Expanding)
.embargo(3)
.generate(40)
.unwrap()
.is_empty());
for fold in &CvFoldGenerator::new()
.min_initial_window(10)
.horizon(2)
.strategy(CVStrategy::Rolling)
.embargo(3)
.generate(40)
.unwrap()
{
assert!(fold.train_end > fold.train_start);
}
}
#[test]
fn fold_generator_embargo_cvconfig_integration() {
let gen = CVConfig::expanding(10, 3)
.with_step_size(3)
.with_embargo(5)
.to_fold_generator();
assert_eq!(gen.embargo, 5);
assert!(!gen.generate(50).unwrap().is_empty());
}
#[test]
fn a1_backward_anchored_expanding_h1() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(1)
.min_initial_window(5)
.generate(10)
.unwrap();
assert_eq!(folds.len(), 5);
assert_eq!(
folds[0],
Fold {
train_start: 0,
train_end: 5,
test_start: 5,
test_end: 6
}
);
assert_eq!(
folds[1],
Fold {
train_start: 0,
train_end: 6,
test_start: 6,
test_end: 7
}
);
assert_eq!(
folds[2],
Fold {
train_start: 0,
train_end: 7,
test_start: 7,
test_end: 8
}
);
assert_eq!(
folds[3],
Fold {
train_start: 0,
train_end: 8,
test_start: 8,
test_end: 9
}
);
assert_eq!(
folds[4],
Fold {
train_start: 0,
train_end: 9,
test_start: 9,
test_end: 10
}
);
}
#[test]
fn a2_backward_anchored_expanding_h3() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(3)
.min_initial_window(5)
.generate(20)
.unwrap();
assert_eq!(folds.len(), 3);
assert_eq!(
folds[0],
Fold {
train_start: 0,
train_end: 11,
test_start: 11,
test_end: 14
}
);
assert_eq!(
folds[1],
Fold {
train_start: 0,
train_end: 14,
test_start: 14,
test_end: 17
}
);
assert_eq!(
folds[2],
Fold {
train_start: 0,
train_end: 17,
test_start: 17,
test_end: 20
}
);
}
#[test]
fn a3_all_indices_within_bounds() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(7)
.min_initial_window(30)
.generate(200)
.unwrap();
for fold in &folds {
assert!(fold.train_start < fold.train_end);
assert!(fold.train_end <= 200);
assert!(fold.test_start < fold.test_end);
assert!(fold.test_end <= 200);
}
}
#[test]
fn a4_no_train_test_overlap() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(7)
.min_initial_window(30)
.generate(200)
.unwrap();
for fold in &folds {
assert!(fold.train_end <= fold.test_start);
}
}
#[test]
fn a5_folds_chronologically_ordered() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(5)
.min_initial_window(20)
.generate(200)
.unwrap();
for i in 1..folds.len() {
assert!(folds[i].test_start > folds[i - 1].test_start);
}
}
#[test]
fn a6_test_sets_non_overlapping() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(7)
.min_initial_window(30)
.generate(200)
.unwrap();
for i in 1..folds.len() {
assert!(
folds[i].test_start >= folds[i - 1].test_end,
"fold {} test_start {} < fold {} test_end {}",
i,
folds[i].test_start,
i - 1,
folds[i - 1].test_end
);
}
}
#[test]
fn a7_test_sets_contiguous() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(3)
.min_initial_window(5)
.generate(30)
.unwrap();
for i in 1..folds.len() {
assert_eq!(folds[i].test_start, folds[i - 1].test_end);
}
}
#[test]
fn b1_expanding_all_start_at_zero() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(5)
.min_initial_window(20)
.generate(200)
.unwrap();
for fold in &folds {
assert_eq!(fold.train_start, 0);
}
}
#[test]
fn b2_expanding_train_size_non_decreasing() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(5)
.min_initial_window(20)
.generate(200)
.unwrap();
for i in 1..folds.len() {
assert!(folds[i].train_size() >= folds[i - 1].train_size());
}
}
#[test]
fn b3_expanding_train_grows_by_horizon() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(7)
.min_initial_window(20)
.generate(200)
.unwrap();
for i in 1..folds.len() {
assert_eq!(folds[i].train_size() - folds[i - 1].train_size(), 7);
}
}
#[test]
fn c1_rolling_exact_indices() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(1)
.min_initial_window(5)
.strategy(CVStrategy::Rolling)
.generate(10)
.unwrap();
assert_eq!(folds.len(), 5);
assert_eq!(
folds[0],
Fold {
train_start: 0,
train_end: 5,
test_start: 5,
test_end: 6
}
);
assert_eq!(
folds[1],
Fold {
train_start: 1,
train_end: 6,
test_start: 6,
test_end: 7
}
);
assert_eq!(
folds[2],
Fold {
train_start: 2,
train_end: 7,
test_start: 7,
test_end: 8
}
);
assert_eq!(
folds[3],
Fold {
train_start: 3,
train_end: 8,
test_start: 8,
test_end: 9
}
);
assert_eq!(
folds[4],
Fold {
train_start: 4,
train_end: 9,
test_start: 9,
test_end: 10
}
);
}
#[test]
fn c2_rolling_all_same_train_size() {
let folds = CvFoldGenerator::new()
.n_folds(10)
.horizon(1)
.min_initial_window(50)
.strategy(CVStrategy::Rolling)
.generate(500)
.unwrap();
for fold in &folds {
assert_eq!(fold.train_size(), 50);
}
}
#[test]
fn c3_rolling_train_start_increases() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(5)
.min_initial_window(20)
.strategy(CVStrategy::Rolling)
.generate(200)
.unwrap();
for i in 1..folds.len() {
assert!(folds[i].train_start > folds[i - 1].train_start);
}
}
#[test]
fn d1_last_fold_expanding() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(12)
.min_initial_window(20)
.generate(144)
.unwrap();
assert_eq!(folds.last().unwrap().test_end, 144);
}
#[test]
fn d2_last_fold_rolling() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(5)
.min_initial_window(20)
.strategy(CVStrategy::Rolling)
.generate(100)
.unwrap();
assert_eq!(folds.last().unwrap().test_end, 100);
}
#[test]
fn d3_last_fold_with_gap() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(5)
.min_initial_window(10)
.gap(3)
.generate(80)
.unwrap();
assert_eq!(folds.last().unwrap().test_end, 80);
}
#[test]
fn d4_single_fold_anchored() {
let folds = CvFoldGenerator::new()
.n_folds(1)
.horizon(5)
.min_initial_window(10)
.generate(50)
.unwrap();
assert_eq!(folds.len(), 1);
assert_eq!(folds[0].test_end, 50);
assert_eq!(folds[0].train_start, 0);
assert_eq!(folds[0].train_end, 45);
}
#[test]
fn e1_expanding_respects_min() {
let folds = CvFoldGenerator::new()
.n_folds(20)
.horizon(7)
.min_initial_window(50)
.generate(200)
.unwrap();
for (i, fold) in folds.iter().enumerate() {
assert!(
fold.train_size() >= 50,
"fold {} train_size {}",
i,
fold.train_size()
);
}
}
#[test]
fn e2_min_constraint_drops_early_folds() {
let folds = CvFoldGenerator::new()
.n_folds(10)
.horizon(5)
.min_initial_window(15)
.generate(30)
.unwrap();
assert_eq!(folds.len(), 3);
assert!(folds[0].train_size() >= 15);
}
#[test]
fn e3_n_folds_caps_not_constraint() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(1)
.min_initial_window(5)
.generate(100)
.unwrap();
assert_eq!(folds.len(), 3);
}
#[test]
fn f1_gap_respected_all_folds() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(5)
.min_initial_window(10)
.gap(3)
.generate(80)
.unwrap();
for fold in &folds {
assert!(fold.test_start >= fold.train_end + 3);
}
}
#[test]
fn f2_gap_zero_means_adjacent() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(1)
.min_initial_window(5)
.gap(0)
.generate(20)
.unwrap();
for fold in &folds {
assert_eq!(fold.test_start, fold.train_end);
}
}
#[test]
fn g1_purge_creates_separation() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(2)
.min_initial_window(5)
.purge(3)
.generate(30)
.unwrap();
for fold in &folds {
assert!(
fold.test_start > fold.train_end,
"purge should separate train and test"
);
assert!(fold.test_start - fold.train_end >= 3);
}
}
#[test]
fn h1_embargo_shifts_subsequent_train_start() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(2)
.min_initial_window(5)
.embargo(4)
.generate(40)
.unwrap();
assert_eq!(folds[0].train_start, 0);
if folds.len() > 1 {
assert!(
folds[1].train_start > 0,
"embargo should shift fold 1 train_start"
);
}
}
#[test]
fn h2_large_embargo_reduces_folds() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(2)
.min_initial_window(3)
.embargo(15)
.generate(30)
.unwrap();
assert!(
folds.len() < 5,
"large embargo should reduce folds, got {}",
folds.len()
);
}
#[test]
fn i1_error_series_too_short() {
let result = CvFoldGenerator::new()
.n_folds(5)
.horizon(5)
.min_initial_window(10)
.generate(10); assert!(result.is_err());
}
#[test]
fn i2_error_with_gap_purge() {
let result = CvFoldGenerator::new()
.n_folds(3)
.horizon(3)
.min_initial_window(5)
.gap(3)
.purge(2)
.generate(12); assert!(result.is_err());
}
#[test]
fn j1_reduces_to_1_fold() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(3)
.min_initial_window(7)
.on_constraint_violation(ConstraintViolation::ReduceFolds)
.generate(10)
.unwrap();
assert_eq!(folds.len(), 1);
assert_eq!(
folds[0],
Fold {
train_start: 0,
train_end: 7,
test_start: 7,
test_end: 10
}
);
}
#[test]
fn j2_no_reduction_when_fits() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(5)
.min_initial_window(10)
.on_constraint_violation(ConstraintViolation::ReduceFolds)
.generate(100)
.unwrap();
assert_eq!(folds.len(), 3);
}
#[test]
fn k1_horizon_larger_than_half_series() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(15)
.min_initial_window(5)
.generate(25)
.unwrap();
assert_eq!(folds.len(), 1);
assert_eq!(folds[0].test_end, 25);
}
#[test]
fn k2_default_values() {
let gen = CvFoldGenerator::new();
assert_eq!(gen.target_n_folds, 5);
assert_eq!(gen.on_constraint_violation, ConstraintViolation::Error);
assert_eq!(gen.min_initial_window, 10);
assert_eq!(gen.horizon, 1);
}
#[test]
fn k3_series_exactly_min_plus_horizon() {
let folds = CvFoldGenerator::new()
.n_folds(1)
.horizon(3)
.min_initial_window(7)
.generate(10)
.unwrap();
assert_eq!(folds.len(), 1);
assert_eq!(
folds[0],
Fold {
train_start: 0,
train_end: 7,
test_start: 7,
test_end: 10
}
);
}
#[test]
fn l1_sklearn_5_splits_h1_on_10() {
let folds = CvFoldGenerator::new()
.n_folds(5)
.horizon(1)
.min_initial_window(5)
.generate(10)
.unwrap();
assert_eq!(folds.len(), 5);
assert_eq!(folds[0].test_start, 5);
assert_eq!(folds[4].test_end, 10);
for i in 1..folds.len() {
assert_eq!(folds[i].test_start, folds[i - 1].test_end);
}
}
#[test]
fn l2_sklearn_3_splits_h2_on_20() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(2)
.min_initial_window(5)
.generate(20)
.unwrap();
assert_eq!(folds.len(), 3);
assert_eq!(
folds[0],
Fold {
train_start: 0,
train_end: 14,
test_start: 14,
test_end: 16
}
);
assert_eq!(
folds[1],
Fold {
train_start: 0,
train_end: 16,
test_start: 16,
test_end: 18
}
);
assert_eq!(
folds[2],
Fold {
train_start: 0,
train_end: 18,
test_start: 18,
test_end: 20
}
);
}
#[test]
fn m1_step_size_smaller_than_horizon_overlapping_tests() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(5)
.step_size(2)
.min_initial_window(5)
.generate(20)
.unwrap();
assert_eq!(folds.len(), 3);
assert_eq!(folds.last().unwrap().test_end, 20);
for fold in &folds {
assert_eq!(fold.test_size(), 5);
}
assert!(folds[1].test_start < folds[0].test_end);
}
#[test]
fn m2_step_size_larger_than_horizon_sparse() {
let folds = CvFoldGenerator::new()
.n_folds(3)
.horizon(3)
.step_size(10)
.min_initial_window(5)
.generate(50)
.unwrap();
assert_eq!(folds.len(), 3);
assert_eq!(folds.last().unwrap().test_end, 50);
assert!(folds[1].test_start > folds[0].test_end);
}
#[test]
fn m3_step_size_equals_horizon_contiguous() {
let folds_default = CvFoldGenerator::new()
.n_folds(4)
.horizon(5)
.min_initial_window(10)
.generate(50)
.unwrap();
let folds_explicit = CvFoldGenerator::new()
.n_folds(4)
.horizon(5)
.step_size(5)
.min_initial_window(10)
.generate(50)
.unwrap();
assert_eq!(folds_default, folds_explicit);
}
#[test]
fn m4_step_size_1_maximum_folds() {
let folds = CvFoldGenerator::new()
.n_folds(100)
.horizon(3)
.step_size(1)
.min_initial_window(5)
.generate(20)
.unwrap();
assert!(folds.len() > 5);
assert_eq!(folds.last().unwrap().test_end, 20);
}
#[test]
fn m5_default_step_is_none() {
assert!(CvFoldGenerator::new().step_size.is_none());
}
}