use super::fingerprint::ForecastabilityFingerprint;
use super::scorers::{score, Scorer};
#[cfg(feature = "parallel")]
use rayon::prelude::*;
#[cfg(feature = "postprocess")]
use crate::validation::aid::{AidAnalyzer, AidDemandType};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum SeriesPattern {
WhiteNoise,
Linear,
Seasonal,
Nonlinear,
Complex,
Intermittent,
}
impl std::fmt::Display for SeriesPattern {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::WhiteNoise => write!(f, "A: White noise"),
Self::Linear => write!(f, "B: Linear / AR-like"),
Self::Seasonal => write!(f, "C: Seasonal / periodic"),
Self::Nonlinear => write!(f, "D: Nonlinear deterministic"),
Self::Complex => write!(f, "E: Complex / mixed"),
Self::Intermittent => write!(f, "F: Intermittent demand"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum ModelFamily {
Skip,
LinearStatistical,
SeasonalStatistical,
NonlinearML,
Ensemble,
Intermittent,
}
impl std::fmt::Display for ModelFamily {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Skip => write!(f, "Skip (Naive)"),
Self::LinearStatistical => write!(f, "ARIMA / ETS / Theta"),
Self::SeasonalStatistical => write!(f, "SeasonalARIMA / ETS(seasonal) / MSTL"),
Self::NonlinearML => write!(f, "MFLES / RegressionForecaster / tree-based"),
Self::Ensemble => write!(f, "AutoForecast / Ensemble"),
Self::Intermittent => write!(f, "Croston / TSB / ADIDA / IMAPA"),
}
}
}
#[derive(Debug, Clone)]
pub struct ExogenousScore {
pub index: usize,
pub best_lag: usize,
pub te_at_best_lag: f64,
pub te_curve: Vec<f64>,
}
#[derive(Debug, Clone)]
pub struct TriageResult {
pub pattern: SeriesPattern,
pub model_family: ModelFamily,
pub fingerprint: Option<ForecastabilityFingerprint>,
pub permutation_entropy: f64,
pub spectral_predictability: f64,
pub recommended_lags: Vec<usize>,
pub is_intermittent: bool,
pub aid_demand_type: Option<String>,
}
#[derive(Debug, Clone)]
pub struct BatchTriageResult {
pub results: Vec<TriageResult>,
pub pattern_counts: [(SeriesPattern, usize); 6],
pub family_counts: [(ModelFamily, usize); 6],
}
#[derive(Debug, Clone)]
pub struct TriageConfig {
pub max_lag: usize,
pub n_surrogates: usize,
pub alpha: f64,
pub seed: Option<u64>,
}
impl Default for TriageConfig {
fn default() -> Self {
Self {
max_lag: 20,
n_surrogates: 50,
alpha: 0.05,
seed: None,
}
}
}
impl TriageConfig {
pub fn max_lag(mut self, v: usize) -> Self {
self.max_lag = v;
self
}
pub fn n_surrogates(mut self, v: usize) -> Self {
self.n_surrogates = v;
self
}
pub fn alpha(mut self, v: f64) -> Self {
self.alpha = v;
self
}
pub fn seed(mut self, v: u64) -> Self {
self.seed = Some(v);
self
}
}
fn classify_pattern(fp: &ForecastabilityFingerprint, pe: f64) -> SeriesPattern {
if fp.informative_horizons.is_empty() || fp.signal_to_noise < 1.5 {
if pe > 0.9 || fp.information_mass < 0.01 {
return SeriesPattern::WhiteNoise;
}
}
if fp.nonlinear_share > 0.5 && fp.signal_to_noise > 2.0 {
return SeriesPattern::Nonlinear;
}
if fp.nonlinear_share < 0.3 && fp.directness_ratio > 0.3 {
return SeriesPattern::Linear;
}
if fp.information_structure > 0.6
&& fp.informative_horizons.len() >= 3
&& fp.nonlinear_share < 0.5
{
return SeriesPattern::Seasonal;
}
SeriesPattern::Complex
}
fn recommend_family(pattern: SeriesPattern) -> ModelFamily {
match pattern {
SeriesPattern::WhiteNoise => ModelFamily::Skip,
SeriesPattern::Linear => ModelFamily::LinearStatistical,
SeriesPattern::Seasonal => ModelFamily::SeasonalStatistical,
SeriesPattern::Nonlinear => ModelFamily::NonlinearML,
SeriesPattern::Complex => ModelFamily::Ensemble,
SeriesPattern::Intermittent => ModelFamily::Intermittent,
}
}
fn check_intermittent(series: &[f64]) -> (bool, Option<String>) {
let n = series.len();
if n == 0 {
return (false, None);
}
let zero_count = series.iter().filter(|&&v| v.abs() < 1e-10).count();
let zero_fraction = zero_count as f64 / n as f64;
if zero_fraction <= 0.3 {
return (false, None);
}
#[cfg(feature = "postprocess")]
{
let result = AidAnalyzer::new().analyze(series);
let summary = result.summary();
let demand_type_str = format!("{:?}", summary.demand_type);
let is_intermittent = matches!(summary.demand_type, AidDemandType::Intermittent);
(
is_intermittent || zero_fraction > 0.5,
Some(demand_type_str),
)
}
#[cfg(not(feature = "postprocess"))]
{
(true, None)
}
}
pub fn run_triage(series: &[f64], config: &TriageConfig) -> TriageResult {
let pe = score(series, Scorer::PermutationEntropy);
let sp = score(series, Scorer::SpectralPredictability);
let (is_intermittent, aid_demand_type) = check_intermittent(series);
if is_intermittent {
return TriageResult {
pattern: SeriesPattern::Intermittent,
model_family: ModelFamily::Intermittent,
fingerprint: None, permutation_entropy: pe,
spectral_predictability: sp,
recommended_lags: vec![],
is_intermittent: true,
aid_demand_type,
};
}
let fp = ForecastabilityFingerprint::compute(
series,
config.max_lag,
config.n_surrogates,
config.alpha,
config.seed,
);
let pattern = classify_pattern(&fp, pe);
let model_family = recommend_family(pattern);
let recommended_lags = fp.informative_horizons.clone();
TriageResult {
pattern,
model_family,
fingerprint: Some(fp),
permutation_entropy: pe,
spectral_predictability: sp,
recommended_lags,
is_intermittent: false,
aid_demand_type,
}
}
pub fn run_batch_triage(all_series: &[Vec<f64>], config: &TriageConfig) -> BatchTriageResult {
#[cfg(feature = "parallel")]
let results: Vec<TriageResult> = all_series
.par_iter()
.map(|s| run_triage(s, config))
.collect();
#[cfg(not(feature = "parallel"))]
let results: Vec<TriageResult> = all_series.iter().map(|s| run_triage(s, config)).collect();
let mut pattern_counts = [
(SeriesPattern::WhiteNoise, 0),
(SeriesPattern::Linear, 0),
(SeriesPattern::Seasonal, 0),
(SeriesPattern::Nonlinear, 0),
(SeriesPattern::Complex, 0),
(SeriesPattern::Intermittent, 0),
];
let mut family_counts = [
(ModelFamily::Skip, 0),
(ModelFamily::LinearStatistical, 0),
(ModelFamily::SeasonalStatistical, 0),
(ModelFamily::NonlinearML, 0),
(ModelFamily::Ensemble, 0),
(ModelFamily::Intermittent, 0),
];
for r in &results {
for pc in &mut pattern_counts {
if pc.0 == r.pattern {
pc.1 += 1;
}
}
for fc in &mut family_counts {
if fc.0 == r.model_family {
fc.1 += 1;
}
}
}
BatchTriageResult {
results,
pattern_counts,
family_counts,
}
}
pub fn screen_exogenous(
target: &[f64],
candidates: &[Vec<f64>],
max_lag: usize,
) -> Vec<ExogenousScore> {
let mut scores: Vec<ExogenousScore> = candidates
.iter()
.enumerate()
.map(|(i, cand)| {
let te_curve = super::transfer_entropy::transfer_entropy_curve(cand, target, max_lag);
let (best_lag, te_at_best_lag) = te_curve
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(lag_idx, &te)| (lag_idx + 1, te)) .unwrap_or((1, 0.0));
ExogenousScore {
index: i,
best_lag,
te_at_best_lag,
te_curve,
}
})
.collect();
scores.sort_by(|a, b| b.te_at_best_lag.partial_cmp(&a.te_at_best_lag).unwrap());
scores
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
fn make_white_noise(n: usize, seed: u64) -> Vec<f64> {
let mut rng = StdRng::seed_from_u64(seed);
(0..n)
.map(|_| {
let u1: f64 = rng.gen::<f64>().max(f64::MIN_POSITIVE);
let u2: f64 = rng.gen();
(-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
})
.collect()
}
fn make_logistic(n: usize) -> Vec<f64> {
let mut x = vec![0.0; n];
x[0] = 0.1;
for t in 1..n {
x[t] = 3.9 * x[t - 1] * (1.0 - x[t - 1]);
}
x
}
#[test]
fn triage_white_noise_classifies_a() {
let series = make_white_noise(500, 42);
let result = run_triage(&series, &TriageConfig::default().seed(1));
assert_eq!(
result.pattern,
SeriesPattern::WhiteNoise,
"white noise should be pattern A, got {}",
result.pattern
);
assert_eq!(result.model_family, ModelFamily::Skip);
}
#[test]
fn triage_logistic_map_classifies_nonlinear() {
let series = make_logistic(1000);
let result = run_triage(&series, &TriageConfig::default().seed(1));
assert!(
result.pattern == SeriesPattern::Nonlinear || result.pattern == SeriesPattern::Complex,
"logistic map should be pattern D or E, got {}",
result.pattern
);
assert!(
result.model_family == ModelFamily::NonlinearML
|| result.model_family == ModelFamily::Ensemble,
);
}
#[test]
fn batch_triage_counts_match() {
let series = vec![
make_white_noise(300, 1),
make_white_noise(300, 2),
make_logistic(500),
];
let config = TriageConfig::default()
.max_lag(10)
.n_surrogates(30)
.seed(42);
let batch = run_batch_triage(&series, &config);
assert_eq!(batch.results.len(), 3);
let total: usize = batch.pattern_counts.iter().map(|(_, c)| c).sum();
assert_eq!(total, 3);
}
#[test]
fn triage_intermittent_classifies_f() {
let mut series = vec![0.0; 70];
series.extend(vec![5.0, 0.0, 12.0, 0.0, 0.0, 8.0, 0.0, 3.0, 0.0, 0.0]);
series.extend(vec![0.0; 220]);
let result = run_triage(&series, &TriageConfig::default().seed(1));
assert_eq!(
result.pattern,
SeriesPattern::Intermittent,
"70% zeros should be pattern F, got {}",
result.pattern
);
assert_eq!(result.model_family, ModelFamily::Intermittent);
assert!(result.is_intermittent);
assert!(
result.fingerprint.is_none(),
"fingerprint should be skipped for intermittent"
);
}
#[test]
fn triage_result_has_recommended_lags() {
let series = make_logistic(1000);
let result = run_triage(&series, &TriageConfig::default().seed(1));
assert!(
!result.recommended_lags.is_empty(),
"logistic map should have recommended lags"
);
for &lag in &result.recommended_lags {
assert!((1..=20).contains(&lag), "lag {} out of range", lag);
}
}
#[test]
fn screen_exogenous_ranks_driver_first() {
let mut rng = StdRng::seed_from_u64(42);
let n = 300;
let driver: Vec<f64> = (0..n).map(|_| (rng.gen::<f64>() - 0.5) * 2.0).collect();
let mut target = vec![0.0; n];
for t in 1..n {
target[t] = 0.7 * driver[t - 1] + (rng.gen::<f64>() - 0.5) * 0.5;
}
let noise: Vec<f64> = (0..n).map(|_| (rng.gen::<f64>() - 0.5) * 2.0).collect();
let scores = screen_exogenous(&target, &[driver, noise], 3);
assert_eq!(scores[0].index, 0, "driver should rank first");
assert!(
scores[0].te_at_best_lag > scores[1].te_at_best_lag,
"driver TE ({:.4}) should exceed noise TE ({:.4})",
scores[0].te_at_best_lag,
scores[1].te_at_best_lag
);
assert_eq!(
scores[0].best_lag, 1,
"driver best lag should be 1, got {}",
scores[0].best_lag
);
}
}