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>
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 int MaximumIterations { get; set; }
@ -63,13 +63,38 @@ namespace MathNet.Numerics.Optimization
/// <param name="initialGuess">The intial guess</param>
/// <returns>The minimum point</returns>
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);
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);
return FindMinimum(objectiveFunction, initialGuess, initalPertubation, convergenceTolerance, maximumIterations);
}
/// <summary>
@ -79,7 +104,7 @@ namespace MathNet.Numerics.Optimization
/// <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)
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
if (objectiveFunction == null)
@ -111,7 +136,7 @@ namespace MathNet.Numerics.Optimization
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 = ExitCondition.Converged;
break;
@ -143,9 +168,9 @@ namespace MathNet.Numerics.Optimization
}
}
// 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);
@ -159,7 +184,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="vertices"></param>
/// <param name="objectiveFunction"></param>
/// <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];
for (int i = 0; i < vertices.Length; i++)
@ -178,7 +203,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)
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);
@ -198,7 +223,7 @@ namespace MathNet.Numerics.Optimization
/// </summary>
/// <param name="errorValues"></param>
/// <returns></returns>
private static ErrorProfile EvaluateSimplex(double[] errorValues)
static ErrorProfile EvaluateSimplex(double[] errorValues)
{
ErrorProfile errorProfile = new ErrorProfile();
if (errorValues[0] > errorValues[1])
@ -239,7 +264,7 @@ namespace MathNet.Numerics.Optimization
/// </summary>
/// <param name="simplexConstants"></param>
/// <returns></returns>
private static Vector<double>[] InitializeVertices(SimplexConstant[] simplexConstants)
static Vector<double>[] InitializeVertices(SimplexConstant[] simplexConstants)
{
int numDimensions = simplexConstants.Length;
Vector<double>[] vertices = new Vector<double>[numDimensions + 1];
@ -273,7 +298,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="errorValues"></param>
/// <param name="objectiveFunction"></param>
/// <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)
{
// find the centroid through which we will reflect
@ -306,7 +331,7 @@ namespace MathNet.Numerics.Optimization
/// <param name="vertices"></param>
/// <param name="errorValues"></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)
{
Vector<double> lowestVertex = vertices[errorProfile.LowestIndex];
@ -327,7 +352,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)
static Vector<double> ComputeCentroid(Vector<double>[] vertices, ErrorProfile errorProfile)
{
int numVertices = vertices.Length;
// 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));
}
private sealed class SimplexConstant
sealed class SimplexConstant
{
private double _value;
private double _initialPerturbation;
public SimplexConstant(double value, double initialPerturbation)
{
_value = value;
_initialPerturbation = initialPerturbation;
Value = value;
InitialPerturbation = initialPerturbation;
}
/// <summary>
/// The value of the constant
/// </summary>
public double Value
{
get { return _value; }
set { _value = value; }
}
public double Value { get; }
// The size of the initial perturbation
public double InitialPerturbation
{
get { return _initialPerturbation; }
set { _initialPerturbation = value; }
}
public double InitialPerturbation { get; }
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;
private int _nextHighestIndex;
private int _lowestIndex;
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; }
}
public int HighestIndex { get; set; }
public int NextHighestIndex { get; set; }
public int LowestIndex { get; set; }
}
}
}

3
src/UnitTests/OptimizationTests/NelderMeadSimplexTests.cs

@ -111,9 +111,8 @@ namespace MathNet.Numerics.UnitTests.OptimizationTests
public void Mgh_Tests(TestFunctions.TestCase test_case)
{
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)
{

Loading…
Cancel
Save