Browse Source

Optimization: NelderMeadSimplex: allow static usage

pull/511/head
Christoph Ruegg 9 years ago
parent
commit
d633f8e608
  1. 98
      src/Numerics/Optimization/NelderMeadSimplex.cs
  2. 3
      src/UnitTests/OptimizationTests/NelderMeadSimplexTests.cs

98
src/Numerics/Optimization/NelderMeadSimplex.cs

@ -43,7 +43,7 @@ namespace MathNet.Numerics.Optimization
/// </summary> /// </summary>
public sealed class NelderMeadSimplex public sealed class NelderMeadSimplex
{ {
private static readonly double JITTER = 1e-10d; // a small value used to protect against floating point noise static readonly double JITTER = 1e-10d; // a small value used to protect against floating point noise
public double ConvergenceTolerance { get; set; } public double ConvergenceTolerance { get; set; }
public int MaximumIterations { get; set; } public int MaximumIterations { get; set; }
@ -63,13 +63,38 @@ namespace MathNet.Numerics.Optimization
/// <param name="initialGuess">The intial guess</param> /// <param name="initialGuess">The intial guess</param>
/// <returns>The minimum point</returns> /// <returns>The minimum point</returns>
public MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector<double> initialGuess) public MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector<double> initialGuess)
{
return FindMinimum(objectiveFunction, initialGuess, ConvergenceTolerance, MaximumIterations);
}
/// <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)
{
return FindMinimum(objectiveFunction, initialGuess, initalPertubation, ConvergenceTolerance, MaximumIterations);
}
/// <summary>
/// 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 static MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector<double> initialGuess, double convergenceTolerance, int maximumIterations)
{ {
var initalPertubation = new LinearAlgebra.Double.DenseVector(initialGuess.Count); var initalPertubation = new LinearAlgebra.Double.DenseVector(initialGuess.Count);
for (int i = 0; i < initialGuess.Count; i++) for (int i = 0; i < initialGuess.Count; i++)
{ {
initalPertubation[i] = initialGuess[i] == 0.0 ? 0.00025 : initialGuess[i] * 0.05; initalPertubation[i] = initialGuess[i] == 0.0 ? 0.00025 : initialGuess[i] * 0.05;
} }
return FindMinimum(objectiveFunction, initialGuess, initalPertubation); return FindMinimum(objectiveFunction, initialGuess, initalPertubation, convergenceTolerance, maximumIterations);
} }
/// <summary> /// <summary>
@ -79,7 +104,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="initialGuess">The intial guess</param> /// <param name="initialGuess">The intial guess</param>
/// <param name="initalPertubation">The inital pertubation</param> /// <param name="initalPertubation">The inital pertubation</param>
/// <returns>The minimum point</returns> /// <returns>The minimum point</returns>
public MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector<double> initialGuess, Vector<double> initalPertubation) public static MinimizationResult FindMinimum(IObjectiveFunction objectiveFunction, Vector<double> initialGuess, Vector<double> initalPertubation, double convergenceTolerance, int maximumIterations)
{ {
// confirm that we are in a position to commence // confirm that we are in a position to commence
if (objectiveFunction == null) if (objectiveFunction == null)
@ -111,7 +136,7 @@ namespace MathNet.Numerics.Optimization
errorProfile = EvaluateSimplex(errorValues); errorProfile = EvaluateSimplex(errorValues);
// see if the range in point heights is small enough to exit // see if the range in point heights is small enough to exit
if (HasConverged(ConvergenceTolerance, errorProfile, errorValues)) if (HasConverged(convergenceTolerance, errorProfile, errorValues))
{ {
exitCondition = ExitCondition.Converged; exitCondition = ExitCondition.Converged;
break; break;
@ -143,9 +168,9 @@ namespace MathNet.Numerics.Optimization
} }
} }
// check to see if we have exceeded our alloted number of evaluations // check to see if we have exceeded our alloted number of evaluations
if (evaluationCount >= MaximumIterations) if (evaluationCount >= maximumIterations)
{ {
throw new MaximumIterationsException(String.Format("Maximum iterations ({0}) reached.", MaximumIterations)); throw new MaximumIterationsException(String.Format("Maximum iterations ({0}) reached.", maximumIterations));
} }
} }
var regressionResult = new MinimizationResult(objectiveFunction, evaluationCount, exitCondition); var regressionResult = new MinimizationResult(objectiveFunction, evaluationCount, exitCondition);
@ -159,7 +184,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="vertices"></param> /// <param name="vertices"></param>
/// <param name="objectiveFunction"></param> /// <param name="objectiveFunction"></param>
/// <returns></returns> /// <returns></returns>
private static double[] InitializeErrorValues(Vector<double>[] vertices, IObjectiveFunction objectiveFunction) static double[] InitializeErrorValues(Vector<double>[] vertices, IObjectiveFunction objectiveFunction)
{ {
double[] errorValues = new double[vertices.Length]; double[] errorValues = new double[vertices.Length];
for (int i = 0; i < vertices.Length; i++) for (int i = 0; i < vertices.Length; i++)
@ -178,7 +203,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="errorProfile"></param> /// <param name="errorProfile"></param>
/// <param name="errorValues"></param> /// <param name="errorValues"></param>
/// <returns></returns> /// <returns></returns>
private static bool HasConverged(double convergenceTolerance, ErrorProfile errorProfile, double[] errorValues) static bool HasConverged(double convergenceTolerance, ErrorProfile errorProfile, double[] errorValues)
{ {
double range = 2 * Math.Abs(errorValues[errorProfile.HighestIndex] - errorValues[errorProfile.LowestIndex]) / double range = 2 * Math.Abs(errorValues[errorProfile.HighestIndex] - errorValues[errorProfile.LowestIndex]) /
(Math.Abs(errorValues[errorProfile.HighestIndex]) + Math.Abs(errorValues[errorProfile.LowestIndex]) + JITTER); (Math.Abs(errorValues[errorProfile.HighestIndex]) + Math.Abs(errorValues[errorProfile.LowestIndex]) + JITTER);
@ -198,7 +223,7 @@ namespace MathNet.Numerics.Optimization
/// </summary> /// </summary>
/// <param name="errorValues"></param> /// <param name="errorValues"></param>
/// <returns></returns> /// <returns></returns>
private static ErrorProfile EvaluateSimplex(double[] errorValues) static ErrorProfile EvaluateSimplex(double[] errorValues)
{ {
ErrorProfile errorProfile = new ErrorProfile(); ErrorProfile errorProfile = new ErrorProfile();
if (errorValues[0] > errorValues[1]) if (errorValues[0] > errorValues[1])
@ -239,7 +264,7 @@ namespace MathNet.Numerics.Optimization
/// </summary> /// </summary>
/// <param name="simplexConstants"></param> /// <param name="simplexConstants"></param>
/// <returns></returns> /// <returns></returns>
private static Vector<double>[] InitializeVertices(SimplexConstant[] simplexConstants) static Vector<double>[] InitializeVertices(SimplexConstant[] simplexConstants)
{ {
int numDimensions = simplexConstants.Length; int numDimensions = simplexConstants.Length;
Vector<double>[] vertices = new Vector<double>[numDimensions + 1]; Vector<double>[] vertices = new Vector<double>[numDimensions + 1];
@ -273,7 +298,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="errorValues"></param> /// <param name="errorValues"></param>
/// <param name="objectiveFunction"></param> /// <param name="objectiveFunction"></param>
/// <returns></returns> /// <returns></returns>
private static double TryToScaleSimplex(double scaleFactor, ref ErrorProfile errorProfile, Vector<double>[] vertices, static double TryToScaleSimplex(double scaleFactor, ref ErrorProfile errorProfile, Vector<double>[] vertices,
double[] errorValues, IObjectiveFunction objectiveFunction) double[] errorValues, IObjectiveFunction objectiveFunction)
{ {
// find the centroid through which we will reflect // find the centroid through which we will reflect
@ -306,7 +331,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="vertices"></param> /// <param name="vertices"></param>
/// <param name="errorValues"></param> /// <param name="errorValues"></param>
/// <param name="objectiveFunction"></param> /// <param name="objectiveFunction"></param>
private static void ShrinkSimplex(ErrorProfile errorProfile, Vector<double>[] vertices, double[] errorValues, static void ShrinkSimplex(ErrorProfile errorProfile, Vector<double>[] vertices, double[] errorValues,
IObjectiveFunction objectiveFunction) IObjectiveFunction objectiveFunction)
{ {
Vector<double> lowestVertex = vertices[errorProfile.LowestIndex]; Vector<double> lowestVertex = vertices[errorProfile.LowestIndex];
@ -327,7 +352,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="vertices"></param> /// <param name="vertices"></param>
/// <param name="errorProfile"></param> /// <param name="errorProfile"></param>
/// <returns></returns> /// <returns></returns>
private static Vector<double> ComputeCentroid(Vector<double>[] vertices, ErrorProfile errorProfile) static Vector<double> ComputeCentroid(Vector<double>[] vertices, ErrorProfile errorProfile)
{ {
int numVertices = vertices.Length; int numVertices = vertices.Length;
// find the centroid of all points except the worst one // find the centroid of all points except the worst one
@ -342,32 +367,21 @@ namespace MathNet.Numerics.Optimization
return centroid.Multiply(1.0d / (numVertices - 1)); return centroid.Multiply(1.0d / (numVertices - 1));
} }
private sealed class SimplexConstant sealed class SimplexConstant
{ {
private double _value;
private double _initialPerturbation;
public SimplexConstant(double value, double initialPerturbation) public SimplexConstant(double value, double initialPerturbation)
{ {
_value = value; Value = value;
_initialPerturbation = initialPerturbation; InitialPerturbation = initialPerturbation;
} }
/// <summary> /// <summary>
/// The value of the constant /// The value of the constant
/// </summary> /// </summary>
public double Value public double Value { get; }
{
get { return _value; }
set { _value = value; }
}
// The size of the initial perturbation // The size of the initial perturbation
public double InitialPerturbation public double InitialPerturbation { get; }
{
get { return _initialPerturbation; }
set { _initialPerturbation = value; }
}
public static SimplexConstant[] CreateSimplexConstantsFromVectors(Vector<double> initialGuess, Vector<double> initialPertubation) public static SimplexConstant[] CreateSimplexConstantsFromVectors(Vector<double> initialGuess, Vector<double> initialPertubation)
{ {
@ -380,29 +394,11 @@ namespace MathNet.Numerics.Optimization
} }
} }
private sealed class ErrorProfile sealed class ErrorProfile
{ {
private int _highestIndex; public int HighestIndex { get; set; }
private int _nextHighestIndex; public int NextHighestIndex { get; set; }
private int _lowestIndex; public int LowestIndex { get; set; }
public int HighestIndex
{
get { return _highestIndex; }
set { _highestIndex = value; }
}
public int NextHighestIndex
{
get { return _nextHighestIndex; }
set { _nextHighestIndex = value; }
}
public int LowestIndex
{
get { return _lowestIndex; }
set { _lowestIndex = value; }
}
} }
} }
} }

3
src/UnitTests/OptimizationTests/NelderMeadSimplexTests.cs

@ -111,9 +111,8 @@ namespace MathNet.Numerics.UnitTests.OptimizationTests
public void Mgh_Tests(TestFunctions.TestCase test_case) public void Mgh_Tests(TestFunctions.TestCase test_case)
{ {
var obj = new MghObjectiveFunction(test_case.Function, true, true); var obj = new MghObjectiveFunction(test_case.Function, true, true);
var solver = new NelderMeadSimplex(1e-8, 1000);
var result = solver.FindMinimum(obj, test_case.InitialGuess); var result = NelderMeadSimplex.FindMinimum(obj, test_case.InitialGuess, 1e-8, 1000);
if (test_case.MinimizingPoint != null) if (test_case.MinimizingPoint != null)
{ {

Loading…
Cancel
Save