Skip to main content

rustyqlib/core/optimization/
nelder_mead.rs

1//! Nelder-Mead simplex: gradient-free local search by reflecting,
2//! expanding and contracting a simplex of `n + 1` points. The tool for
3//! noisy or non-smooth objectives where gradients mislead.
4
5use super::{OptimConfig, OptimResult};
6
7/// Minimize `f` from `x0` with the standard Nelder-Mead coefficients
8/// (reflection 1, expansion 2, contraction 1/2, shrink 1/2). Converged
9/// when the simplex's value spread falls below `tol` (relative to the
10/// best value).
11pub fn nelder_mead(cfg: &OptimConfig, f: &dyn Fn(&[f64]) -> f64, x0: &[f64]) -> OptimResult {
12    let n = x0.len();
13    // initial simplex: x0 plus a 5% step per coordinate (0.00025 at zero)
14    let mut simplex: Vec<Vec<f64>> = vec![x0.to_vec()];
15    for i in 0..n {
16        let mut v = x0.to_vec();
17        v[i] += if x0[i] != 0.0 { 0.05 * x0[i] } else { 0.00025 };
18        simplex.push(v);
19    }
20    let mut values: Vec<f64> = simplex.iter().map(|v| f(v)).collect();
21
22    for it in 0..cfg.max_iter {
23        // order best -> worst
24        let mut order: Vec<usize> = (0..=n).collect();
25        order.sort_by(|&i, &j| values[i].total_cmp(&values[j]));
26        let best = order[0];
27        let worst = order[n];
28        let second_worst = order[n - 1];
29        if (values[worst] - values[best]).abs() <= cfg.tol * (1.0 + values[best].abs()) {
30            return OptimResult {
31                x: simplex[best].clone(),
32                value: values[best],
33                iterations: it,
34                converged: true,
35            };
36        }
37        // centroid of all but the worst
38        let mut centroid = vec![0.0; n];
39        for (idx, v) in simplex.iter().enumerate() {
40            if idx != worst {
41                for (c, vi) in centroid.iter_mut().zip(v) {
42                    *c += vi / n as f64;
43                }
44            }
45        }
46        let point = |t: f64| -> Vec<f64> {
47            centroid
48                .iter()
49                .zip(&simplex[worst])
50                .map(|(c, w)| c + t * (c - w))
51                .collect()
52        };
53
54        let reflected = point(1.0);
55        let f_r = f(&reflected);
56        if f_r < values[best] {
57            let expanded = point(2.0);
58            let f_e = f(&expanded);
59            if f_e < f_r {
60                simplex[worst] = expanded;
61                values[worst] = f_e;
62            } else {
63                simplex[worst] = reflected;
64                values[worst] = f_r;
65            }
66        } else if f_r < values[second_worst] {
67            simplex[worst] = reflected;
68            values[worst] = f_r;
69        } else {
70            let contracted = if f_r < values[worst] { point(0.5) } else { point(-0.5) };
71            let f_c = f(&contracted);
72            if f_c < values[worst].min(f_r) {
73                simplex[worst] = contracted;
74                values[worst] = f_c;
75            } else {
76                // shrink toward the best vertex
77                let anchor = simplex[best].clone();
78                for idx in 0..=n {
79                    if idx != best {
80                        simplex[idx] = simplex[idx]
81                            .iter()
82                            .zip(&anchor)
83                            .map(|(v, b)| b + 0.5 * (v - b))
84                            .collect();
85                        values[idx] = f(&simplex[idx]);
86                    }
87                }
88            }
89        }
90    }
91    let (best, &value) =
92        values.iter().enumerate().min_by(|a, b| a.1.total_cmp(b.1)).expect("non-empty");
93    OptimResult { x: simplex[best].clone(), value, iterations: cfg.max_iter, converged: false }
94}
95
96#[cfg(test)]
97mod tests {
98    use super::*;
99
100    #[test]
101    fn minimizes_rosenbrock_without_gradients() {
102        let f = |x: &[f64]| (1.0 - x[0]).powi(2) + 100.0 * (x[1] - x[0] * x[0]).powi(2);
103        let r = nelder_mead(&OptimConfig::new(1e-12, 2000), &f, &[-1.2, 1.0]);
104        assert!((r.x[0] - 1.0).abs() < 1e-4 && (r.x[1] - 1.0).abs() < 1e-4, "{r:?}");
105    }
106
107    #[test]
108    fn handles_a_non_smooth_objective() {
109        // |x| + |y - 2|: kinked at the optimum, gradients undefined there
110        let f = |x: &[f64]| x[0].abs() + (x[1] - 2.0).abs();
111        let r = nelder_mead(&OptimConfig::new(1e-12, 2000), &f, &[3.0, -3.0]);
112        assert!(r.x[0].abs() < 1e-4 && (r.x[1] - 2.0).abs() < 1e-4, "{r:?}");
113    }
114}