use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::exponential::AutoETS;
use crate::models::theta::AutoTheta;
use crate::models::traits::{validate_series_complete, FittedParams, Forecaster};
use crate::utils::ols::OLSResult;
use std::collections::HashMap;
#[cfg(all(feature = "distributional", feature = "postprocess"))]
use crate::models::laplace::LaplaceForecaster;
#[cfg(feature = "postprocess")]
use crate::validation::aid::AidAnalyzer;
#[cfg(feature = "postprocess")]
use anofox_regression::solvers::{DemandDistribution, DemandType};
const MIN_HISTORY_FOR_LAPLACE: usize = 60;
const TREND_R2_TRIGGER: f64 = 0.30;
const SEASONAL_AUTOCORR_TRIGGER: f64 = 0.40;
#[derive(Debug, Clone, PartialEq)]
pub enum SelectedFamily {
IntermittentPoisson,
IntermittentNegBinomial,
IntermittentRectifiedNormal,
IntermittentPositive,
RegularCount,
RegularPositive,
RegularNormal,
AutoETSStructural,
AutoThetaShortHistory,
Fallback,
}
pub struct SmartForecaster {
inner: Option<Box<dyn Forecaster + Send>>,
selected: Option<SelectedFamily>,
seasonal_period: usize,
}
impl SmartForecaster {
pub fn new() -> Self {
Self {
inner: None,
selected: None,
seasonal_period: 7,
}
}
pub fn with_seasonal_period(mut self, period: usize) -> Self {
self.seasonal_period = period.max(2);
self
}
pub fn selected_family(&self) -> Option<&SelectedFamily> {
self.selected.as_ref()
}
fn commit(
&mut self,
mut inner: Box<dyn Forecaster + Send>,
selected: SelectedFamily,
series: &TimeSeries,
) -> Result<()> {
inner.fit(series)?;
self.inner = Some(inner);
self.selected = Some(selected);
Ok(())
}
}
impl Default for SmartForecaster {
fn default() -> Self {
Self::new()
}
}
fn trend_r_squared(values: &[f64]) -> f64 {
let n = values.len();
if n < 3 {
return 0.0;
}
let n_f = n as f64;
let mean_t = (n_f - 1.0) / 2.0;
let mean_y = values.iter().sum::<f64>() / n_f;
let ss_tot: f64 = values.iter().map(|v| (v - mean_y).powi(2)).sum();
if ss_tot < 1e-9 {
return 0.0;
}
let num: f64 = values
.iter()
.enumerate()
.map(|(i, y)| (i as f64 - mean_t) * (y - mean_y))
.sum();
let den: f64 = (0..n).map(|i| (i as f64 - mean_t).powi(2)).sum();
if den < 1e-9 {
return 0.0;
}
let slope = num / den;
let intercept = mean_y - slope * mean_t;
let ss_res: f64 = values
.iter()
.enumerate()
.map(|(i, y)| (y - intercept - slope * i as f64).powi(2))
.sum();
(1.0 - ss_res / ss_tot).clamp(0.0, 1.0)
}
fn seasonal_autocorr_abs(values: &[f64], lag: usize) -> f64 {
let n = values.len();
if lag == 0 || lag >= n || n < lag + 3 {
return 0.0;
}
let n_f = n as f64;
let mean = values.iter().sum::<f64>() / n_f;
let var: f64 = values.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / n_f;
if var < 1e-9 {
return 0.0;
}
let cov: f64 = (lag..n)
.map(|i| (values[i] - mean) * (values[i - lag] - mean))
.sum::<f64>()
/ (n - lag) as f64;
(cov / var).abs()
}
fn is_autoets_favorable(values: &[f64], seasonal_period: usize) -> bool {
if values.len() < MIN_HISTORY_FOR_LAPLACE {
return false;
}
let has_trend = trend_r_squared(values) > TREND_R2_TRIGGER;
let has_seasonal = seasonal_period >= 2
&& seasonal_autocorr_abs(values, seasonal_period) > SEASONAL_AUTOCORR_TRIGGER;
has_trend || has_seasonal
}
#[cfg(all(feature = "distributional", feature = "postprocess"))]
fn build_from_aid(
demand_type: DemandType,
distribution: DemandDistribution,
seasonal_period: usize,
) -> (Box<dyn Forecaster + Send>, SelectedFamily) {
match (demand_type, distribution) {
(DemandType::Intermittent, DemandDistribution::Poisson)
| (DemandType::Intermittent, DemandDistribution::Geometric) => {
let m = LaplaceForecaster::new()
.with_poisson_defaults()
.with_seasonal_intermittent_defaults(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::IntermittentPoisson)
}
(DemandType::Intermittent, DemandDistribution::NegativeBinomial) => {
let m = LaplaceForecaster::new()
.with_negative_binomial_defaults()
.with_seasonal_intermittent_defaults(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::IntermittentNegBinomial)
}
(DemandType::Intermittent, DemandDistribution::RectifiedNormal) => {
let m = LaplaceForecaster::new()
.with_rectified_normal_defaults()
.with_seasonal_intermittent_defaults(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::IntermittentRectifiedNormal)
}
(DemandType::Intermittent, DemandDistribution::LogNormal) => {
let m = LaplaceForecaster::new()
.with_lognormal_defaults()
.with_seasonal_intermittent_defaults(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::IntermittentPositive)
}
(DemandType::Intermittent, DemandDistribution::Gamma) => {
let m = LaplaceForecaster::new()
.with_gamma_defaults()
.with_seasonal_intermittent_defaults(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::IntermittentPositive)
}
(DemandType::Intermittent, DemandDistribution::Normal) => {
let m = LaplaceForecaster::new()
.with_intermittent_defaults()
.non_negative()
.auto();
(Box::new(m), SelectedFamily::IntermittentPositive)
}
(DemandType::Regular, DemandDistribution::Poisson)
| (DemandType::Regular, DemandDistribution::Geometric) => {
let m = LaplaceForecaster::new()
.with_poisson_defaults()
.with_seasonal(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::RegularCount)
}
(DemandType::Regular, DemandDistribution::NegativeBinomial) => {
let m = LaplaceForecaster::new()
.with_negative_binomial_defaults()
.with_seasonal(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::RegularCount)
}
(DemandType::Regular, DemandDistribution::LogNormal) => {
let m = LaplaceForecaster::new()
.with_lognormal_defaults()
.with_seasonal(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::RegularPositive)
}
(DemandType::Regular, DemandDistribution::Gamma) => {
let m = LaplaceForecaster::new()
.with_gamma_defaults()
.with_seasonal(seasonal_period)
.non_negative();
(Box::new(m), SelectedFamily::RegularPositive)
}
(DemandType::Regular, DemandDistribution::RectifiedNormal) => {
let m = LaplaceForecaster::new()
.with_rectified_normal_defaults()
.non_negative();
(Box::new(m), SelectedFamily::RegularPositive)
}
(DemandType::Regular, DemandDistribution::Normal) => {
let m = LaplaceForecaster::new().auto();
(Box::new(m), SelectedFamily::RegularNormal)
}
}
}
impl Forecaster for SmartForecaster {
#[allow(unused_variables)] fn fit(&mut self, series: &TimeSeries) -> Result<()> {
validate_series_complete(series)?;
let values = series.primary_values();
if values.is_empty() {
return Err(ForecastError::InvalidParameter(
"SmartForecaster requires at least one observation".into(),
));
}
let period_for_classical = if self.seasonal_period >= 2
&& seasonal_autocorr_abs(values, self.seasonal_period) > SEASONAL_AUTOCORR_TRIGGER
{
Some(self.seasonal_period)
} else {
None
};
if values.len() < MIN_HISTORY_FOR_LAPLACE {
let m: Box<dyn Forecaster + Send> = match period_for_classical {
Some(p) => Box::new(AutoTheta::seasonal(p)),
None => Box::new(AutoTheta::new()),
};
return self.commit(m, SelectedFamily::AutoThetaShortHistory, series);
}
#[cfg(all(feature = "distributional", feature = "postprocess"))]
let (inner, selected) = {
let aid = AidAnalyzer::new().analyze(values);
let summary = aid.summary();
if matches!(summary.demand_type, DemandType::Regular)
&& is_autoets_favorable(values, self.seasonal_period)
{
let m: Box<dyn Forecaster + Send> = match period_for_classical {
Some(p) => Box::new(AutoETS::with_period(p)),
None => Box::new(AutoETS::new()),
};
(m, SelectedFamily::AutoETSStructural)
} else {
build_from_aid(
summary.demand_type,
summary.distribution,
self.seasonal_period,
)
}
};
#[cfg(not(all(feature = "distributional", feature = "postprocess")))]
let (inner, selected): (Box<dyn Forecaster + Send>, SelectedFamily) = {
(
Box::new(crate::models::theta::AutoTheta::new()),
SelectedFamily::Fallback,
)
};
self.commit(inner, selected, series)
}
fn predict(&self, horizon: usize) -> Result<Forecast> {
match &self.inner {
Some(m) => m.predict(horizon),
None => Err(ForecastError::FitRequired {
model: Some("SmartForecaster".into()),
}),
}
}
fn predict_with_intervals(&self, horizon: usize, level: f64) -> Result<Forecast> {
match &self.inner {
Some(m) => m.predict_with_intervals(horizon, level),
None => Err(ForecastError::FitRequired {
model: Some("SmartForecaster".into()),
}),
}
}
fn fitted_values(&self) -> Option<&[f64]> {
self.inner.as_ref().and_then(|m| m.fitted_values())
}
fn residuals(&self) -> Option<&[f64]> {
self.inner.as_ref().and_then(|m| m.residuals())
}
fn training_values(&self) -> Result<&[f64]> {
match &self.inner {
Some(m) => m.training_values(),
None => Err(ForecastError::FitRequired {
model: Some("SmartForecaster".into()),
}),
}
}
fn fitted_params(&self) -> Option<FittedParams> {
self.inner.as_ref().and_then(|m| m.fitted_params())
}
fn training_regressors(&self) -> Option<&HashMap<String, Vec<f64>>> {
self.inner.as_ref().and_then(|m| m.training_regressors())
}
fn exog_coefficients(&self) -> Option<&OLSResult> {
self.inner.as_ref().and_then(|m| m.exog_coefficients())
}
fn name(&self) -> &str {
"SmartForecaster"
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, TimeZone, Utc};
fn make_ts(vals: Vec<f64>) -> TimeSeries {
let base = Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap();
let stamps: Vec<_> = (0..vals.len())
.map(|i| base + Duration::hours(i as i64))
.collect();
TimeSeries::univariate(stamps, vals).unwrap()
}
#[cfg(all(feature = "distributional", feature = "postprocess"))]
#[test]
fn intermittent_count_series_routes_to_intermittent_family() {
let mut vals = vec![0.0; 200];
for i in (0..200).step_by(3) {
vals[i] = 2.0;
}
let ts = make_ts(vals);
let mut f = SmartForecaster::new();
f.fit(&ts).unwrap();
assert!(matches!(
f.selected_family(),
Some(SelectedFamily::IntermittentPoisson)
| Some(SelectedFamily::IntermittentNegBinomial)
| Some(SelectedFamily::IntermittentRectifiedNormal)
| Some(SelectedFamily::IntermittentPositive)
));
let fc = f.predict(10).unwrap();
for v in fc.primary() {
assert!(*v >= 0.0);
}
}
#[cfg(all(feature = "distributional", feature = "postprocess"))]
#[test]
fn continuous_regular_normal_routes_to_auto_or_ets() {
let vals: Vec<f64> = (0..200)
.map(|i| 50.0 + (i as f64 * 0.05).sin() * 5.0)
.collect();
let ts = make_ts(vals);
let mut f = SmartForecaster::new();
f.fit(&ts).unwrap();
assert!(matches!(
f.selected_family(),
Some(SelectedFamily::RegularNormal)
| Some(SelectedFamily::RegularPositive)
| Some(SelectedFamily::RegularCount)
| Some(SelectedFamily::AutoETSStructural)
));
}
#[cfg(all(feature = "distributional", feature = "postprocess"))]
#[test]
fn short_history_routes_to_autotheta() {
let vals: Vec<f64> = (0..40).map(|i| 50.0 + i as f64 * 0.1).collect();
let ts = make_ts(vals);
let mut f = SmartForecaster::new();
f.fit(&ts).unwrap();
assert_eq!(
f.selected_family(),
Some(&SelectedFamily::AutoThetaShortHistory)
);
let fc = f.predict(6).unwrap();
assert_eq!(fc.primary().len(), 6);
}
#[cfg(all(feature = "distributional", feature = "postprocess"))]
#[test]
fn structural_trend_routes_to_auto_ets() {
let vals: Vec<f64> = (0..200).map(|i| 50.0 + 0.5 * i as f64).collect();
let ts = make_ts(vals);
let mut f = SmartForecaster::new();
f.fit(&ts).unwrap();
assert_eq!(
f.selected_family(),
Some(&SelectedFamily::AutoETSStructural)
);
}
#[cfg(all(feature = "distributional", feature = "postprocess"))]
#[test]
fn random_walk_does_not_route_to_auto_ets() {
let mut x = 50.0;
let mut rng_state: u64 = 42;
let vals: Vec<f64> = (0..300)
.map(|_| {
rng_state = rng_state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let u1 = ((rng_state >> 33) as f64) / (u32::MAX as f64);
rng_state = rng_state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let u2 = ((rng_state >> 33) as f64) / (u32::MAX as f64);
let z = (-2.0 * u1.max(1e-9).ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
x += z;
x
})
.collect();
let ts = make_ts(vals);
let mut f = SmartForecaster::new().with_seasonal_period(12);
f.fit(&ts).unwrap();
let fc = f.predict(10).unwrap();
assert_eq!(fc.primary().len(), 10);
}
#[test]
fn trend_r_squared_close_to_one_on_pure_trend() {
let vals: Vec<f64> = (0..100).map(|i| i as f64).collect();
assert!((trend_r_squared(&vals) - 1.0).abs() < 1e-6);
}
#[test]
fn trend_r_squared_near_zero_on_pure_noise() {
let mut s: u64 = 123;
let vals: Vec<f64> = (0..200)
.map(|_| {
s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((s >> 33) as f64) / (u32::MAX as f64) - 0.5
})
.collect();
assert!(trend_r_squared(&vals) < 0.05);
}
}