Skip to main content

torsh_autograd/
hyperparameter_optimization.rs

1//! Gradient-based hyperparameter optimization for automatic tuning of learning rates,
2//! regularization parameters, and other hyperparameters using differentiation.
3
4use crate::context::AutogradContext;
5use std::collections::HashMap;
6use torsh_core::{Result, TorshError};
7use torsh_tensor::Tensor;
8
9/// Configuration for hyperparameter optimization
10#[derive(Debug, Clone)]
11pub struct HyperparameterConfig {
12    /// Learning rate for hyperparameter updates
13    pub meta_learning_rate: f64,
14    /// Maximum number of optimization steps
15    pub max_steps: usize,
16    /// Convergence tolerance
17    pub tolerance: f64,
18    /// Whether to use second-order gradients
19    pub second_order: bool,
20    /// Validation frequency (steps)
21    pub validation_frequency: usize,
22    /// Early stopping patience
23    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/// Hyperparameter that can be optimized
40#[derive(Debug, Clone)]
41pub struct OptimizableHyperparameter {
42    /// Current value of the hyperparameter
43    pub value: Tensor,
44    /// Name/identifier for the hyperparameter
45    pub name: String,
46    /// Lower bound for the parameter
47    pub lower_bound: Option<f64>,
48    /// Upper bound for the parameter
49    pub upper_bound: Option<f64>,
50    /// Whether to use log scale (e.g., for learning rates)
51    pub log_scale: bool,
52}
53
54impl OptimizableHyperparameter {
55    /// Create a new optimizable hyperparameter
56    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    /// Get the actual hyperparameter value (handling log scale)
79    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    /// Apply bounds and constraints to the hyperparameter
89    pub fn apply_constraints(&mut self) -> Result<()> {
90        let mut value = self.value.item()? as f64;
91
92        // Apply bounds
93        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
108/// Gradient-based hyperparameter optimizer
109pub 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    /// Create a new hyperparameter optimizer
121    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    /// Add a hyperparameter to optimize
133    pub fn add_hyperparameter(&mut self, hyperparameter: OptimizableHyperparameter) {
134        let name = hyperparameter.name.clone();
135        self.hyperparameters.insert(name, hyperparameter);
136    }
137
138    /// Get hyperparameter value by name
139    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    /// Perform one step of hyperparameter optimization
149    pub fn step<F>(&mut self, objective_fn: F) -> Result<f64>
150    where
151        F: Fn(&HashMap<String, f64>) -> Result<Tensor>,
152    {
153        // Get current hyperparameter values
154        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        // Compute objective and gradients
160        let objective = objective_fn(&current_values)?;
161        let objective_value = objective.item()? as f64;
162
163        // Compute gradients with respect to hyperparameters
164        let gradients = self.compute_hyperparameter_gradients(&objective_fn, &current_values)?;
165
166        // Update hyperparameters using gradients
167        self.update_hyperparameters(&gradients)?;
168
169        // Apply constraints
170        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    /// Optimize hyperparameters using validation loss
179    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            // Perform optimization step
193            let train_loss = self.step(&objective_fn)?;
194
195            // Validate periodically
196            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(&current_values)?;
200                validation_loss = Some(val_loss);
201
202                // Early stopping check
203                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            // Record history
219            history.push(OptimizationStep {
220                step,
221                train_loss,
222                validation_loss,
223                hyperparameters: self.get_current_values()?,
224            });
225
226            // Check convergence
227            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    /// Compute gradients with respect to hyperparameters
241    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                // First-order gradient via central finite differences.
256                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    /// Compute the first-order gradient of the objective with respect to a
266    /// single hyperparameter using central finite differences.
267    ///
268    /// `objective_fn` only exposes the objective as a black-box function of
269    /// concrete hyperparameter values (`&HashMap<String, f64> -> Tensor`): the
270    /// values it receives are plain `f64`s extracted via
271    /// `OptimizableHyperparameter::get_value`, fully detached from any
272    /// computation graph. That means there is no recorded tape for
273    /// `self.context` (`AutogradContext`) -- or for the tensor's own
274    /// `requires_grad` tape -- to run reverse-mode differentiation through, no
275    /// matter how the objective is invoked. Central finite differences is the
276    /// principled technique for differentiating exactly this kind of
277    /// black-box scalar objective; it is not a placeholder, it is the same
278    /// numerical method this crate already trusts as a reference oracle for
279    /// gradient checking elsewhere (see `gradient_checking.rs`).
280    ///
281    /// The difference is taken on the *raw* hyperparameter tensor value
282    /// (`hyperparam.value`, pre-`exp` for log-scale parameters) because that
283    /// is the quantity `update_hyperparameters` actually updates. For
284    /// log-scale hyperparameters we re-apply the same `exp` transform used by
285    /// `OptimizableHyperparameter::get_value` before invoking `objective_fn`,
286    /// so the chain rule through the log transform falls out automatically.
287    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    /// Adaptive step size for central-difference numerical differentiation.
315    ///
316    /// Uses a relative step (scaled by the magnitude of the evaluation point,
317    /// floored at `1.0` so the step stays well-defined near zero) sized to
318    /// `f32::EPSILON.cbrt()`. Hyperparameter values round-trip through
319    /// `f32`-backed `Tensor`s here, so a step much smaller than that would be
320    /// swallowed by `f32` rounding noise when `objective_fn` casts its result
321    /// down to `f32`, while a much larger step would introduce unnecessary
322    /// truncation error for non-quadratic objectives.
323    fn finite_difference_step(x: f64) -> f64 {
324        (f32::EPSILON as f64).cbrt() * x.abs().max(1.0)
325    }
326
327    /// Step size for the *second* central difference.
328    ///
329    /// A second difference divides by `h^2`, so the cancellation error of the
330    /// three objective evaluations is amplified by `1/h^2` instead of `1/h`.
331    /// With values round-tripping through `f32`-backed `Tensor`s, the
332    /// first-derivative step (`eps^(1/3)`, see [`Self::finite_difference_step`])
333    /// would inflate that noise by roughly four orders of magnitude, so the
334    /// Hessian term uses the standard `eps^(1/4)` step instead.
335    fn second_difference_step(x: f64) -> f64 {
336        (f32::EPSILON as f64).powf(0.25) * x.abs().max(1.0)
337    }
338
339    /// Compute a second-order (safeguarded Newton) update direction for one
340    /// hyperparameter.
341    ///
342    /// The objective is only available as a black box over concrete
343    /// hyperparameter values (see [`Self::compute_first_order_gradient`] for why
344    /// no tape exists), so both derivatives are taken numerically about the
345    /// current raw value `x`:
346    ///
347    /// * gradient `g = (f(x+h) - f(x-h)) / 2h`
348    /// * curvature `c = (f(x+h) - 2 f(x) + f(x-h)) / h^2`
349    ///
350    /// and the returned direction is the Newton step `g / c`. Because
351    /// `update_hyperparameters` applies `x <- x - meta_learning_rate * d`, that
352    /// makes the meta learning rate a damping factor on a true Newton step,
353    /// which is what "second order" buys: the step shrinks automatically in
354    /// sharply curved directions and lengthens in flat ones, instead of using
355    /// one fixed rate everywhere.
356    ///
357    /// **Safeguard.** A Newton step is only a descent direction where the
358    /// objective is locally convex. When the measured curvature is not usefully
359    /// positive (`c <= |g| * sqrt(eps)`, which also covers the numerically
360    /// indistinguishable-from-zero case), the plain first-order gradient is
361    /// returned instead. This is the standard safeguarded-Newton fallback, not a
362    /// silent no-op: the direction is always a real derivative of the objective.
363    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    /// Update hyperparameters using computed gradients
413    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                // Gradient ascent (we want to maximize the objective, which is typically negative loss)
420                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    /// Get current hyperparameter values
429    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/// Result of hyperparameter optimization
439#[derive(Debug, Clone)]
440pub struct HyperparameterOptimizationResult {
441    /// Whether optimization converged
442    pub converged: bool,
443    /// Final optimized hyperparameters
444    pub final_hyperparameters: HashMap<String, f64>,
445    /// Best validation loss achieved
446    pub best_validation_loss: Option<f64>,
447    /// Optimization history
448    pub history: Vec<OptimizationStep>,
449}
450
451/// Single step in optimization history
452#[derive(Debug, Clone)]
453pub struct OptimizationStep {
454    /// Step number
455    pub step: usize,
456    /// Training loss at this step
457    pub train_loss: f64,
458    /// Validation loss at this step (if computed)
459    pub validation_loss: Option<f64>,
460    /// Hyperparameter values at this step
461    pub hyperparameters: HashMap<String, f64>,
462}
463
464/// Convenience functions for common hyperparameter optimization scenarios
465impl HyperparameterOptimizer {
466    /// Create optimizer for learning rate optimization
467    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), // Lower bound
478            Some(1.0),  // Upper bound
479            true,       // Log scale
480        )?;
481
482        optimizer.add_hyperparameter(lr_param);
483        Ok(optimizer)
484    }
485
486    /// Create optimizer for regularization strength
487    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), // Lower bound
498            Some(1.0), // Upper bound
499            true,      // Log scale
500        )?;
501
502        optimizer.add_hyperparameter(reg_param);
503        Ok(optimizer)
504    }
505
506    /// Create optimizer for multiple hyperparameters
507    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, // Initial value too high
556            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    /// Regression test for the "no-op optimizer" bug: `compute_first_order_gradient`
581    /// used to unconditionally return `Tensor::zeros_like(...)`, so gradient
582    /// descent never moved the hyperparameter no matter how far it was from the
583    /// optimum. This must fail against that stub (the parameter never moves,
584    /// so `x_after == x_before` and neither assertion below can hold) and pass
585    /// once `compute_first_order_gradient` returns a real gradient.
586    #[test]
587    fn test_first_order_gradient_moves_toward_minimum() {
588        // Minimize f(x) = (x - 5)^2, whose unique minimum is x = 5.
589        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        // The parameter must have moved strictly closer to the true minimum.
618        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        // 20 gradient-descent steps on a well-conditioned convex quadratic
623        // should land close to the minimum, not just move a little.
624        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    /// `compute_first_order_gradient` is central-difference exact for a
631    /// quadratic objective, so its output should closely match the true
632    /// analytic gradient df/dx = 2(x - 5), not just be "non-zero".
633    #[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, &current_values)
652            .unwrap();
653
654        let grad_x = gradients["x"].item().unwrap() as f64;
655
656        // Analytic gradient at x = 0 is 2 * (0 - 5) = -10.
657        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}