Skip to main content

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}