use crate::error::Result;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub enum Explanation {
Regression(RegressionExplanation),
Ets(EtsExplanation),
Arima(ArimaExplanation),
Mfles(MflesExplanation),
Theta(ThetaExplanation),
Tbats(TbatsExplanation),
Mstl(MstlExplanation),
#[cfg(feature = "distributional")]
Laplace(LaplaceExplanation),
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct RegressionExplanation {
pub feature_names: Vec<String>,
pub coefficients: Vec<f64>,
pub intercept: f64,
pub r_squared: f64,
pub backend: String,
pub coef_std_errors: Vec<f64>,
pub intercept_std_error: f64,
pub fitted_values: Vec<f64>,
pub residuals: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct EtsExplanation {
pub spec: String,
pub alpha: f64,
pub beta: Option<f64>,
pub gamma: Option<f64>,
pub phi: Option<f64>,
pub seasonal_period: usize,
pub fitted_values: Vec<f64>,
pub trend_component: Vec<f64>,
pub seasonal_component: Option<Vec<f64>>,
pub residuals: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ArimaExplanation {
pub order: (usize, usize, usize),
pub seasonal_order: Option<(usize, usize, usize, usize)>,
pub coefficients: Vec<f64>,
pub aic: f64,
pub bic: f64,
pub fitted_values: Vec<f64>,
pub residuals: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct MflesExplanation {
pub seasonal_period: usize,
pub max_rounds: usize,
pub multiplicative: bool,
pub penalty: Option<f64>,
pub fitted_values: Vec<f64>,
pub trend_component: Vec<f64>,
pub seasonal_component: Vec<f64>,
pub residuals: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct ThetaExplanation {
pub variant: String,
pub theta: f64,
pub alpha: Option<f64>,
pub seasonal_period: usize,
pub fitted_values: Vec<f64>,
pub residuals: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct TbatsExplanation {
pub seasonal_periods: Vec<usize>,
pub box_cox_lambda: Option<f64>,
pub selected_config: String,
pub aic: f64,
pub fitted_values: Vec<f64>,
pub residuals: Vec<f64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct MstlExplanation {
pub seasonal_periods: Vec<usize>,
pub iterations: usize,
pub fitted_values: Vec<f64>,
pub trend_component: Vec<f64>,
pub seasonal_component: Vec<f64>,
pub residuals: Vec<f64>,
}
#[cfg(feature = "distributional")]
#[derive(Debug, Clone, PartialEq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct LaplaceExplanation {
pub horizon_dists: Vec<crate::models::laplace::GaussianMixture>,
pub leaf_weights: Vec<f64>,
pub leaf_names: Vec<String>,
pub fitted_values: Vec<f64>,
pub residuals: Vec<f64>,
}
pub trait Inspectable {
fn explanation(&self) -> Result<Explanation>;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn explanation_variants_are_owned_and_clonable() {
let e = Explanation::Regression(RegressionExplanation {
feature_names: vec!["x".into()],
coefficients: vec![1.0],
intercept: 0.5,
r_squared: 0.9,
backend: "ols".into(),
coef_std_errors: vec![0.05],
intercept_std_error: 0.02,
fitted_values: vec![1.0, 2.0],
residuals: vec![0.1, -0.1],
});
let e2 = e.clone();
assert_eq!(e, e2);
}
#[cfg(feature = "serde")]
#[test]
fn regression_explanation_round_trips_through_serde_json() {
let e = RegressionExplanation {
feature_names: vec!["intercept".into(), "lag_1".into()],
coefficients: vec![1.5, 0.8],
intercept: 1.5,
r_squared: 0.92,
backend: "ridge".into(),
coef_std_errors: vec![0.10, 0.04],
intercept_std_error: 0.07,
fitted_values: vec![1.0, 2.0, 3.0],
residuals: vec![0.0, 0.1, -0.1],
};
let json = serde_json::to_string(&e).unwrap();
let e2: RegressionExplanation = serde_json::from_str(&json).unwrap();
assert_eq!(e, e2);
}
#[cfg(feature = "serde")]
#[test]
fn explanation_enum_round_trips_through_serde_json() {
let e = Explanation::Ets(EtsExplanation {
spec: "AAA".into(),
alpha: 0.3,
beta: Some(0.1),
gamma: Some(0.2),
phi: None,
seasonal_period: 12,
fitted_values: vec![10.0; 5],
trend_component: vec![1.0; 5],
seasonal_component: Some(vec![0.5; 5]),
residuals: vec![0.0; 5],
});
let json = serde_json::to_string(&e).unwrap();
let e2: Explanation = serde_json::from_str(&json).unwrap();
assert_eq!(e, e2);
}
}