stats_claw/optimizers/gradient/gradient_descent.rs
1//! Vanilla (batch) gradient descent — the worked pattern every learning-rate
2//! optimizer follows.
3
4use crate::optimizers::{ConvergenceStatus, Objective, OptimizeResult, norm};
5
6/// Minimizes `obj` by full-gradient steps of fixed learning rate.
7///
8/// At each step the point moves against the gradient: `x ← x − lr·∇f(x)`. The
9/// run stops early (reporting [`ConvergenceStatus::Converged`]) when the gradient
10/// norm drops below `tol`, otherwise it runs the full `max_iter` budget and
11/// reports [`ConvergenceStatus::MaxIterReached`].
12///
13/// # Arguments
14///
15/// * `obj` — the objective to minimize.
16/// * `x0` — the starting point; its length is the problem dimension.
17/// * `lr` — the learning rate (step size); must be positive and small enough for
18/// the problem's curvature or the iterates diverge.
19/// * `max_iter` — the maximum number of gradient steps.
20/// * `tol` — the gradient-norm convergence threshold.
21///
22/// # Returns
23///
24/// An [`OptimizeResult`] with the located point, its objective value, the number
25/// of iterations performed, and the convergence status.
26///
27/// # Examples
28///
29/// ```
30/// use stats_claw::optimizers::gradient::gradient_descent;
31/// use stats_claw::optimizers::objectives::Quadratic;
32/// use stats_claw::optimizers::ConvergenceStatus;
33///
34/// let obj = Quadratic::new(vec![3.0, -2.0]);
35/// let r = gradient_descent(&obj, &[0.0, 0.0], 0.1, 10_000, 1e-12);
36/// assert!(matches!(r.status, ConvergenceStatus::Converged));
37/// assert!((r.x[0] - 3.0).abs() < 1e-6);
38/// ```
39#[must_use]
40pub fn gradient_descent(
41 obj: &impl Objective,
42 x0: &[f64],
43 lr: f64,
44 max_iter: usize,
45 tol: f64,
46) -> OptimizeResult {
47 let mut x = x0.to_vec();
48 let mut status = ConvergenceStatus::MaxIterReached;
49 let mut iterations = 0;
50 for step in 0..max_iter {
51 iterations = step + 1;
52 let g = obj.grad(&x);
53 if norm(&g) < tol {
54 status = ConvergenceStatus::Converged;
55 break;
56 }
57 for (xi, gi) in x.iter_mut().zip(&g) {
58 *xi -= lr * gi;
59 }
60 }
61 let fx = obj.value(&x);
62 OptimizeResult {
63 x,
64 fx,
65 iterations,
66 status,
67 }
68}
69
70#[cfg(test)]
71mod tests {
72 use super::*;
73 use crate::optimizers::objectives::Quadratic;
74
75 #[test]
76 fn budget_limited_run_reports_max_iter() {
77 let obj = Quadratic::new(vec![3.0, -2.0]);
78 let r = gradient_descent(&obj, &[0.0, 0.0], 0.1, 1, 1e-12);
79 assert_eq!(r.status, ConvergenceStatus::MaxIterReached);
80 assert_eq!(r.iterations, 1, "iterations was {}", r.iterations);
81 }
82}