1use chrono::NaiveDate;
2use ndarray::Array1;
3use serde::{Deserialize, Serialize};
4
5use crate::decomposition::{ProphetDecomposition, SeasonalityMode};
6use crate::errors::{ChronosError, Result};
7
8#[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#[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#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
39pub enum OptimizationMetric {
40 MAE,
41 RMSE,
42 MAPE,
43}
44
45#[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
54pub 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 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 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 for spec in &mut tuned_model.seasonalities {
139 spec.prior_scale = candidate.seasonality_prior_scale;
140 }
141
142 if tuned_model
144 .fit(train_dates, &train_vals, None, None)
145 .is_err()
146 {
147 continue;
148 }
149
150 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 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
226fn 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}