Browse Source

Added more documentation to NelderMeadSimplex and changed style to match existing code

unified_optimization
Erik Ovegard 11 years ago
parent
commit
c942f60eef
  1. 74
      src/Numerics/Optimization/NelderMeadSimplex.cs

74
src/Numerics/Optimization/NelderMeadSimplex.cs

@ -38,6 +38,13 @@ using System.Text;
namespace MathNet.Numerics.Optimization
{
/// <summary>
/// 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
/// </summary>
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
}
/// <summary>
/// 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
/// </summary>
/// <param name="objectiveFunction">The objective function, no gradient or hessian needed</param>
/// <param name="initialGuess">The intial guess</param>
/// <returns>The minimum point</returns>
public MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector<double> 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);
}
/// <summary>
/// Finds the minimum of the objective function with an intial pertubation
/// </summary>
/// <param name="objectiveFunction">The objective function, no gradient or hessian needed</param>
/// <param name="initialGuess">The intial guess</param>
/// <param name="initalPertubation">The inital pertubation</param>
/// <returns>The minimum point</returns>
public MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector<double> initialGuess, Vector<double> 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<double>[] vertices = _initializeVertices(simplexConstants);
Vector<double>[] 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
/// </summary>
/// <param name="vertices"></param>
/// <returns></returns>
private static double[] _initializeErrorValues(Vector<double>[] vertices, IObjectiveFunction objectiveFunction)
private static double[] InitializeErrorValues(Vector<double>[] 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
/// <param name="errorProfile"></param>
/// <param name="errorValues"></param>
/// <returns></returns>
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
/// </summary>
/// <param name="errorValues"></param>
/// <returns></returns>
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
/// </summary>
/// <param name="simplexConstants"></param>
/// <returns></returns>
private static Vector<double>[] _initializeVertices(SimplexConstant[] simplexConstants)
private static Vector<double>[] InitializeVertices(SimplexConstant[] simplexConstants)
{
int numDimensions = simplexConstants.Length;
Vector<double>[] vertices = new Vector<double>[numDimensions + 1];
@ -246,11 +274,11 @@ namespace MathNet.Numerics.Optimization
/// <param name="vertices"></param>
/// <param name="errorValues"></param>
/// <returns></returns>
private static double _tryToScaleSimplex(double scaleFactor, ref ErrorProfile errorProfile, Vector<double>[] vertices,
private static double TryToScaleSimplex(double scaleFactor, ref ErrorProfile errorProfile, Vector<double>[] vertices,
double[] errorValues, IObjectiveFunction objectiveFunction)
{
// find the centroid through which we will reflect
Vector<double> centroid = _computeCentroid(vertices, errorProfile);
Vector<double> centroid = ComputeCentroid(vertices, errorProfile);
// define the vector from the centroid to the high point
Vector<double> centroidToHighPoint = vertices[errorProfile.HighestIndex].Subtract(centroid);
@ -278,7 +306,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="errorProfile"></param>
/// <param name="vertices"></param>
/// <param name="errorValues"></param>
private static void _shrinkSimplex(ErrorProfile errorProfile, Vector<double>[] vertices, double[] errorValues,
private static void ShrinkSimplex(ErrorProfile errorProfile, Vector<double>[] vertices, double[] errorValues,
IObjectiveFunction objectiveFunction)
{
Vector<double> lowestVertex = vertices[errorProfile.LowestIndex];
@ -299,7 +327,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="vertices"></param>
/// <param name="errorProfile"></param>
/// <returns></returns>
private static Vector<double> _computeCentroid(Vector<double>[] vertices, ErrorProfile errorProfile)
private static Vector<double> ComputeCentroid(Vector<double>[] 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<double> initialGuess)
public static SimplexConstant[] CreateSimplexConstantsFromVectors(Vector<double> initialGuess, Vector<double> 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;
}

Loading…
Cancel
Save