1use crate::context::AutogradContext;
5use std::collections::HashMap;
6use torsh_core::{Result, TorshError};
7use torsh_tensor::Tensor;
8
9#[derive(Debug, Clone)]
11pub struct HyperparameterConfig {
12 pub meta_learning_rate: f64,
14 pub max_steps: usize,
16 pub tolerance: f64,
18 pub second_order: bool,
20 pub validation_frequency: usize,
22 pub early_stopping_patience: usize,
24}
25
26impl Default for HyperparameterConfig {
27 fn default() -> Self {
28 Self {
29 meta_learning_rate: 0.01,
30 max_steps: 1000,
31 tolerance: 1e-6,
32 second_order: true,
33 validation_frequency: 10,
34 early_stopping_patience: 50,
35 }
36 }
37}
38
39#[derive(Debug, Clone)]
41pub struct OptimizableHyperparameter {
42 pub value: Tensor,
44 pub name: String,
46 pub lower_bound: Option<f64>,
48 pub upper_bound: Option<f64>,
50 pub log_scale: bool,
52}
53
54impl OptimizableHyperparameter {
55 pub fn new(
57 name: String,
58 initial_value: f64,
59 lower_bound: Option<f64>,
60 upper_bound: Option<f64>,
61 log_scale: bool,
62 ) -> Result<Self> {
63 let value = if log_scale {
64 Tensor::scalar(initial_value.ln() as f32)
65 } else {
66 Tensor::scalar(initial_value as f32)
67 }?;
68
69 Ok(Self {
70 value,
71 name,
72 lower_bound,
73 upper_bound,
74 log_scale,
75 })
76 }
77
78 pub fn get_value(&self) -> Result<f64> {
80 let raw_value = self.value.item()? as f64;
81 if self.log_scale {
82 Ok(raw_value.exp())
83 } else {
84 Ok(raw_value)
85 }
86 }
87
88 pub fn apply_constraints(&mut self) -> Result<()> {
90 let mut value = self.value.item()? as f64;
91
92 if let Some(lower) = self.lower_bound {
94 let bound = if self.log_scale { lower.ln() } else { lower };
95 value = value.max(bound);
96 }
97
98 if let Some(upper) = self.upper_bound {
99 let bound = if self.log_scale { upper.ln() } else { upper };
100 value = value.min(bound);
101 }
102
103 self.value = Tensor::scalar(value as f32)?;
104 Ok(())
105 }
106}
107
108pub struct HyperparameterOptimizer {
110 config: HyperparameterConfig,
111 hyperparameters: HashMap<String, OptimizableHyperparameter>,
112 #[allow(dead_code)]
113 context: AutogradContext,
114 step_count: usize,
115 best_validation_loss: Option<f64>,
116 patience_counter: usize,
117}
118
119impl HyperparameterOptimizer {
120 pub fn new(config: HyperparameterConfig) -> Self {
122 Self {
123 config,
124 hyperparameters: HashMap::new(),
125 context: AutogradContext::new(),
126 step_count: 0,
127 best_validation_loss: None,
128 patience_counter: 0,
129 }
130 }
131
132 pub fn add_hyperparameter(&mut self, hyperparameter: OptimizableHyperparameter) {
134 let name = hyperparameter.name.clone();
135 self.hyperparameters.insert(name, hyperparameter);
136 }
137
138 pub fn get_hyperparameter(&self, name: &str) -> Result<f64> {
140 self.hyperparameters
141 .get(name)
142 .ok_or_else(|| {
143 TorshError::AutogradError(format!("Hyperparameter '{}' not found", name))
144 })?
145 .get_value()
146 }
147
148 pub fn step<F>(&mut self, objective_fn: F) -> Result<f64>
150 where
151 F: Fn(&HashMap<String, f64>) -> Result<Tensor>,
152 {
153 let mut current_values = HashMap::new();
155 for (name, hyperparam) in &self.hyperparameters {
156 current_values.insert(name.clone(), hyperparam.get_value()?);
157 }
158
159 let objective = objective_fn(¤t_values)?;
161 let objective_value = objective.item()? as f64;
162
163 let gradients = self.compute_hyperparameter_gradients(&objective_fn, ¤t_values)?;
165
166 self.update_hyperparameters(&gradients)?;
168
169 for hyperparam in self.hyperparameters.values_mut() {
171 hyperparam.apply_constraints()?;
172 }
173
174 self.step_count += 1;
175 Ok(objective_value)
176 }
177
178 pub fn optimize<F, V>(
180 &mut self,
181 objective_fn: F,
182 validation_fn: V,
183 ) -> Result<HyperparameterOptimizationResult>
184 where
185 F: Fn(&HashMap<String, f64>) -> Result<Tensor>,
186 V: Fn(&HashMap<String, f64>) -> Result<f64>,
187 {
188 let mut history = Vec::new();
189 let mut converged = false;
190
191 for step in 0..self.config.max_steps {
192 let train_loss = self.step(&objective_fn)?;
194
195 let mut validation_loss = None;
197 if step % self.config.validation_frequency == 0 {
198 let current_values = self.get_current_values()?;
199 let val_loss = validation_fn(¤t_values)?;
200 validation_loss = Some(val_loss);
201
202 if let Some(best_loss) = self.best_validation_loss {
204 if val_loss < best_loss - self.config.tolerance {
205 self.best_validation_loss = Some(val_loss);
206 self.patience_counter = 0;
207 } else {
208 self.patience_counter += 1;
209 if self.patience_counter >= self.config.early_stopping_patience {
210 converged = true;
211 }
212 }
213 } else {
214 self.best_validation_loss = Some(val_loss);
215 }
216 }
217
218 history.push(OptimizationStep {
220 step,
221 train_loss,
222 validation_loss,
223 hyperparameters: self.get_current_values()?,
224 });
225
226 if converged {
228 break;
229 }
230 }
231
232 Ok(HyperparameterOptimizationResult {
233 converged,
234 final_hyperparameters: self.get_current_values()?,
235 best_validation_loss: self.best_validation_loss,
236 history,
237 })
238 }
239
240 fn compute_hyperparameter_gradients<F>(
242 &self,
243 objective_fn: &F,
244 current_values: &HashMap<String, f64>,
245 ) -> Result<HashMap<String, Tensor>>
246 where
247 F: Fn(&HashMap<String, f64>) -> Result<Tensor>,
248 {
249 let mut gradients = HashMap::new();
250
251 for (name, hyperparam) in &self.hyperparameters {
252 let grad = if self.config.second_order {
253 self.compute_second_order_gradient(objective_fn, current_values, name, hyperparam)?
254 } else {
255 self.compute_first_order_gradient(objective_fn, current_values, name, hyperparam)?
257 };
258
259 gradients.insert(name.clone(), grad);
260 }
261
262 Ok(gradients)
263 }
264
265 fn compute_first_order_gradient<F>(
288 &self,
289 objective_fn: &F,
290 current_values: &HashMap<String, f64>,
291 name: &str,
292 hyperparam: &OptimizableHyperparameter,
293 ) -> Result<Tensor>
294 where
295 F: Fn(&HashMap<String, f64>) -> Result<Tensor>,
296 {
297 let raw_value = hyperparam.value.item()? as f64;
298 let step = Self::finite_difference_step(raw_value);
299
300 let evaluate_at = |raw: f64| -> Result<f64> {
301 let actual_value = if hyperparam.log_scale { raw.exp() } else { raw };
302 let mut perturbed = current_values.clone();
303 perturbed.insert(name.to_string(), actual_value);
304 Ok(objective_fn(&perturbed)?.item()? as f64)
305 };
306
307 let objective_plus = evaluate_at(raw_value + step)?;
308 let objective_minus = evaluate_at(raw_value - step)?;
309 let gradient = (objective_plus - objective_minus) / (2.0 * step);
310
311 Tensor::scalar(gradient as f32)
312 }
313
314 fn finite_difference_step(x: f64) -> f64 {
324 (f32::EPSILON as f64).cbrt() * x.abs().max(1.0)
325 }
326
327 fn second_difference_step(x: f64) -> f64 {
336 (f32::EPSILON as f64).powf(0.25) * x.abs().max(1.0)
337 }
338
339 fn compute_second_order_gradient<F>(
364 &self,
365 objective_fn: &F,
366 current_values: &HashMap<String, f64>,
367 name: &str,
368 hyperparam: &OptimizableHyperparameter,
369 ) -> Result<Tensor>
370 where
371 F: Fn(&HashMap<String, f64>) -> Result<Tensor>,
372 {
373 let raw_value = hyperparam.value.item()? as f64;
374 let step = Self::second_difference_step(raw_value);
375
376 let evaluate_at = |raw: f64| -> Result<f64> {
377 let actual_value = if hyperparam.log_scale { raw.exp() } else { raw };
378 let mut perturbed = current_values.clone();
379 perturbed.insert(name.to_string(), actual_value);
380 Ok(objective_fn(&perturbed)?.item()? as f64)
381 };
382
383 let objective_center = evaluate_at(raw_value)?;
384 let objective_plus = evaluate_at(raw_value + step)?;
385 let objective_minus = evaluate_at(raw_value - step)?;
386
387 let gradient = (objective_plus - objective_minus) / (2.0 * step);
388 let curvature = (objective_plus - 2.0 * objective_center + objective_minus) / (step * step);
389
390 let curvature_floor = gradient.abs() * (f32::EPSILON as f64).sqrt();
391 let direction = if curvature > curvature_floor && curvature.is_finite() {
392 gradient / curvature
393 } else {
394 tracing::debug!(
395 "Hyperparameter '{name}': curvature {curvature} is not usefully positive, \
396 falling back to the first-order direction"
397 );
398 gradient
399 };
400
401 if !direction.is_finite() {
402 return Err(TorshError::AutogradError(format!(
403 "second-order hyperparameter gradient for '{name}' is not finite \
404 (gradient {gradient}, curvature {curvature}); the objective is likely \
405 discontinuous at this point"
406 )));
407 }
408
409 Tensor::scalar(direction as f32)
410 }
411
412 fn update_hyperparameters(&mut self, gradients: &HashMap<String, Tensor>) -> Result<()> {
414 for (name, gradient) in gradients {
415 if let Some(hyperparam) = self.hyperparameters.get_mut(name) {
416 let grad_value = gradient.item()? as f64;
417 let current_value = hyperparam.value.item()? as f64;
418
419 let new_value =
421 current_value - (self.config.meta_learning_rate as f64 * grad_value);
422 hyperparam.value = Tensor::scalar(new_value as f32)?;
423 }
424 }
425 Ok(())
426 }
427
428 fn get_current_values(&self) -> Result<HashMap<String, f64>> {
430 let mut values = HashMap::new();
431 for (name, hyperparam) in &self.hyperparameters {
432 values.insert(name.clone(), hyperparam.get_value()?);
433 }
434 Ok(values)
435 }
436}
437
438#[derive(Debug, Clone)]
440pub struct HyperparameterOptimizationResult {
441 pub converged: bool,
443 pub final_hyperparameters: HashMap<String, f64>,
445 pub best_validation_loss: Option<f64>,
447 pub history: Vec<OptimizationStep>,
449}
450
451#[derive(Debug, Clone)]
453pub struct OptimizationStep {
454 pub step: usize,
456 pub train_loss: f64,
458 pub validation_loss: Option<f64>,
460 pub hyperparameters: HashMap<String, f64>,
462}
463
464impl HyperparameterOptimizer {
466 pub fn for_learning_rate(
468 initial_lr: f64,
469 config: Option<HyperparameterConfig>,
470 ) -> Result<Self> {
471 let config = config.unwrap_or_default();
472 let mut optimizer = Self::new(config);
473
474 let lr_param = OptimizableHyperparameter::new(
475 "learning_rate".to_string(),
476 initial_lr,
477 Some(1e-6), Some(1.0), true, )?;
481
482 optimizer.add_hyperparameter(lr_param);
483 Ok(optimizer)
484 }
485
486 pub fn for_regularization(
488 initial_reg: f64,
489 config: Option<HyperparameterConfig>,
490 ) -> Result<Self> {
491 let config = config.unwrap_or_default();
492 let mut optimizer = Self::new(config);
493
494 let reg_param = OptimizableHyperparameter::new(
495 "regularization".to_string(),
496 initial_reg,
497 Some(0.0), Some(1.0), true, )?;
501
502 optimizer.add_hyperparameter(reg_param);
503 Ok(optimizer)
504 }
505
506 pub fn for_multiple_params(
508 params: Vec<(&str, f64, Option<f64>, Option<f64>, bool)>,
509 config: Option<HyperparameterConfig>,
510 ) -> Result<Self> {
511 let config = config.unwrap_or_default();
512 let mut optimizer = Self::new(config);
513
514 for (name, initial_value, lower_bound, upper_bound, log_scale) in params {
515 let param = OptimizableHyperparameter::new(
516 name.to_string(),
517 initial_value,
518 lower_bound,
519 upper_bound,
520 log_scale,
521 )?;
522 optimizer.add_hyperparameter(param);
523 }
524
525 Ok(optimizer)
526 }
527}
528
529#[cfg(test)]
530mod tests {
531 use super::*;
532
533 #[test]
534 fn test_optimizable_hyperparameter_creation() {
535 let param = OptimizableHyperparameter::new(
536 "learning_rate".to_string(),
537 0.01,
538 Some(1e-6),
539 Some(1.0),
540 true,
541 )
542 .unwrap();
543
544 assert_eq!(param.name, "learning_rate");
545 assert!((param.get_value().unwrap() - 0.01).abs() < 1e-6);
546 assert_eq!(param.lower_bound, Some(1e-6));
547 assert_eq!(param.upper_bound, Some(1.0));
548 assert!(param.log_scale);
549 }
550
551 #[test]
552 fn test_hyperparameter_bounds() {
553 let mut param = OptimizableHyperparameter::new(
554 "test".to_string(),
555 10.0, Some(1e-6),
557 Some(1.0),
558 false,
559 )
560 .unwrap();
561
562 param.apply_constraints().unwrap();
563 assert!((param.get_value().unwrap() - 1.0).abs() < 1e-6);
564 }
565
566 #[test]
567 fn test_optimizer_creation() {
568 let config = HyperparameterConfig::default();
569 let optimizer = HyperparameterOptimizer::new(config);
570 assert_eq!(optimizer.step_count, 0);
571 }
572
573 #[test]
574 fn test_learning_rate_optimizer() {
575 let optimizer = HyperparameterOptimizer::for_learning_rate(0.01, None).unwrap();
576 let lr = optimizer.get_hyperparameter("learning_rate").unwrap();
577 assert!((lr - 0.01).abs() < 1e-6);
578 }
579
580 #[test]
587 fn test_first_order_gradient_moves_toward_minimum() {
588 fn objective(values: &HashMap<String, f64>) -> Result<Tensor> {
590 let x = values["x"];
591 Tensor::scalar(((x - 5.0) * (x - 5.0)) as f32)
592 }
593
594 let config = HyperparameterConfig {
595 meta_learning_rate: 0.1,
596 max_steps: 1,
597 tolerance: 1e-6,
598 second_order: false,
599 validation_frequency: 1,
600 early_stopping_patience: 50,
601 };
602
603 let mut optimizer = HyperparameterOptimizer::new(config);
604 optimizer.add_hyperparameter(
605 OptimizableHyperparameter::new("x".to_string(), 0.0, None, None, false).unwrap(),
606 );
607
608 let x_before = optimizer.get_hyperparameter("x").unwrap();
609 assert!((x_before - 0.0).abs() < 1e-6);
610
611 for _ in 0..20 {
612 optimizer.step(objective).unwrap();
613 }
614
615 let x_after = optimizer.get_hyperparameter("x").unwrap();
616
617 assert!(
619 (x_after - 5.0).abs() < (x_before - 5.0).abs(),
620 "expected x to move toward the minimum at 5.0 (from {x_before}), got {x_after}"
621 );
622 assert!(
625 (x_after - 5.0).abs() < 0.5,
626 "expected x to converge near the minimum 5.0 after 20 steps, got {x_after}"
627 );
628 }
629
630 #[test]
634 fn test_first_order_gradient_matches_analytic_gradient() {
635 fn objective(values: &HashMap<String, f64>) -> Result<Tensor> {
636 let x = values["x"];
637 Tensor::scalar(((x - 5.0) * (x - 5.0)) as f32)
638 }
639
640 let config = HyperparameterConfig {
641 second_order: false,
642 ..HyperparameterConfig::default()
643 };
644 let mut optimizer = HyperparameterOptimizer::new(config);
645 optimizer.add_hyperparameter(
646 OptimizableHyperparameter::new("x".to_string(), 0.0, None, None, false).unwrap(),
647 );
648
649 let current_values = optimizer.get_current_values().unwrap();
650 let gradients = optimizer
651 .compute_hyperparameter_gradients(&objective, ¤t_values)
652 .unwrap();
653
654 let grad_x = gradients["x"].item().unwrap() as f64;
655
656 assert!(
658 (grad_x - (-10.0)).abs() < 1e-2,
659 "expected gradient close to -10.0, got {grad_x}"
660 );
661 assert!(grad_x.abs() > 1e-6, "gradient must not be (near-)zero");
662 }
663}