rustyqlib/core/optimization/
nelder_mead.rs1use super::{OptimConfig, OptimResult};
6
7pub fn nelder_mead(cfg: &OptimConfig, f: &dyn Fn(&[f64]) -> f64, x0: &[f64]) -> OptimResult {
12 let n = x0.len();
13 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 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 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 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 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}