#[derive(Debug, Clone)]
pub struct NelderMeadResult {
pub optimal_point: Vec<f64>,
pub optimal_value: f64,
pub iterations: usize,
pub converged: bool,
}
#[derive(Debug, Clone)]
pub struct NelderMeadConfig {
pub max_iter: usize,
pub tolerance: f64,
pub alpha: f64,
pub gamma: f64,
pub rho: f64,
pub sigma: f64,
pub initial_step: f64,
}
impl Default for NelderMeadConfig {
fn default() -> Self {
Self {
max_iter: 1000,
tolerance: 1e-8,
alpha: 1.0,
gamma: 2.0,
rho: 0.5,
sigma: 0.5,
initial_step: 0.05,
}
}
}
#[inline]
fn sanitize_objective(value: f64) -> f64 {
if value.is_finite() {
value
} else {
f64::MAX
}
}
pub fn nelder_mead<F>(
objective: F,
initial: &[f64],
bounds: Option<&[(f64, f64)]>,
config: NelderMeadConfig,
) -> NelderMeadResult
where
F: Fn(&[f64]) -> f64,
{
let n = initial.len();
if n == 0 {
return NelderMeadResult {
optimal_point: vec![],
optimal_value: f64::NAN,
iterations: 0,
converged: false,
};
}
let mut simplex: Vec<Vec<f64>> = Vec::with_capacity(n + 1);
let mut first = initial.to_vec();
apply_bounds_in_place(&mut first, bounds);
simplex.push(first);
for i in 0..n {
let mut vertex = initial.to_vec();
let step = if initial[i].abs() > 1e-10 {
config.initial_step * initial[i].abs()
} else {
config.initial_step
};
vertex[i] += step;
apply_bounds_in_place(&mut vertex, bounds);
simplex.push(vertex);
}
let mut values: Vec<f64> = simplex
.iter()
.map(|v| sanitize_objective(objective(v)))
.collect();
let mut indices: Vec<usize> = (0..=n).collect();
let mut centroid = vec![0.0; n];
let mut reflected = vec![0.0; n];
let mut expanded = vec![0.0; n];
let mut contracted = vec![0.0; n];
let mut temp = vec![0.0; n];
let mut iterations = 0;
let mut converged = false;
while iterations < config.max_iter {
iterations += 1;
sort_simplex_indices(&mut indices, &values);
let best_idx = indices[0];
let worst_idx = indices[n];
let second_worst_idx = indices[n - 1];
if check_convergence(
&simplex,
&values,
best_idx,
worst_idx,
config.tolerance,
&mut centroid,
) {
converged = true;
break;
}
let reflected_value = match try_reflection_expansion(
&objective,
&config,
bounds,
&mut simplex,
&mut values,
worst_idx,
best_idx,
second_worst_idx,
¢roid,
&mut reflected,
&mut expanded,
) {
None => continue, Some(rv) => rv, };
if try_contraction(
&objective,
&config,
bounds,
&mut simplex,
&mut values,
worst_idx,
¢roid,
&reflected,
reflected_value,
&mut contracted,
) {
continue;
}
shrink_simplex(
&objective,
&config,
bounds,
&mut simplex,
&mut values,
best_idx,
&mut temp,
);
}
let best_idx = values
.iter()
.enumerate()
.min_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(i, _)| i)
.unwrap_or(0);
NelderMeadResult {
optimal_point: simplex[best_idx].clone(),
optimal_value: values[best_idx],
iterations,
converged,
}
}
#[inline]
fn sort_simplex_indices(indices: &mut [usize], values: &[f64]) {
for (i, idx) in indices.iter_mut().enumerate() {
*idx = i;
}
indices.sort_by(|&a, &b| {
values[a]
.partial_cmp(&values[b])
.unwrap_or(std::cmp::Ordering::Equal)
});
}
#[inline]
fn check_convergence(
simplex: &[Vec<f64>],
values: &[f64],
best_idx: usize,
worst_idx: usize,
tolerance: f64,
centroid: &mut [f64],
) -> bool {
let range = values[worst_idx] - values[best_idx];
if range < tolerance {
return true;
}
compute_centroid_into(simplex, worst_idx, centroid);
let max_dist = simplex
.iter()
.map(|v| euclidean_distance(v, centroid))
.fold(0.0, f64::max);
max_dist < tolerance
}
#[inline]
fn try_reflection_expansion<F: Fn(&[f64]) -> f64>(
objective: &F,
config: &NelderMeadConfig,
bounds: Option<&[(f64, f64)]>,
simplex: &mut [Vec<f64>],
values: &mut [f64],
worst_idx: usize,
best_idx: usize,
second_worst_idx: usize,
centroid: &[f64],
reflected: &mut [f64],
expanded: &mut [f64],
) -> Option<f64> {
reflect_into(&simplex[worst_idx], centroid, config.alpha, reflected);
apply_bounds_in_place(reflected, bounds);
let reflected_value = sanitize_objective(objective(reflected));
if reflected_value < values[second_worst_idx] && reflected_value >= values[best_idx] {
simplex[worst_idx].copy_from_slice(reflected);
values[worst_idx] = reflected_value;
return None; }
if reflected_value < values[best_idx] {
expand_into(centroid, reflected, config.gamma, expanded);
apply_bounds_in_place(expanded, bounds);
let expanded_value = sanitize_objective(objective(expanded));
if expanded_value < reflected_value {
simplex[worst_idx].copy_from_slice(expanded);
values[worst_idx] = expanded_value;
} else {
simplex[worst_idx].copy_from_slice(reflected);
values[worst_idx] = reflected_value;
}
return None; }
Some(reflected_value) }
#[inline]
fn try_contraction<F: Fn(&[f64]) -> f64>(
objective: &F,
config: &NelderMeadConfig,
bounds: Option<&[(f64, f64)]>,
simplex: &mut [Vec<f64>],
values: &mut [f64],
worst_idx: usize,
centroid: &[f64],
reflected: &[f64],
reflected_value: f64,
contracted: &mut [f64],
) -> bool {
if reflected_value < values[worst_idx] {
contract_into(centroid, reflected, config.rho, contracted);
apply_bounds_in_place(contracted, bounds);
let contracted_value = sanitize_objective(objective(contracted));
if contracted_value <= reflected_value {
simplex[worst_idx].copy_from_slice(contracted);
values[worst_idx] = contracted_value;
return true;
}
} else {
contract_into(centroid, &simplex[worst_idx], config.rho, contracted);
apply_bounds_in_place(contracted, bounds);
let contracted_value = sanitize_objective(objective(contracted));
if contracted_value < values[worst_idx] {
simplex[worst_idx].copy_from_slice(contracted);
values[worst_idx] = contracted_value;
return true;
}
}
false
}
#[inline]
fn shrink_simplex<F: Fn(&[f64]) -> f64>(
objective: &F,
config: &NelderMeadConfig,
bounds: Option<&[(f64, f64)]>,
simplex: &mut [Vec<f64>],
values: &mut [f64],
best_idx: usize,
temp: &mut [f64],
) {
let n = temp.len();
temp.copy_from_slice(&simplex[best_idx]);
for i in 0..=n {
if i != best_idx {
for j in 0..n {
simplex[i][j] = temp[j] + config.sigma * (simplex[i][j] - temp[j]);
}
apply_bounds_in_place(&mut simplex[i], bounds);
values[i] = sanitize_objective(objective(&simplex[i]));
}
}
}
fn compute_centroid_into(simplex: &[Vec<f64>], exclude_idx: usize, out: &mut [f64]) {
let count = simplex.len() - 1;
for o in out.iter_mut() {
*o = 0.0;
}
for (i, vertex) in simplex.iter().enumerate() {
if i != exclude_idx {
for (o, v) in out.iter_mut().zip(vertex.iter()) {
*o += v;
}
}
}
let inv = 1.0 / count as f64;
for o in out.iter_mut() {
*o *= inv;
}
}
fn reflect_into(point: &[f64], centroid: &[f64], alpha: f64, out: &mut [f64]) {
for ((o, c), p) in out.iter_mut().zip(centroid.iter()).zip(point.iter()) {
*o = c + alpha * (c - p);
}
}
fn expand_into(centroid: &[f64], reflected: &[f64], gamma: f64, out: &mut [f64]) {
for ((o, c), r) in out.iter_mut().zip(centroid.iter()).zip(reflected.iter()) {
*o = c + gamma * (r - c);
}
}
fn contract_into(centroid: &[f64], point: &[f64], rho: f64, out: &mut [f64]) {
for ((o, c), p) in out.iter_mut().zip(centroid.iter()).zip(point.iter()) {
*o = c + rho * (p - c);
}
}
fn apply_bounds_in_place(point: &mut [f64], bounds: Option<&[(f64, f64)]>) {
if let Some(b) = bounds {
for (i, x) in point.iter_mut().enumerate() {
if i < b.len() {
*x = x.clamp(b[i].0, b[i].1);
}
}
}
}
fn euclidean_distance(a: &[f64], b: &[f64]) -> f64 {
a.iter()
.zip(b.iter())
.map(|(x, y)| (x - y).powi(2))
.sum::<f64>()
.sqrt()
}
#[cfg(test)]
mod tests {
use super::*;
use approx::assert_relative_eq;
#[test]
fn nelder_mead_quadratic_2d() {
let result = nelder_mead(
|x| (x[0] - 2.0).powi(2) + (x[1] - 3.0).powi(2),
&[0.0, 0.0],
None,
NelderMeadConfig::default(),
);
assert!(result.converged);
assert_relative_eq!(result.optimal_point[0], 2.0, epsilon = 1e-4);
assert_relative_eq!(result.optimal_point[1], 3.0, epsilon = 1e-4);
assert_relative_eq!(result.optimal_value, 0.0, epsilon = 1e-6);
}
#[test]
fn nelder_mead_rosenbrock() {
let config = NelderMeadConfig {
max_iter: 5000,
tolerance: 1e-10,
..Default::default()
};
let result = nelder_mead(
|x| (1.0 - x[0]).powi(2) + 100.0 * (x[1] - x[0].powi(2)).powi(2),
&[0.0, 0.0],
None,
config,
);
assert_relative_eq!(result.optimal_point[0], 1.0, epsilon = 1e-3);
assert_relative_eq!(result.optimal_point[1], 1.0, epsilon = 1e-3);
}
#[test]
fn nelder_mead_1d() {
let result = nelder_mead(
|x| (x[0] - 5.0).powi(2),
&[0.0],
None,
NelderMeadConfig::default(),
);
assert!(result.converged);
assert_relative_eq!(result.optimal_point[0], 5.0, epsilon = 0.1);
}
#[test]
fn nelder_mead_with_bounds() {
let result = nelder_mead(
|x| (x[0] - 5.0).powi(2),
&[1.0],
Some(&[(0.0, 3.0)]),
NelderMeadConfig::default(),
);
assert_relative_eq!(result.optimal_point[0], 3.0, epsilon = 1e-4);
}
#[test]
fn nelder_mead_with_bounds_2d() {
let result = nelder_mead(
|x| (x[0] - 2.0).powi(2) + (x[1] - 3.0).powi(2),
&[0.5, 0.5],
Some(&[(0.0, 1.0), (0.0, 1.0)]),
NelderMeadConfig::default(),
);
assert_relative_eq!(result.optimal_point[0], 1.0, epsilon = 1e-4);
assert_relative_eq!(result.optimal_point[1], 1.0, epsilon = 1e-4);
}
#[test]
fn nelder_mead_exponential_smoothing_alpha() {
let data = [10.0, 12.0, 11.0, 13.0, 14.0, 13.0, 15.0, 16.0];
let sse = |params: &[f64]| {
let alpha = params[0];
let mut level = data[0];
let mut error_sum = 0.0;
for &y in &data[1..] {
let forecast = level;
let error = y - forecast;
error_sum += error * error;
level = alpha * y + (1.0 - alpha) * level;
}
error_sum
};
let result = nelder_mead(
sse,
&[0.5],
Some(&[(0.01, 0.99)]),
NelderMeadConfig::default(),
);
assert!(result.converged);
assert!(result.optimal_point[0] > 0.01 && result.optimal_point[0] < 0.99);
}
#[test]
fn nelder_mead_empty_initial() {
let result = nelder_mead(|_| 0.0, &[], None, NelderMeadConfig::default());
assert!(!result.converged);
assert!(result.optimal_value.is_nan());
}
#[test]
fn nelder_mead_already_optimal() {
let result = nelder_mead(
|x| (x[0] - 2.0).powi(2),
&[2.0],
None,
NelderMeadConfig::default(),
);
assert!(result.converged);
assert_relative_eq!(result.optimal_point[0], 2.0, epsilon = 1e-4);
}
#[test]
fn nelder_mead_3d() {
let result = nelder_mead(
|x| x[0].powi(2) + x[1].powi(2) + x[2].powi(2),
&[1.0, 2.0, 3.0],
None,
NelderMeadConfig::default(),
);
assert!(result.converged);
assert_relative_eq!(result.optimal_point[0], 0.0, epsilon = 1e-4);
assert_relative_eq!(result.optimal_point[1], 0.0, epsilon = 1e-4);
assert_relative_eq!(result.optimal_point[2], 0.0, epsilon = 1e-4);
}
#[test]
fn nelder_mead_config_custom() {
let config = NelderMeadConfig {
max_iter: 100,
tolerance: 1e-4,
alpha: 1.5,
gamma: 2.5,
rho: 0.4,
sigma: 0.4,
initial_step: 0.1,
};
let result = nelder_mead(|x| (x[0] - 1.0).powi(2), &[0.0], None, config);
assert_relative_eq!(result.optimal_point[0], 1.0, epsilon = 0.01);
}
#[test]
fn nelder_mead_nan_objective_handled() {
let result = nelder_mead(
|x| {
if x[0] < 0.0 {
f64::NAN
} else {
(x[0] - 3.0).powi(2)
}
},
&[1.0],
None,
NelderMeadConfig::default(),
);
assert!(result.converged);
assert!(result.optimal_value.is_finite());
assert_relative_eq!(result.optimal_point[0], 3.0, epsilon = 0.1);
}
#[test]
fn nelder_mead_inf_objective_handled() {
let result = nelder_mead(
|x| {
if x[0] < -1.0 {
f64::INFINITY
} else {
(x[0] - 2.0).powi(2)
}
},
&[1.0],
None,
NelderMeadConfig::default(),
);
assert!(result.converged);
assert!(result.optimal_value.is_finite());
assert_relative_eq!(result.optimal_point[0], 2.0, epsilon = 0.1);
}
}