use crate::core::{Forecast, TimeSeries};
use crate::error::{ForecastError, Result};
use crate::models::{BoxedForecaster, FittedParams, Forecaster};
use crate::utils::ols::OLSResult;
use std::collections::HashMap;
use std::fmt;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum InverseMode {
Fitted,
Predict,
}
pub trait Transform: fmt::Debug + Send + Sync {
fn fit_transform(&mut self, values: &[f64]) -> Result<Vec<f64>>;
fn inverse(&self, values: &[f64], mode: InverseMode) -> Result<Vec<f64>>;
fn offset(&self) -> usize;
fn name(&self) -> &str;
fn clone_box(&self) -> Box<dyn Transform>;
}
impl Clone for Box<dyn Transform> {
fn clone(&self) -> Self {
self.clone_box()
}
}
pub struct Pipeline {
transforms: Vec<Box<dyn Transform>>,
model: BoxedForecaster,
original_len: usize,
total_offset: usize,
fitted: Option<Vec<f64>>,
residuals_cache: Option<Vec<f64>>,
}
impl fmt::Debug for Pipeline {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Pipeline")
.field("transforms", &self.transforms)
.field("model", &self.model.name())
.field("total_offset", &self.total_offset)
.finish()
}
}
impl Pipeline {
pub fn builder() -> PipelineBuilder {
PipelineBuilder {
transforms: Vec::new(),
model: None,
}
}
}
pub struct PipelineBuilder {
transforms: Vec<Box<dyn Transform>>,
model: Option<BoxedForecaster>,
}
impl PipelineBuilder {
pub fn transform(mut self, t: impl Transform + 'static) -> Self {
self.transforms.push(Box::new(t));
self
}
pub fn transform_boxed(mut self, t: Box<dyn Transform>) -> Self {
self.transforms.push(t);
self
}
pub fn model(mut self, model: BoxedForecaster) -> Self {
self.model = Some(model);
self
}
pub fn build(self) -> Pipeline {
Pipeline {
transforms: self.transforms,
model: self.model.expect("Pipeline requires a model"),
original_len: 0,
total_offset: 0,
fitted: None,
residuals_cache: None,
}
}
}
impl Forecaster for Pipeline {
fn fit(&mut self, series: &TimeSeries) -> Result<()> {
let original = series.primary_values();
let n = original.len();
if n == 0 {
return Err(ForecastError::EmptyData);
}
self.original_len = n;
self.total_offset = 0;
self.fitted = None;
self.residuals_cache = None;
let mut current = original.to_vec();
for t in &mut self.transforms {
current = t.fit_transform(¤t)?;
self.total_offset += t.offset();
}
let inner_ts = series.slice(self.total_offset, n)?;
let inner_ts = TimeSeries::univariate(inner_ts.timestamps().to_vec(), current.clone())?;
self.model.fit(&inner_ts)?;
if let Some(inner_fitted) = self.model.fitted_values() {
let mut inv = inner_fitted.to_vec();
for t in self.transforms.iter().rev() {
inv = t.inverse(&inv, InverseMode::Fitted)?;
}
let fitted_full = if inv.len() >= n {
inv[inv.len() - n..].to_vec()
} else {
let mut padded = vec![f64::NAN; n - inv.len()];
padded.extend_from_slice(&inv);
padded
};
let residuals: Vec<f64> = original
.iter()
.zip(fitted_full.iter())
.map(|(&o, &f)| if f.is_nan() { f64::NAN } else { o - f })
.collect();
self.fitted = Some(fitted_full);
self.residuals_cache = Some(residuals);
}
Ok(())
}
fn predict(&self, horizon: usize) -> Result<Forecast> {
let forecast = self.model.predict(horizon)?;
self.inverse_forecast(forecast)
}
fn predict_with_intervals(&self, horizon: usize, level: f64) -> Result<Forecast> {
let forecast = self.model.predict_with_intervals(horizon, level)?;
self.inverse_forecast(forecast)
}
fn fitted_values(&self) -> Option<&[f64]> {
self.fitted.as_deref()
}
fn residuals(&self) -> Option<&[f64]> {
self.residuals_cache.as_deref()
}
fn name(&self) -> &str {
"Pipeline"
}
fn fitted_params(&self) -> Option<FittedParams> {
self.model.fitted_params()
}
fn supports_exog(&self) -> bool {
self.model.supports_exog()
}
fn has_exog(&self) -> bool {
self.model.has_exog()
}
fn exog_names(&self) -> Option<&[String]> {
self.model.exog_names()
}
fn exog_coefficients(&self) -> Option<&OLSResult> {
self.model.exog_coefficients()
}
fn predict_with_exog(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
) -> Result<Forecast> {
let forecast = self.model.predict_with_exog(horizon, future_regressors)?;
self.inverse_forecast(forecast)
}
fn predict_with_exog_intervals(
&self,
horizon: usize,
future_regressors: &HashMap<String, Vec<f64>>,
level: f64,
) -> Result<Forecast> {
let forecast = self
.model
.predict_with_exog_intervals(horizon, future_regressors, level)?;
self.inverse_forecast(forecast)
}
}
impl Pipeline {
fn inverse_forecast(&self, forecast: Forecast) -> Result<Forecast> {
let point = self.inverse_series(forecast.primary())?;
if forecast.has_lower() && forecast.has_upper() {
let lower = self.inverse_series(forecast.lower_series(0)?)?;
let upper = self.inverse_series(forecast.upper_series(0)?)?;
Ok(Forecast::from_values_with_intervals(point, lower, upper))
} else {
Ok(Forecast::from_values(point))
}
}
fn inverse_series(&self, values: &[f64]) -> Result<Vec<f64>> {
let mut current = values.to_vec();
for t in self.transforms.iter().rev() {
current = t.inverse(¤t, InverseMode::Predict)?;
}
Ok(current)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::models::baseline::Naive;
use crate::transform::transforms::{
BoxCoxTransform, DifferenceTransform, LogTransform, ScaleMethod, ScaleTransform,
};
use chrono::{TimeZone, Utc};
fn make_timestamps(n: usize) -> Vec<chrono::DateTime<Utc>> {
(0..n)
.map(|i| {
Utc.with_ymd_and_hms(2020, 1, 1, 0, 0, 0).unwrap()
+ chrono::Duration::days(i as i64)
})
.collect()
}
fn make_ts(values: Vec<f64>) -> TimeSeries {
let n = values.len();
TimeSeries::univariate(make_timestamps(n), values).unwrap()
}
#[test]
fn pipeline_with_no_transforms() {
let values: Vec<f64> = (1..=30).map(|i| i as f64).collect();
let ts = make_ts(values);
let mut pipeline = Pipeline::builder().model(Box::new(Naive::new())).build();
pipeline.fit(&ts).unwrap();
assert!(pipeline.is_fitted());
let forecast = pipeline.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
for &v in forecast.primary() {
assert!((v - 30.0).abs() < 1e-10);
}
}
#[test]
fn pipeline_difference_naive() {
let values: Vec<f64> = (1..=30).map(|i| i as f64).collect();
let ts = make_ts(values);
let mut pipeline = Pipeline::builder()
.transform(DifferenceTransform::new(1))
.model(Box::new(Naive::new()))
.build();
pipeline.fit(&ts).unwrap();
let forecast = pipeline.predict(5).unwrap();
let expected: Vec<f64> = (31..=35).map(|i| i as f64).collect();
for (a, b) in forecast.primary().iter().zip(expected.iter()) {
assert!((a - b).abs() < 1e-10, "expected {}, got {}", b, a);
}
}
#[test]
fn pipeline_boxcox_difference_naive() {
let values: Vec<f64> = (1..=30).map(|i| (i as f64).powi(2)).collect();
let ts = make_ts(values.clone());
let mut pipeline = Pipeline::builder()
.transform(BoxCoxTransform::auto())
.transform(DifferenceTransform::new(1))
.model(Box::new(Naive::new()))
.build();
pipeline.fit(&ts).unwrap();
assert!(pipeline.is_fitted());
let forecast = pipeline.predict(3).unwrap();
assert_eq!(forecast.horizon(), 3);
for &v in forecast.primary() {
assert!(v > 800.0, "forecast {} should be > 800", v);
}
}
#[test]
fn pipeline_fitted_values_same_length_as_input() {
let values: Vec<f64> = (1..=20).map(|i| i as f64).collect();
let ts = make_ts(values.clone());
let mut pipeline = Pipeline::builder()
.transform(DifferenceTransform::new(1))
.model(Box::new(Naive::new()))
.build();
pipeline.fit(&ts).unwrap();
let fitted = pipeline.fitted_values().unwrap();
assert_eq!(fitted.len(), values.len());
let residuals = pipeline.residuals().unwrap();
assert_eq!(residuals.len(), values.len());
}
#[test]
fn pipeline_with_intervals() {
let values: Vec<f64> = (1..=30).map(|i| i as f64).collect();
let ts = make_ts(values);
let mut pipeline = Pipeline::builder()
.transform(DifferenceTransform::new(1))
.model(Box::new(Naive::new()))
.build();
pipeline.fit(&ts).unwrap();
let forecast = pipeline.predict_with_intervals(5, 0.95).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(forecast.has_lower());
assert!(forecast.has_upper());
}
#[test]
fn pipeline_name() {
let pipeline = Pipeline::builder().model(Box::new(Naive::new())).build();
assert_eq!(pipeline.name(), "Pipeline");
}
#[test]
fn pipeline_scale_naive() {
let values: Vec<f64> = (1..=20).map(|i| i as f64 * 10.0).collect();
let ts = make_ts(values.clone());
let mut pipeline = Pipeline::builder()
.transform(ScaleTransform::new(ScaleMethod::Standardize))
.model(Box::new(Naive::new()))
.build();
pipeline.fit(&ts).unwrap();
let forecast = pipeline.predict(3).unwrap();
for &v in forecast.primary() {
assert!((v - 200.0).abs() < 1e-8, "expected ~200, got {}", v);
}
}
#[test]
fn pipeline_log_naive() {
let values: Vec<f64> = (1..=20).map(|i| i as f64).collect();
let ts = make_ts(values);
let mut pipeline = Pipeline::builder()
.transform(LogTransform::new())
.model(Box::new(Naive::new()))
.build();
pipeline.fit(&ts).unwrap();
let forecast = pipeline.predict(3).unwrap();
for &v in forecast.primary() {
assert!((v - 20.0).abs() < 1e-8, "expected ~20, got {}", v);
}
}
#[test]
fn pipeline_fit_predict() {
let values: Vec<f64> = (1..=30).map(|i| i as f64).collect();
let ts = make_ts(values);
let mut pipeline = Pipeline::builder()
.transform(DifferenceTransform::new(1))
.model(Box::new(Naive::new()))
.build();
let forecast = pipeline.fit_predict(&ts, 5).unwrap();
assert_eq!(forecast.horizon(), 5);
assert!(pipeline.is_fitted());
}
#[test]
fn pipeline_predict_before_fit() {
let pipeline = Pipeline::builder()
.transform(DifferenceTransform::new(1))
.model(Box::new(Naive::new()))
.build();
assert!(!pipeline.is_fitted());
assert!(pipeline.predict(5).is_err());
}
#[test]
fn pipeline_insufficient_data_for_transform() {
let values = vec![1.0, 2.0];
let ts = make_ts(values);
let mut pipeline = Pipeline::builder()
.transform(DifferenceTransform::new(2))
.model(Box::new(Naive::new()))
.build();
assert!(pipeline.fit(&ts).is_err());
}
#[test]
fn pipeline_as_model_spec() {
use crate::models::ModelSpec;
let spec = ModelSpec::new(
"Pipeline(Diff→Naive)",
|| {
Box::new(
Pipeline::builder()
.transform(DifferenceTransform::new(1))
.model(Box::new(Naive::new()))
.build(),
)
},
true,
);
let mut model = spec.create();
let values: Vec<f64> = (1..=30).map(|i| i as f64).collect();
let ts = make_ts(values);
model.fit(&ts).unwrap();
let forecast = model.predict(5).unwrap();
assert_eq!(forecast.horizon(), 5);
}
}