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;
}