From 739c0f7230bb75d04cd2d92db17ba32a84fd1b44 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Erik=20Oveg=C3=A5rd?= Date: Tue, 30 Apr 2019 15:00:14 +0200 Subject: [PATCH] Improve convergence for symmetrical functions --- .../NelderMeadSimplexTests.cs | 36 +++++++++++++++---- .../Optimization/NelderMeadSimplex.cs | 10 ++++++ 2 files changed, 40 insertions(+), 6 deletions(-) diff --git a/src/Numerics.Tests/OptimizationTests/NelderMeadSimplexTests.cs b/src/Numerics.Tests/OptimizationTests/NelderMeadSimplexTests.cs index 30e9eac2..c1d161d0 100644 --- a/src/Numerics.Tests/OptimizationTests/NelderMeadSimplexTests.cs +++ b/src/Numerics.Tests/OptimizationTests/NelderMeadSimplexTests.cs @@ -42,17 +42,19 @@ namespace MathNet.Numerics.UnitTests.OptimizationTests [TestFixture] public class NelderMeadSimplexTests { + private const double Tolerance = 1.0e-5; + [Test] public void NMS_FindMinimum_Rosenbrock_Easy() { var obj = ObjectiveFunction.Value(RosenbrockFunction.Value); - var solver = new NelderMeadSimplex(1e-5, maximumIterations: 1000); + var solver = new NelderMeadSimplex(Tolerance * 0.1, maximumIterations: 1000); var initialGuess = new DenseVector(new[] { 1.2, 1.2 }); var result = solver.FindMinimum(obj, initialGuess); - Assert.That(Math.Abs(result.MinimizingPoint[0] - 1.0), Is.LessThan(1e-3)); - Assert.That(Math.Abs(result.MinimizingPoint[1] - 1.0), Is.LessThan(1e-3)); + Assert.That(Math.Abs(result.MinimizingPoint[0] - 1.0), Is.LessThan(Tolerance)); + Assert.That(Math.Abs(result.MinimizingPoint[1] - 1.0), Is.LessThan(Tolerance)); } @@ -60,14 +62,14 @@ namespace MathNet.Numerics.UnitTests.OptimizationTests public void NMS_FindMinimum_Rosenbrock_Hard() { var obj = ObjectiveFunction.Value(RosenbrockFunction.Value); - var solver = new NelderMeadSimplex(1e-5, maximumIterations: 1000); + var solver = new NelderMeadSimplex(Tolerance * 0.1, maximumIterations: 1000); var initialGuess = new DenseVector(new[] { -1.2, 1.0 }); var result = solver.FindMinimum(obj,initialGuess); - Assert.That(Math.Abs(result.MinimizingPoint[0] - 1.0), Is.LessThan(1e-3)); - Assert.That(Math.Abs(result.MinimizingPoint[1] - 1.0), Is.LessThan(1e-3)); + Assert.That(Math.Abs(result.MinimizingPoint[0] - 1.0), Is.LessThan(Tolerance)); + Assert.That(Math.Abs(result.MinimizingPoint[1] - 1.0), Is.LessThan(Tolerance)); } private class MghTestCaseEnumerator : IEnumerable @@ -127,5 +129,27 @@ namespace MathNet.Numerics.UnitTests.OptimizationTests var success = (abs_min <= 1 && abs_err < 1e-3) || (abs_min > 1 && rel_err < 1e-3); Assert.That(success, "Minimal function value is not as expected."); } + + [Test] + public void SymmetricalOneDimensionalFunction() + { + var minimizer = new NelderMeadSimplex(Tolerance*0.1, 500); + + var initialVec = Vector.Build.DenseOfEnumerable(new[] { 1.0 }); + var objFunc = ObjectiveFunction.Value(xSq); + + var min = minimizer.FindMinimum(objFunc, initialVec); + var xForMinimum = min.MinimizingPoint.ToArray(); + var minimum = xSq(min.MinimizingPoint); + + Assert.AreEqual(1.0, minimum, Tolerance, "Minimal function value is not as expected."); + Assert.AreEqual(0.0, xForMinimum[0], Tolerance, "x at minimum is not as expected."); + } + + private double xSq(IEnumerable parameters) + { + var beta = parameters.ToArray(); + return 1.0 + Math.Pow(beta[0], 2); + } } } diff --git a/src/Numerics/Optimization/NelderMeadSimplex.cs b/src/Numerics/Optimization/NelderMeadSimplex.cs index 3d201f02..2b2c278a 100644 --- a/src/Numerics/Optimization/NelderMeadSimplex.cs +++ b/src/Numerics/Optimization/NelderMeadSimplex.cs @@ -129,6 +129,7 @@ namespace MathNet.Numerics.Optimization ErrorProfile errorProfile; errorValues = InitializeErrorValues(vertices, objectiveFunction); + int numTimesHasConverged = 0; // iterate until we converge, or complete our permitted number of iterations while (true) @@ -136,7 +137,16 @@ namespace MathNet.Numerics.Optimization errorProfile = EvaluateSimplex(errorValues); // see if the range in point heights is small enough to exit + // to handle the case when the function is symmetrical and extra iteration is performed if (HasConverged(convergenceTolerance, errorProfile, errorValues)) + { + numTimesHasConverged++; + } + else + { + numTimesHasConverged = 0; + } + if (numTimesHasConverged == 2) { exitCondition = ExitCondition.Converged; break;