Skip to main content

chronos_ts/
tuning.rs

1use chrono::NaiveDate;
2use ndarray::Array1;
3use serde::{Deserialize, Serialize};
4
5use crate::decomposition::{ProphetDecomposition, SeasonalityMode};
6use crate::errors::{ChronosError, Result};
7
8/// Hyperparameter grid target parameters
9#[derive(Debug, Clone, Serialize, Deserialize)]
10pub struct HyperparameterGrid {
11    pub changepoint_prior_scales: Vec<f64>,
12    pub seasonality_prior_scales: Vec<f64>,
13    pub holidays_prior_scales: Vec<f64>,
14    pub seasonality_modes: Vec<SeasonalityMode>,
15}
16
17impl Default for HyperparameterGrid {
18    fn default() -> Self {
19        Self {
20            changepoint_prior_scales: vec![0.001, 0.01, 0.05, 0.1, 0.5],
21            seasonality_prior_scales: vec![0.01, 0.1, 1.0, 10.0],
22            holidays_prior_scales: vec![0.01, 0.1, 1.0, 10.0],
23            seasonality_modes: vec![SeasonalityMode::Additive, SeasonalityMode::Multiplicative],
24        }
25    }
26}
27
28/// A specific candidate set of hyperparameter configurations
29#[derive(Debug, Clone, Serialize, Deserialize)]
30pub struct HyperparameterCandidate {
31    pub changepoint_prior_scale: f64,
32    pub seasonality_prior_scale: f64,
33    pub holidays_prior_scale: f64,
34    pub seasonality_mode: SeasonalityMode,
35}
36
37/// Cross-validation metric choice
38#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
39pub enum OptimizationMetric {
40    MAE,
41    RMSE,
42    MAPE,
43}
44
45/// Result returned from auto-tuning evaluation
46#[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct AutoTuneResult {
48    pub best_params: HyperparameterCandidate,
49    pub best_score: f64,
50    pub metric: OptimizationMetric,
51    pub all_evaluated_scores: Vec<(HyperparameterCandidate, f64)>,
52}
53
54/// AutoTuner handles parameter search over rolling time-series cutoffs
55pub struct AutoTuner {
56    grid: HyperparameterGrid,
57    metric: OptimizationMetric,
58    initial_window_pct: f64,
59    horizon_days: usize,
60    step_days: usize,
61}
62
63impl AutoTuner {
64    pub fn new(grid: HyperparameterGrid) -> Self {
65        Self {
66            grid,
67            metric: OptimizationMetric::MAE,
68            initial_window_pct: 0.6,
69            horizon_days: 30,
70            step_days: 15,
71        }
72    }
73
74    pub fn with_metric(mut self, metric: OptimizationMetric) -> Self {
75        self.metric = metric;
76        self
77    }
78
79    pub fn with_validation_split(
80        mut self,
81        initial_window_pct: f64,
82        horizon_days: usize,
83        step_days: usize,
84    ) -> Self {
85        self.initial_window_pct = initial_window_pct.clamp(0.1, 0.9);
86        self.horizon_days = horizon_days.max(1);
87        self.step_days = step_days.max(1);
88        self
89    }
90
91    /// Run grid search optimization across historical data cuts
92    pub fn fit_and_tune(
93        &self,
94        base_model: &ProphetDecomposition,
95        dates: &[NaiveDate],
96        values: &Array1<f64>,
97    ) -> Result<AutoTuneResult> {
98        let n = dates.len();
99        if n < 30 || values.len() != n {
100            return Err(ChronosError::InvalidParameters(
101                "Auto-tuning requires at least 30 data points".into(),
102            ));
103        }
104
105        let candidates = self.generate_candidates();
106        let cutoffs = self.generate_cross_validation_cutoffs(dates)?;
107
108        if cutoffs.is_empty() {
109            return Err(ChronosError::InvalidParameters(
110                "Dataset duration is too short for the configured horizon and initial window"
111                    .into(),
112            ));
113        }
114
115        let mut evaluated_scores = Vec::new();
116        let mut best_score = f64::INFINITY;
117        let mut best_candidate = candidates[0].clone();
118
119        for candidate in candidates {
120            let mut fold_scores = Vec::new();
121
122            for (train_end_idx, val_end_idx) in &cutoffs {
123                let train_dates = &dates[..*train_end_idx];
124                let train_vals = values.slice(ndarray::s![..*train_end_idx]).to_owned();
125
126                let val_dates = &dates[*train_end_idx..*val_end_idx];
127                let val_vals = values
128                    .slice(ndarray::s![*train_end_idx..*val_end_idx])
129                    .to_owned();
130
131                // Clone base template and apply hyperparameter candidate
132                let mut tuned_model = base_model.clone();
133                tuned_model.changepoint_prior_scale = candidate.changepoint_prior_scale;
134                tuned_model.holiday_prior_scale = candidate.holidays_prior_scale;
135                tuned_model.seasonality_mode = candidate.seasonality_mode;
136
137                // Update seasonality prior scales across registered seasonality specs
138                for spec in &mut tuned_model.seasonalities {
139                    spec.prior_scale = candidate.seasonality_prior_scale;
140                }
141
142                // Fit fold with 4 parameters: dates, y, cap, floor
143                if tuned_model
144                    .fit(train_dates, &train_vals, None, None)
145                    .is_err()
146                {
147                    continue;
148                }
149
150                // Predict validation step
151                if let Ok(pred) = tuned_model.predict(val_dates) {
152                    let score = evaluate_metric(&pred.yhat, &val_vals, self.metric);
153                    if !score.is_nan() && !score.is_infinite() {
154                        fold_scores.push(score);
155                    }
156                }
157            }
158
159            if fold_scores.is_empty() {
160                continue;
161            }
162
163            // Average score across cross-validation windows
164            let mean_score = fold_scores.iter().sum::<f64>() / fold_scores.len() as f64;
165            evaluated_scores.push((candidate.clone(), mean_score));
166
167            if mean_score < best_score {
168                best_score = mean_score;
169                best_candidate = candidate;
170            }
171        }
172
173        if evaluated_scores.is_empty() {
174            return Err(ChronosError::InvalidParameters(
175                "All candidate evaluations failed during tuning".into(),
176            ));
177        }
178
179        Ok(AutoTuneResult {
180            best_params: best_candidate,
181            best_score,
182            metric: self.metric,
183            all_evaluated_scores: evaluated_scores,
184        })
185    }
186
187    fn generate_candidates(&self) -> Vec<HyperparameterCandidate> {
188        let mut candidates = Vec::new();
189        for &cp in &self.grid.changepoint_prior_scales {
190            for &s in &self.grid.seasonality_prior_scales {
191                for &h in &self.grid.holidays_prior_scales {
192                    for &m in &self.grid.seasonality_modes {
193                        candidates.push(HyperparameterCandidate {
194                            changepoint_prior_scale: cp,
195                            seasonality_prior_scale: s,
196                            holidays_prior_scale: h,
197                            seasonality_mode: m,
198                        });
199                    }
200                }
201            }
202        }
203        candidates
204    }
205
206    fn generate_cross_validation_cutoffs(
207        &self,
208        dates: &[NaiveDate],
209    ) -> Result<Vec<(usize, usize)>> {
210        let n = dates.len();
211        let initial_train_size = ((n as f64) * self.initial_window_pct).floor() as usize;
212
213        let mut cutoffs = Vec::new();
214        let mut current_train_end = initial_train_size;
215
216        while current_train_end + self.horizon_days <= n {
217            let val_end = current_train_end + self.horizon_days;
218            cutoffs.push((current_train_end, val_end));
219            current_train_end += self.step_days;
220        }
221
222        Ok(cutoffs)
223    }
224}
225
226/// Helper function to evaluate accuracy error metrics
227fn evaluate_metric(pred: &Array1<f64>, actual: &Array1<f64>, metric: OptimizationMetric) -> f64 {
228    let n = pred.len();
229    if n == 0 || actual.len() != n {
230        return f64::NAN;
231    }
232
233    match metric {
234        OptimizationMetric::MAE => (pred - actual).mapv(|x| x.abs()).sum() / (n as f64),
235        OptimizationMetric::RMSE => ((pred - actual).mapv(|x| x * x).sum() / (n as f64)).sqrt(),
236        OptimizationMetric::MAPE => {
237            let sum_err: f64 = pred
238                .iter()
239                .zip(actual.iter())
240                .map(|(&p, &a)| {
241                    if a.abs() < 1e-8 {
242                        0.0
243                    } else {
244                        ((p - a) / a).abs()
245                    }
246                })
247                .sum();
248            (sum_err / (n as f64)) * 100.0
249        }
250    }
251}