From c942f60eef267d92794083454e3cfb27a46ada4d Mon Sep 17 00:00:00 2001 From: Erik Ovegard Date: Sun, 30 Aug 2015 13:35:57 +0200 Subject: [PATCH] Added more documentation to NelderMeadSimplex and changed style to match existing code --- .../Optimization/NelderMeadSimplex.cs | 74 +++++++++++++------ 1 file changed, 50 insertions(+), 24 deletions(-) diff --git a/src/Numerics/Optimization/NelderMeadSimplex.cs b/src/Numerics/Optimization/NelderMeadSimplex.cs index 2afa1113..72bdb2ed 100644 --- a/src/Numerics/Optimization/NelderMeadSimplex.cs +++ b/src/Numerics/Optimization/NelderMeadSimplex.cs @@ -38,6 +38,13 @@ using System.Text; namespace MathNet.Numerics.Optimization { + /// + /// Class implementing the Nelder-Mead simplex algorithm, used to find a minima when no gradient is available. + /// Called fminsearch() in Matlab. A description of the algorithm can be found at + /// http://se.mathworks.com/help/matlab/math/optimizing-nonlinear-functions.html#bsgpq6p-11 + /// or + /// https://en.wikipedia.org/wiki/Nelder%E2%80%93Mead_method + /// public sealed class NelderMeadSimplex { private static readonly double JITTER = 1e-10d; // a small value used to protect against floating point noise @@ -52,12 +59,31 @@ namespace MathNet.Numerics.Optimization } /// - /// Finds the minimum of the objective function + /// Finds the minimum of the objective function without an intial pertubation, the default values used + /// by fminsearch() in Matlab are used instead + /// http://se.mathworks.com/help/matlab/math/optimizing-nonlinear-functions.html#bsgpq6p-11 /// /// The objective function, no gradient or hessian needed /// The intial guess /// The minimum point public MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector initialGuess) + { + var initalPertubation = new MathNet.Numerics.LinearAlgebra.Double.DenseVector(initialGuess.Count); + for (int i = 0; i < initialGuess.Count; i++) + { + initalPertubation[i] = initialGuess[i] == 0.0 ? 0.00025 : initialGuess[i] * 0.05; + } + return FindMinimum(objectiveFunction, initialGuess, initalPertubation); + } + + /// + /// Finds the minimum of the objective function with an intial pertubation + /// + /// The objective function, no gradient or hessian needed + /// The intial guess + /// The inital pertubation + /// The minimum point + public MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector initialGuess, Vector initalPertubation) { // confirm that we are in a position to commence if (objectiveFunction == null) @@ -66,39 +92,42 @@ namespace MathNet.Numerics.Optimization if (initialGuess == null) throw new ArgumentNullException("initialGuess", "initialGuess must be initialized"); - SimplexConstant[] simplexConstants = SimplexConstant.CreateFromVector(initialGuess); + if (initialGuess == null) + throw new ArgumentNullException("initalPertubation", "initalPertubation must be initialized, if unknown use overloaded version of FindMinimum()"); + + SimplexConstant[] simplexConstants = SimplexConstant.CreateSimplexConstantsFromVectors(initialGuess,initalPertubation); // create the initial simplex int numDimensions = simplexConstants.Length; int numVertices = numDimensions + 1; - Vector[] vertices = _initializeVertices(simplexConstants); + Vector[] vertices = InitializeVertices(simplexConstants); double[] errorValues = new double[numVertices]; int evaluationCount = 0; MinimizationResult.ExitCondition exitCondition = MinimizationResult.ExitCondition.None; ErrorProfile errorProfile; - errorValues = _initializeErrorValues(vertices, objectiveFunction); + errorValues = InitializeErrorValues(vertices, objectiveFunction); // iterate until we converge, or complete our permitted number of iterations while (true) { - errorProfile = _evaluateSimplex(errorValues); + errorProfile = EvaluateSimplex(errorValues); // see if the range in point heights is small enough to exit - if (_hasConverged(ConvergenceTolerance, errorProfile, errorValues)) + if (HasConverged(ConvergenceTolerance, errorProfile, errorValues)) { exitCondition = MinimizationResult.ExitCondition.Converged; break; } // attempt a reflection of the simplex - double reflectionPointValue = _tryToScaleSimplex(-1.0, ref errorProfile, vertices, errorValues, objectiveFunction); + double reflectionPointValue = TryToScaleSimplex(-1.0, ref errorProfile, vertices, errorValues, objectiveFunction); ++evaluationCount; if (reflectionPointValue <= errorValues[errorProfile.LowestIndex]) { // it's better than the best point, so attempt an expansion of the simplex - double expansionPointValue = _tryToScaleSimplex(2.0, ref errorProfile, vertices, errorValues, objectiveFunction); + double expansionPointValue = TryToScaleSimplex(2.0, ref errorProfile, vertices, errorValues, objectiveFunction); ++evaluationCount; } else if (reflectionPointValue >= errorValues[errorProfile.NextHighestIndex]) @@ -106,22 +135,21 @@ namespace MathNet.Numerics.Optimization // it would be worse than the second best point, so attempt a contraction to look // for an intermediate point double currentWorst = errorValues[errorProfile.HighestIndex]; - double contractionPointValue = _tryToScaleSimplex(0.5, ref errorProfile, vertices, errorValues, objectiveFunction); + double contractionPointValue = TryToScaleSimplex(0.5, ref errorProfile, vertices, errorValues, objectiveFunction); ++evaluationCount; if (contractionPointValue >= currentWorst) { // that would be even worse, so let's try to contract uniformly towards the low point; // don't bother to update the error profile, we'll do it at the start of the // next iteration - _shrinkSimplex(errorProfile, vertices, errorValues, objectiveFunction); + ShrinkSimplex(errorProfile, vertices, errorValues, objectiveFunction); evaluationCount += numVertices; // that required one function evaluation for each vertex; keep track } } // check to see if we have exceeded our alloted number of evaluations if (evaluationCount >= MaximumIterations) { - exitCondition = MinimizationResult.ExitCondition.LackOfProgress; - break; + throw new MaximumIterationsException(String.Format("Maximum iterations ({0}) reached.", MaximumIterations)); } } var regressionResult = new MinimizationResult(objectiveFunction, evaluationCount, exitCondition); @@ -134,7 +162,7 @@ namespace MathNet.Numerics.Optimization /// /// /// - private static double[] _initializeErrorValues(Vector[] vertices, IObjectiveFunction objectiveFunction) + private static double[] InitializeErrorValues(Vector[] vertices, IObjectiveFunction objectiveFunction) { double[] errorValues = new double[vertices.Length]; for (int i = 0; i < vertices.Length; i++) @@ -152,7 +180,7 @@ namespace MathNet.Numerics.Optimization /// /// /// - private static bool _hasConverged(double convergenceTolerance, ErrorProfile errorProfile, double[] errorValues) + private static bool HasConverged(double convergenceTolerance, ErrorProfile errorProfile, double[] errorValues) { double range = 2 * Math.Abs(errorValues[errorProfile.HighestIndex] - errorValues[errorProfile.LowestIndex]) / (Math.Abs(errorValues[errorProfile.HighestIndex]) + Math.Abs(errorValues[errorProfile.LowestIndex]) + JITTER); @@ -172,7 +200,7 @@ namespace MathNet.Numerics.Optimization /// /// /// - private static ErrorProfile _evaluateSimplex(double[] errorValues) + private static ErrorProfile EvaluateSimplex(double[] errorValues) { ErrorProfile errorProfile = new ErrorProfile(); if (errorValues[0] > errorValues[1]) @@ -213,7 +241,7 @@ namespace MathNet.Numerics.Optimization /// /// /// - private static Vector[] _initializeVertices(SimplexConstant[] simplexConstants) + private static Vector[] InitializeVertices(SimplexConstant[] simplexConstants) { int numDimensions = simplexConstants.Length; Vector[] vertices = new Vector[numDimensions + 1]; @@ -246,11 +274,11 @@ namespace MathNet.Numerics.Optimization /// /// /// - private static double _tryToScaleSimplex(double scaleFactor, ref ErrorProfile errorProfile, Vector[] vertices, + private static double TryToScaleSimplex(double scaleFactor, ref ErrorProfile errorProfile, Vector[] vertices, double[] errorValues, IObjectiveFunction objectiveFunction) { // find the centroid through which we will reflect - Vector centroid = _computeCentroid(vertices, errorProfile); + Vector centroid = ComputeCentroid(vertices, errorProfile); // define the vector from the centroid to the high point Vector centroidToHighPoint = vertices[errorProfile.HighestIndex].Subtract(centroid); @@ -278,7 +306,7 @@ namespace MathNet.Numerics.Optimization /// /// /// - private static void _shrinkSimplex(ErrorProfile errorProfile, Vector[] vertices, double[] errorValues, + private static void ShrinkSimplex(ErrorProfile errorProfile, Vector[] vertices, double[] errorValues, IObjectiveFunction objectiveFunction) { Vector lowestVertex = vertices[errorProfile.LowestIndex]; @@ -299,7 +327,7 @@ namespace MathNet.Numerics.Optimization /// /// /// - private static Vector _computeCentroid(Vector[] vertices, ErrorProfile errorProfile) + private static Vector ComputeCentroid(Vector[] vertices, ErrorProfile errorProfile) { int numVertices = vertices.Length; // find the centroid of all points except the worst one @@ -341,14 +369,12 @@ namespace MathNet.Numerics.Optimization set { _initialPerturbation = value; } } - public static SimplexConstant[] CreateFromVector(Vector initialGuess) + public static SimplexConstant[] CreateSimplexConstantsFromVectors(Vector initialGuess, Vector initialPertubation) { var constants = new SimplexConstant[initialGuess.Count]; - for (int i = 0; i < constants.Length;i++ ) { - double pertubation = initialGuess[i]==0.0 ? 1e-5 : initialGuess[i]*1e-5; - constants[i] = new SimplexConstant(initialGuess[i], pertubation); + constants[i] = new SimplexConstant(initialGuess[i], initialPertubation[i]); } return constants; }