use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::arima::diff::suggest_differencing;
use crate::models::arima::model::{ARIMA, SARIMA};
use crate::models::inspect::{ArimaExplanation, Explanation, Inspectable};
use crate::models::{validate_series_complete, Forecaster};
use crate::utils::ols::OLSResult;
#[cfg(feature = "parallel")]
use rayon::prelude::*;
#[derive(Debug, Clone)]
pub struct AutoARIMAConfig {
pub max_p: usize,
pub max_q: usize,
pub max_d: usize,
pub max_cap_p: usize,
pub max_cap_q: usize,
pub max_cap_d: usize,
pub seasonal_period: usize,
pub stepwise: bool,
pub true_stepwise: bool,
pub use_aic: bool,
}
impl Default for AutoARIMAConfig {
fn default() -> Self {
Self {
max_p: 5,
max_q: 5,
max_d: 2,
max_cap_p: 1,
max_cap_q: 1,
max_cap_d: 1,
seasonal_period: 0,
stepwise: true,
true_stepwise: false, use_aic: true,
}
}
}
impl AutoARIMAConfig {
pub fn with_max_orders(mut self, max_p: usize, max_d: usize, max_q: usize) -> Self {
self.max_p = max_p;
self.max_d = max_d;
self.max_q = max_q;
self
}
pub fn with_seasonal_orders(mut self, max_p: usize, max_d: usize, max_q: usize) -> Self {
self.max_cap_p = max_p;
self.max_cap_d = max_d;
self.max_cap_q = max_q;
self
}
pub fn with_seasonal_period(mut self, period: usize) -> Self {
self.seasonal_period = period;
self
}
pub fn exhaustive(mut self) -> Self {
self.stepwise = false;
self
}
pub fn with_true_stepwise(mut self) -> Self {
self.stepwise = true;
self.true_stepwise = true;
self
}
}
#[derive(Debug, Clone)]
enum SelectedModel {
ARIMA(ARIMA),
SARIMA(SARIMA),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ModelOrder {
pub p: usize,
pub d: usize,
pub q: usize,
pub cap_p: usize,
pub cap_d: usize,
pub cap_q: usize,
pub s: usize,
}
impl ModelOrder {
pub fn is_seasonal(&self) -> bool {
self.s > 1 && (self.cap_p > 0 || self.cap_d > 0 || self.cap_q > 0)
}
}
#[derive(Debug, Clone)]
pub struct AutoARIMA {
config: AutoARIMAConfig,
selected_model: Option<SelectedModel>,
selected_order: Option<ModelOrder>,
model_scores: Vec<(ModelOrder, f64)>,
training_values_store: Option<Vec<f64>>,
training_regressors_store: Option<std::collections::HashMap<String, Vec<f64>>>,
}
impl AutoARIMA {
pub fn new() -> Self {
Self {
config: AutoARIMAConfig::default(),
selected_model: None,
selected_order: None,
model_scores: Vec::new(),
training_values_store: None,
training_regressors_store: None,
}
}
pub fn with_config(config: AutoARIMAConfig) -> Self {
Self {
config,
selected_model: None,
selected_order: None,
model_scores: Vec::new(),
training_values_store: None,
training_regressors_store: None,
}
}
pub fn seasonal(period: usize) -> Self {
let config = AutoARIMAConfig::default().with_seasonal_period(period);
Self::with_config(config)
}
pub fn selected_order(&self) -> Option<(usize, usize, usize)> {
self.selected_order.map(|o| (o.p, o.d, o.q))
}
pub fn selected_full_order(&self) -> Option<ModelOrder> {
self.selected_order
}
pub fn model_scores(&self) -> &[(ModelOrder, f64)] {
&self.model_scores
}
fn has_seasonal_pattern(values: &[f64], period: usize) -> bool {
use crate::seasonality::STL;
if period < 2 || values.len() < 3 * period {
return false;
}
let stl = STL::new(period);
if let Some(result) = stl.decompose(values) {
result.seasonal_strength() > 0.64
} else {
false
}
}
fn suggest_seasonal_differencing(values: &[f64], period: usize) -> usize {
if period < 2 || values.len() < 2 * period {
return 0;
}
let seasonal_diffs: Vec<f64> = (period..values.len())
.map(|i| values[i] - values[i - period])
.collect();
let orig_mean = values.iter().sum::<f64>() / values.len() as f64;
let orig_var =
values.iter().map(|v| (v - orig_mean).powi(2)).sum::<f64>() / values.len() as f64;
let diff_mean = seasonal_diffs.iter().sum::<f64>() / seasonal_diffs.len() as f64;
let diff_var = seasonal_diffs
.iter()
.map(|v| (v - diff_mean).powi(2))
.sum::<f64>()
/ seasonal_diffs.len() as f64;
if diff_var < orig_var * 0.7 {
1
} else {
0
}
}
fn stepwise_candidates(&self, d: usize, cap_d: usize) -> Vec<ModelOrder> {
let s = self.config.seasonal_period;
let nonseasonal = vec![
(0, 0),
(1, 0),
(0, 1),
(1, 1),
(2, 0),
(0, 2),
(2, 1),
(1, 2),
(2, 2),
(3, 0),
(0, 3),
(3, 1),
(1, 3),
(3, 2),
(2, 3),
];
let mut candidates = Vec::new();
for &(p, q) in &nonseasonal {
if p <= self.config.max_p && q <= self.config.max_q {
candidates.push(ModelOrder {
p,
d,
q,
cap_p: 0,
cap_d,
cap_q: 0,
s,
});
}
}
if s > 1 {
let seasonal = vec![
(0, 1),
(1, 0),
(1, 1),
(2, 0),
(0, 2),
(2, 1),
(1, 2),
(2, 2),
];
let nonseasonal_with_seasonal = vec![
(0, 0),
(1, 0),
(0, 1),
(1, 1),
(2, 0),
(0, 2),
(2, 1),
(1, 2),
(3, 0),
(0, 3),
(2, 2),
(3, 1),
(1, 3),
];
for &(p, q) in &nonseasonal_with_seasonal {
for &(cap_p, cap_q) in &seasonal {
if p <= self.config.max_p
&& q <= self.config.max_q
&& cap_p <= self.config.max_cap_p
&& cap_q <= self.config.max_cap_q
{
candidates.push(ModelOrder {
p,
d,
q,
cap_p,
cap_d,
cap_q,
s,
});
}
}
}
}
candidates
}
fn exhaustive_candidates(&self, d: usize, cap_d: usize) -> Vec<ModelOrder> {
let s = self.config.seasonal_period;
let mut candidates = Vec::new();
for p in 0..=self.config.max_p {
for q in 0..=self.config.max_q {
if s > 1 {
for cap_p in 0..=self.config.max_cap_p {
for cap_q in 0..=self.config.max_cap_q {
candidates.push(ModelOrder {
p,
d,
q,
cap_p,
cap_d,
cap_q,
s,
});
}
}
} else {
candidates.push(ModelOrder {
p,
d,
q,
cap_p: 0,
cap_d: 0,
cap_q: 0,
s: 0,
});
}
}
}
candidates
}
fn evaluate_model_static(
series: &TimeSeries,
order: ModelOrder,
use_aic: bool,
) -> Option<(SelectedModel, f64)> {
if order.is_seasonal() {
let mut model = SARIMA::new(
order.p,
order.d,
order.q,
order.cap_p,
order.cap_d,
order.cap_q,
order.s,
);
if model.fit(series).is_ok() {
let score = if use_aic { model.aic() } else { model.bic() };
if let Some(s) = score {
if s.is_finite() {
return Some((SelectedModel::SARIMA(model), s));
}
}
}
} else {
let mut model = ARIMA::new(order.p, order.d, order.q);
if model.fit(series).is_ok() {
let score = if use_aic { model.aic() } else { model.bic() };
if let Some(s) = score {
if s.is_finite() {
return Some((SelectedModel::ARIMA(model), s));
}
}
}
}
None
}
fn score_order_static(order: ModelOrder, diff_series: &[f64], use_aic: bool) -> Option<f64> {
if order.is_seasonal() {
SARIMA::score_order(
order.p,
order.q,
order.cap_p,
order.cap_q,
order.s,
diff_series,
use_aic,
)
} else {
ARIMA::score_order(order.p, order.q, diff_series, use_aic)
}
}
fn get_neighbors(&self, order: ModelOrder) -> Vec<ModelOrder> {
let mut neighbors = Vec::new();
let s = self.config.seasonal_period;
if s > 1 {
if order.cap_p > 0 {
neighbors.push(ModelOrder {
cap_p: order.cap_p - 1,
..order
});
}
if order.cap_q > 0 {
neighbors.push(ModelOrder {
cap_q: order.cap_q - 1,
..order
});
}
if order.cap_p < self.config.max_cap_p {
neighbors.push(ModelOrder {
cap_p: order.cap_p + 1,
..order
});
}
if order.cap_q < self.config.max_cap_q {
neighbors.push(ModelOrder {
cap_q: order.cap_q + 1,
..order
});
}
if order.cap_p > 0 && order.cap_q > 0 {
neighbors.push(ModelOrder {
cap_p: order.cap_p - 1,
cap_q: order.cap_q - 1,
..order
});
}
if order.cap_p > 0 && order.cap_q < self.config.max_cap_q {
neighbors.push(ModelOrder {
cap_p: order.cap_p - 1,
cap_q: order.cap_q + 1,
..order
});
}
if order.cap_p < self.config.max_cap_p && order.cap_q > 0 {
neighbors.push(ModelOrder {
cap_p: order.cap_p + 1,
cap_q: order.cap_q - 1,
..order
});
}
if order.cap_p < self.config.max_cap_p && order.cap_q < self.config.max_cap_q {
neighbors.push(ModelOrder {
cap_p: order.cap_p + 1,
cap_q: order.cap_q + 1,
..order
});
}
}
if order.p > 0 {
neighbors.push(ModelOrder {
p: order.p - 1,
..order
});
}
if order.q > 0 {
neighbors.push(ModelOrder {
q: order.q - 1,
..order
});
}
if order.p < self.config.max_p {
neighbors.push(ModelOrder {
p: order.p + 1,
..order
});
}
if order.q < self.config.max_q {
neighbors.push(ModelOrder {
q: order.q + 1,
..order
});
}
if order.p > 0 && order.q > 0 {
neighbors.push(ModelOrder {
p: order.p - 1,
q: order.q - 1,
..order
});
}
if order.p > 0 && order.q < self.config.max_q {
neighbors.push(ModelOrder {
p: order.p - 1,
q: order.q + 1,
..order
});
}
if order.p < self.config.max_p && order.q > 0 {
neighbors.push(ModelOrder {
p: order.p + 1,
q: order.q - 1,
..order
});
}
if order.p < self.config.max_p && order.q < self.config.max_q {
neighbors.push(ModelOrder {
p: order.p + 1,
q: order.q + 1,
..order
});
}
neighbors
}
fn true_stepwise_search(
&mut self,
diff_series: &[f64],
d: usize,
cap_d: usize,
) -> Option<(ModelOrder, f64)> {
let s = self.config.seasonal_period;
let use_aic = self.config.use_aic;
let max_models = 94;
let initial_orders = vec![
ModelOrder {
p: 2,
d,
q: 2,
cap_p: if s > 1 { 1 } else { 0 },
cap_d,
cap_q: if s > 1 { 1 } else { 0 },
s,
},
ModelOrder {
p: 0,
d,
q: 0,
cap_p: 0,
cap_d,
cap_q: 0,
s,
},
ModelOrder {
p: if self.config.max_p > 0 { 1 } else { 0 },
d,
q: 0,
cap_p: if s > 1 && self.config.max_cap_p > 0 {
1
} else {
0
},
cap_d,
cap_q: 0,
s,
},
ModelOrder {
p: 0,
d,
q: if self.config.max_q > 0 { 1 } else { 0 },
cap_p: 0,
cap_d,
cap_q: if s > 1 && self.config.max_cap_q > 0 {
1
} else {
0
},
s,
},
];
let mut best_order: Option<ModelOrder> = None;
let mut best_score = f64::INFINITY;
let mut n_models = 0usize;
let mut visited = std::collections::HashSet::new();
for order in initial_orders {
let key = (order.p, order.q, order.cap_p, order.cap_q);
if visited.contains(&key) {
continue;
}
visited.insert(key);
n_models += 1;
if let Some(score) = Self::score_order_static(order, diff_series, use_aic) {
self.model_scores.push((order, score));
if score < best_score {
best_score = score;
best_order = Some(order);
}
}
}
let mut current_order = best_order?;
let mut current_score = best_score;
loop {
if n_models >= max_models {
break;
}
let neighbors = self.get_neighbors(current_order);
let mut improved = false;
for neighbor in neighbors {
let key = (neighbor.p, neighbor.q, neighbor.cap_p, neighbor.cap_q);
if visited.contains(&key) {
continue;
}
visited.insert(key);
n_models += 1;
if let Some(score) = Self::score_order_static(neighbor, diff_series, use_aic) {
self.model_scores.push((neighbor, score));
if score < current_score {
current_score = score;
current_order = neighbor;
improved = true;
break; }
}
if n_models >= max_models {
break;
}
}
if !improved {
break;
}
}
Some((current_order, current_score))
}
fn evaluate_candidates_fast(
&self,
series: &TimeSeries,
#[cfg_attr(not(feature = "parallel"), allow(unused_variables))]
diff_series_map: &std::collections::HashMap<(usize, usize), Vec<f64>>,
candidates: &[ModelOrder],
n_values: usize,
) -> (
Vec<(ModelOrder, f64)>,
Option<(SelectedModel, ModelOrder, f64)>,
) {
let use_aic = self.config.use_aic;
let valid_candidates: Vec<_> = candidates
.iter()
.filter(|order| {
let min_len = order.d
+ order.cap_d * order.s
+ order
.p
.max(order.q)
.max(order.cap_p.max(order.cap_q) * order.s.max(1))
+ 5;
n_values >= min_len
})
.copied()
.collect();
let mut sorted_candidates = valid_candidates;
sorted_candidates.sort_by_key(|o| {
let total_params = o.p + o.q + o.cap_p + o.cap_q;
let is_seasonal = if o.cap_p > 0 || o.cap_q > 0 { 1 } else { 0 };
(is_seasonal, total_params)
});
#[cfg(feature = "parallel")]
{
let scores: Vec<(ModelOrder, f64)> = sorted_candidates
.par_iter()
.filter_map(|&order| {
let diff_series = diff_series_map.get(&(order.d, order.cap_d))?;
Self::score_order_static(order, diff_series, use_aic)
.map(|score| (order, score))
})
.collect();
let best = scores
.iter()
.min_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.and_then(|&(order, score)| {
Self::evaluate_model_static(series, order, use_aic)
.map(|(model, _)| (model, order, score))
});
(scores, best)
}
#[cfg(not(feature = "parallel"))]
{
let mut scores = Vec::with_capacity(sorted_candidates.len());
let mut best: Option<(SelectedModel, ModelOrder, f64)> = None;
let mut best_score = f64::INFINITY;
for &order in &sorted_candidates {
if let Some(diff_series) = diff_series_map.get(&(order.d, order.cap_d)) {
if let Some(quick_score) = Self::score_order_static(order, diff_series, use_aic)
{
scores.push((order, quick_score));
if quick_score < best_score {
if let Some((model, _)) =
Self::evaluate_model_static(series, order, use_aic)
{
best_score = quick_score;
best = Some((model, order, quick_score));
}
}
}
}
}
(scores, best)
}
}
}
impl Default for AutoARIMA {
fn default() -> Self {
Self::new()
}
}
impl Forecaster for AutoARIMA {
fn fit(&mut self, series: &TimeSeries) -> Result<()> {
validate_series_complete(series)?;
let values = series.primary_values();
let s = self.config.seasonal_period;
let min_required = if s > 1 {
3 * s } else {
10
};
if values.len() < min_required {
return Err(ForecastError::InsufficientData {
needed: min_required,
got: values.len(),
hint: Some(if s > 1 {
format!(
"AutoARIMA needs at least 3 seasonal cycles (3*{}={})",
s, min_required
)
} else {
"AutoARIMA needs at least 10 observations for model selection".into()
}),
});
}
let s = if s > 1 && !Self::has_seasonal_pattern(values, s) {
0 } else {
s
};
let suggested_d = suggest_differencing(values).min(self.config.max_d);
let suggested_cap_d = if s > 1 {
Self::suggest_seasonal_differencing(values, s).min(self.config.max_cap_d)
} else {
0
};
let d_range = vec![suggested_d];
let cap_d_range: Vec<usize> = if s > 1 {
vec![suggested_cap_d]
} else {
vec![0]
};
use crate::models::arima::diff::difference;
let mut diff_series_map = std::collections::HashMap::new();
for &d in &d_range {
for &cap_d in &cap_d_range {
let nonseasonal_diff = difference(values, d);
let diff_series = if s > 1 && cap_d > 0 {
SARIMA::seasonal_difference(&nonseasonal_diff, cap_d, s)
} else {
nonseasonal_diff
};
diff_series_map.insert((d, cap_d), diff_series);
}
}
self.model_scores.clear();
let mut best_order: Option<ModelOrder> = None;
let mut best_score = f64::INFINITY;
if self.config.stepwise && self.config.true_stepwise {
for &d in &d_range {
for &cap_d in &cap_d_range {
if let Some(diff_series) = diff_series_map.get(&(d, cap_d)) {
if let Some((order, score)) =
self.true_stepwise_search(diff_series, d, cap_d)
{
if score < best_score {
best_score = score;
best_order = Some(order);
}
}
}
}
}
} else {
let mut candidates = Vec::new();
for &d in &d_range {
for &cap_d in &cap_d_range {
let new_candidates = if self.config.stepwise {
self.stepwise_candidates(d, cap_d)
} else {
self.exhaustive_candidates(d, cap_d)
};
candidates.extend(new_candidates);
}
}
candidates.sort_by(|a, b| {
(a.p, a.d, a.q, a.cap_p, a.cap_d, a.cap_q)
.cmp(&(b.p, b.d, b.q, b.cap_p, b.cap_d, b.cap_q))
});
candidates.dedup();
let (results, best) =
self.evaluate_candidates_fast(series, &diff_series_map, &candidates, values.len());
for (order, score) in results {
self.model_scores.push((order, score));
}
if let Some((model, order, _score)) = best {
best_order = Some(order);
self.selected_model = Some(model);
}
}
self.model_scores
.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
if self.selected_model.is_none() {
if let Some(order) = best_order {
if let Some((model, _)) =
Self::evaluate_model_static(series, order, self.config.use_aic)
{
self.selected_model = Some(model);
}
}
}
self.selected_order = best_order;
if self.selected_model.is_none() {
return Err(ForecastError::ConvergenceFailure(
"No valid ARIMA/SARIMA model could be fitted".to_string(),
));
}
self.training_values_store = Some(values.to_vec());
let regs = series.all_regressors();
self.training_regressors_store = if regs.is_empty() {
None
} else {
Some(regs.clone())
};
Ok(())
}
fn predict(&self, horizon: usize) -> Result<Forecast> {
match self.selected_model.as_ref() {
Some(SelectedModel::ARIMA(model)) => model.predict(horizon),
Some(SelectedModel::SARIMA(model)) => model.predict(horizon),
None => Err(ForecastError::FitRequired { model: None }),
}
}
fn predict_with_intervals(&self, horizon: usize, level: f64) -> Result<Forecast> {
match self.selected_model.as_ref() {
Some(SelectedModel::ARIMA(model)) => model.predict_with_intervals(horizon, level),
Some(SelectedModel::SARIMA(model)) => model.predict_with_intervals(horizon, level),
None => Err(ForecastError::FitRequired { model: None }),
}
}
fn fitted_values(&self) -> Option<&[f64]> {
match self.selected_model.as_ref()? {
SelectedModel::ARIMA(model) => model.fitted_values(),
SelectedModel::SARIMA(model) => model.fitted_values(),
}
}
fn fitted_values_with_intervals(&self, level: f64) -> Option<Forecast> {
match self.selected_model.as_ref()? {
SelectedModel::ARIMA(model) => model.fitted_values_with_intervals(level),
SelectedModel::SARIMA(model) => model.fitted_values_with_intervals(level),
}
}
fn residuals(&self) -> Option<&[f64]> {
match self.selected_model.as_ref()? {
SelectedModel::ARIMA(model) => model.residuals(),
SelectedModel::SARIMA(model) => model.residuals(),
}
}
fn training_values(&self) -> Result<&[f64]> {
self.training_values_store
.as_deref()
.ok_or(ForecastError::FitRequired {
model: Some("AutoARIMA".into()),
})
}
fn training_regressors(&self) -> Option<&std::collections::HashMap<String, Vec<f64>>> {
self.training_regressors_store.as_ref()
}
fn trend_component(&self) -> Result<&[f64]> {
self.fitted_values().ok_or(ForecastError::FitRequired {
model: Some("AutoARIMA".into()),
})
}
fn name(&self) -> &str {
match &self.selected_model {
Some(SelectedModel::SARIMA(_)) => "AutoARIMA (SARIMA)",
_ => "AutoARIMA",
}
}
fn explanation(&self) -> Result<Explanation> {
<Self as Inspectable>::explanation(self)
}
fn supports_exog(&self) -> bool {
true
}
fn has_exog(&self) -> bool {
match self.selected_model.as_ref() {
Some(SelectedModel::ARIMA(model)) => model.has_exog(),
Some(SelectedModel::SARIMA(model)) => model.has_exog(),
None => false,
}
}
fn exog_names(&self) -> Option<&[String]> {
match self.selected_model.as_ref()? {
SelectedModel::ARIMA(model) => model.exog_names(),
SelectedModel::SARIMA(model) => model.exog_names(),
}
}
fn exog_coefficients(&self) -> Option<&OLSResult> {
match self.selected_model.as_ref()? {
SelectedModel::ARIMA(model) => model.exog_coefficients(),
SelectedModel::SARIMA(model) => model.exog_coefficients(),
}
}
fn predict_with_exog(
&self,
horizon: usize,
future_regressors: &std::collections::HashMap<String, Vec<f64>>,
) -> Result<Forecast> {
match self.selected_model.as_ref() {
Some(SelectedModel::ARIMA(model)) => {
model.predict_with_exog(horizon, future_regressors)
}
Some(SelectedModel::SARIMA(model)) => {
model.predict_with_exog(horizon, future_regressors)
}
None => Err(ForecastError::FitRequired { model: None }),
}
}
fn predict_with_exog_intervals(
&self,
horizon: usize,
future_regressors: &std::collections::HashMap<String, Vec<f64>>,
level: f64,
) -> Result<Forecast> {
match self.selected_model.as_ref() {
Some(SelectedModel::ARIMA(model)) => {
model.predict_with_exog_intervals(horizon, future_regressors, level)
}
Some(SelectedModel::SARIMA(model)) => {
model.predict_with_exog_intervals(horizon, future_regressors, level)
}
None => Err(ForecastError::FitRequired { model: None }),
}
}
}
impl Inspectable for AutoARIMA {
fn explanation(&self) -> Result<Explanation> {
let model = self
.selected_model
.as_ref()
.ok_or_else(|| ForecastError::FitRequired {
model: Some("AutoARIMA".to_string()),
})?;
let order = self
.selected_order
.ok_or_else(|| ForecastError::FitRequired {
model: Some("AutoARIMA".to_string()),
})?;
let (coefficients, aic, bic, fitted_values, residuals) = match model {
SelectedModel::ARIMA(m) => {
let mut coeffs = Vec::new();
coeffs.extend_from_slice(m.ar_coefficients());
coeffs.extend_from_slice(m.ma_coefficients());
let f = m.fitted_values().map(|v| v.to_vec()).unwrap_or_default();
let r = m.residuals().map(|v| v.to_vec()).unwrap_or_default();
(
coeffs,
m.aic().unwrap_or(f64::NAN),
m.bic().unwrap_or(f64::NAN),
f,
r,
)
}
SelectedModel::SARIMA(m) => {
let mut coeffs = Vec::new();
coeffs.extend_from_slice(m.ar_coefficients());
coeffs.extend_from_slice(m.ma_coefficients());
coeffs.extend_from_slice(m.seasonal_ar_coefficients());
coeffs.extend_from_slice(m.seasonal_ma_coefficients());
let f = m.fitted_values().map(|v| v.to_vec()).unwrap_or_default();
let r = m.residuals().map(|v| v.to_vec()).unwrap_or_default();
(
coeffs,
m.aic().unwrap_or(f64::NAN),
m.bic().unwrap_or(f64::NAN),
f,
r,
)
}
};
let seasonal_order = if order.is_seasonal() {
Some((order.cap_p, order.cap_d, order.cap_q, order.s))
} else {
None
};
Ok(Explanation::Arima(ArimaExplanation {
order: (order.p, order.d, order.q),
seasonal_order,
coefficients,
aic,
bic,
fitted_values,
residuals,
}))
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, TimeZone, Utc};
fn make_timestamps(n: usize) -> Vec<chrono::DateTime<Utc>> {
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 auto_arima_selects_model() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100).map(|i| 10.0 + (i as f64 * 0.2).sin()).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
assert!(model.selected_order().is_some());
assert!(!model.model_scores().is_empty());
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn auto_arima_with_trend() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 10.0 + 1.5 * i as f64 + (i as f64 * 0.2).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
assert!(model.selected_order().is_some());
}
#[test]
fn auto_arima_ar_process() {
let timestamps = make_timestamps(100);
let mut values = vec![10.0];
for i in 1..100 {
values.push(0.8 * values[i - 1] + (i as f64 * 0.05).sin());
}
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
let (p, _, _) = model.selected_order().unwrap();
assert!(p >= 1);
}
#[test]
fn auto_arima_exhaustive() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 10.0 + i as f64 * 0.5 + (i as f64 * 0.3).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = AutoARIMAConfig::default().exhaustive();
let mut model = AutoARIMA::with_config(config);
model.fit(&ts).unwrap();
assert!(model.selected_order().is_some());
assert!(model.model_scores().len() > 3);
}
#[test]
fn auto_arima_true_stepwise() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 10.0 + i as f64 * 0.5 + (i as f64 * 0.3).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = AutoARIMAConfig::default().with_true_stepwise();
let mut model = AutoARIMA::with_config(config);
model.fit(&ts).unwrap();
assert!(model.selected_order().is_some());
let n_models = model.model_scores().len();
assert!(
n_models > 0,
"Should evaluate at least some models, got {}",
n_models
);
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
#[test]
fn auto_arima_model_scores_sorted() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100).map(|i| 10.0 + (i as f64 * 0.3).sin()).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
let scores = model.model_scores();
for i in 1..scores.len() {
assert!(scores[i].1 >= scores[i - 1].1);
}
}
#[test]
fn auto_arima_confidence_intervals() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 10.0 + i as f64 * 0.5 + (i as f64 * 0.3).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
let forecast = model.predict_with_intervals(5, 0.95).unwrap();
assert!(forecast.has_lower());
assert!(forecast.has_upper());
}
#[test]
fn auto_arima_insufficient_data() {
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 mut model = AutoARIMA::new();
assert!(matches!(
model.fit(&ts),
Err(ForecastError::InsufficientData { .. })
));
}
#[test]
fn auto_arima_requires_fit() {
let model = AutoARIMA::new();
assert!(matches!(
model.predict(5),
Err(ForecastError::FitRequired { .. })
));
}
#[test]
fn auto_arima_fitted_and_residuals() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 10.0 + i as f64 + (i as f64 * 0.2).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
assert!(model.fitted_values().is_some());
assert!(model.residuals().is_some());
}
#[test]
fn auto_arima_name() {
let model = AutoARIMA::new();
assert_eq!(model.name(), "AutoARIMA");
}
#[test]
fn auto_arima_config() {
let config = AutoARIMAConfig::default()
.with_max_orders(5, 2, 5)
.exhaustive();
assert_eq!(config.max_p, 5);
assert_eq!(config.max_d, 2);
assert_eq!(config.max_q, 5);
assert!(!config.stepwise);
}
#[test]
fn auto_arima_seasonal() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| {
50.0 + 0.5 * i as f64 + 10.0 * (2.0 * std::f64::consts::PI * i as f64 / 12.0).sin()
})
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::seasonal(12);
model.fit(&ts).unwrap();
assert!(model.selected_full_order().is_some());
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn auto_arima_seasonal_config() {
let config = AutoARIMAConfig::default()
.with_seasonal_period(12)
.with_seasonal_orders(2, 1, 2);
assert_eq!(config.seasonal_period, 12);
assert_eq!(config.max_cap_p, 2);
assert_eq!(config.max_cap_d, 1);
assert_eq!(config.max_cap_q, 2);
}
#[test]
fn auto_arima_seasonal_selects_sarima() {
let timestamps = make_timestamps(100);
let values: Vec<f64> = (0..100)
.map(|i| 50.0 + 15.0 * (2.0 * std::f64::consts::PI * i as f64 / 12.0).sin())
.collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let config = AutoARIMAConfig::default()
.with_seasonal_period(12)
.exhaustive();
let mut model = AutoARIMA::with_config(config);
model.fit(&ts).unwrap();
if model.selected_full_order().is_some() {
assert!(model.model_scores().len() > 1);
}
let forecast = model.predict(12).unwrap();
assert_eq!(forecast.horizon(), 12);
}
#[test]
fn auto_arima_training_values_retained() {
let timestamps = make_timestamps(60);
let values: Vec<f64> = (0..60).map(|i| 10.0 + (i as f64 * 0.3).sin()).collect();
let ts = TimeSeries::univariate(timestamps, values.clone()).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
let training = model.training_values().unwrap();
assert_eq!(training, values.as_slice());
}
#[test]
fn auto_arima_training_regressors_none_without_regs() {
let timestamps = make_timestamps(60);
let values: Vec<f64> = (0..60).map(|i| 10.0 + (i as f64 * 0.3).sin()).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
assert!(model.training_regressors().is_none());
}
#[test]
fn auto_arima_trend_equals_fitted_values() {
let timestamps = make_timestamps(80);
let values: Vec<f64> = (0..80).map(|i| 10.0 + (i as f64 * 0.3).sin()).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
let trend = model.trend_component().unwrap();
let fitted = model.fitted_values().unwrap();
assert_eq!(trend.len(), fitted.len());
for (t, f) in trend.iter().zip(fitted.iter()) {
assert_eq!(t.is_nan(), f.is_nan(), "NaN-ness must agree");
if !t.is_nan() {
assert_eq!(t, f);
}
}
}
#[test]
fn auto_arima_seasonal_component_returns_err() {
let timestamps = make_timestamps(60);
let values: Vec<f64> = (0..60).map(|i| 10.0 + (i as f64 * 0.3).sin()).collect();
let ts = TimeSeries::univariate(timestamps, values).unwrap();
let mut model = AutoARIMA::new();
model.fit(&ts).unwrap();
assert!(matches!(
model.seasonal_component(),
Err(ForecastError::InvalidParameter(_))
));
}
#[test]
fn auto_arima_trend_component_requires_fit() {
let model = AutoARIMA::new();
assert!(matches!(
model.trend_component(),
Err(ForecastError::FitRequired { .. })
));
}
}