Browse Source

Merge pull request #340 from eriove/unified_optimization

Nelder-Mead simplex algorithm (unified optimization)
unified_optimization
Christoph Ruegg 11 years ago
parent
commit
87b333896d
  1. 1
      src/Numerics/Numerics.csproj
  2. 3
      src/Numerics/Optimization/MinimizationResult.cs
  3. 415
      src/Numerics/Optimization/NelderMeadSimplex.cs
  4. 107
      src/UnitTests/OptimizationTests/NelderMeadSimplexTests.cs
  5. 1
      src/UnitTests/UnitTests.csproj

1
src/Numerics/Numerics.csproj

@ -102,6 +102,7 @@
<Compile Include="LinearAlgebra\Solvers\DelegateStopCriterion.cs" />
<Compile Include="LinearAlgebra\CreateVector.cs" />
<Compile Include="LinearRegression\Options.cs" />
<Compile Include="Optimization\NelderMeadSimplex.cs" />
<Compile Include="Optimization\ObjectiveFunctions\LazyObjectiveFunctionBase.cs" />
<Compile Include="Optimization\Exceptions.cs" />
<Compile Include="Optimization\ObjectiveFunctions\ObjectiveFunctionBase.cs" />

3
src/Numerics/Optimization/MinimizationResult.cs

@ -13,7 +13,8 @@ namespace MathNet.Numerics.Optimization
WeakWolfeCriteria,
BoundTolerance,
StrongWolfeCriteria,
LackOfFunctionImprovement
LackOfFunctionImprovement,
Converged
}
public Vector<double> MinimizingPoint { get { return FunctionInfoAtMinimum.Point; } }

415
src/Numerics/Optimization/NelderMeadSimplex.cs

@ -0,0 +1,415 @@
// <copyright file="NelderMeadSimplex.cs" company="Math.NET">
// Math.NET Numerics, part of the Math.NET Project
// http://numerics.mathdotnet.com
// http://github.com/mathnet/mathnet-numerics
// http://mathnetnumerics.codeplex.com
//
// Copyright (c) 2009-2015 Math.NET
//
// Permission is hereby granted, free of charge, to any person
// obtaining a copy of this software and associated documentation
// files (the "Software"), to deal in the Software without
// restriction, including without limitation the rights to use,
// copy, modify, merge, publish, distribute, sublicense, and/or sell
// copies of the Software, and to permit persons to whom the
// Software is furnished to do so, subject to the following
// conditions:
//
// The above copyright notice and this permission notice shall be
// included in all copies or substantial portions of the Software.
//
// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
// EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES
// OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
// NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT
// HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY,
// WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING
// FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR
// OTHER DEALINGS IN THE SOFTWARE.
// </copyright>
// Converted from code relased with a MIT liscense available at https://code.google.com/p/nelder-mead-simplex/
using MathNet.Numerics.LinearAlgebra;
using System;
using System.Collections.Generic;
using System.Linq;
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
public double ConvergenceTolerance { get; set; }
public int MaximumIterations { get; set; }
public NelderMeadSimplex(double convergenceTolerance, int maximumIterations)
{
ConvergenceTolerance = convergenceTolerance;
MaximumIterations = 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 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)
throw new ArgumentNullException("objectiveFunction","ObjectiveFunction must be set to a valid ObjectiveFunctionDelegate");
if (initialGuess == null)
throw new ArgumentNullException("initialGuess", "initialGuess must be initialized");
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);
double[] errorValues = new double[numVertices];
int evaluationCount = 0;
MinimizationResult.ExitCondition exitCondition = MinimizationResult.ExitCondition.None;
ErrorProfile errorProfile;
errorValues = InitializeErrorValues(vertices, objectiveFunction);
// iterate until we converge, or complete our permitted number of iterations
while (true)
{
errorProfile = EvaluateSimplex(errorValues);
// see if the range in point heights is small enough to exit
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);
++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);
++evaluationCount;
}
else if (reflectionPointValue >= errorValues[errorProfile.NextHighestIndex])
{
// 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);
++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);
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)
{
throw new MaximumIterationsException(String.Format("Maximum iterations ({0}) reached.", MaximumIterations));
}
}
var regressionResult = new MinimizationResult(objectiveFunction, evaluationCount, exitCondition);
return regressionResult;
}
/// <summary>
/// Evaluate the objective function at each vertex to create a corresponding
/// list of error values for each vertex
/// </summary>
/// <param name="vertices"></param>
/// <param name="objectiveFunction"></param>
/// <returns></returns>
private static double[] InitializeErrorValues(Vector<double>[] vertices, IObjectiveFunction objectiveFunction)
{
double[] errorValues = new double[vertices.Length];
for (int i = 0; i < vertices.Length; i++)
{
objectiveFunction.EvaluateAt(vertices[i]);
errorValues[i] = objectiveFunction.Value;
}
return errorValues;
}
/// <summary>
/// Check whether the points in the error profile have so little range that we
/// consider ourselves to have converged
/// </summary>
/// <param name="convergenceTolerance"></param>
/// <param name="errorProfile"></param>
/// <param name="errorValues"></param>
/// <returns></returns>
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);
if (range < convergenceTolerance)
{
return true;
}
else
{
return false;
}
}
/// <summary>
/// Examine all error values to determine the ErrorProfile
/// </summary>
/// <param name="errorValues"></param>
/// <returns></returns>
private static ErrorProfile EvaluateSimplex(double[] errorValues)
{
ErrorProfile errorProfile = new ErrorProfile();
if (errorValues[0] > errorValues[1])
{
errorProfile.HighestIndex = 0;
errorProfile.NextHighestIndex = 1;
}
else
{
errorProfile.HighestIndex = 1;
errorProfile.NextHighestIndex = 0;
}
for (int index = 0; index < errorValues.Length; index++)
{
double errorValue = errorValues[index];
if (errorValue <= errorValues[errorProfile.LowestIndex])
{
errorProfile.LowestIndex = index;
}
if (errorValue > errorValues[errorProfile.HighestIndex])
{
errorProfile.NextHighestIndex = errorProfile.HighestIndex; // downgrade the current highest to next highest
errorProfile.HighestIndex = index;
}
else if (errorValue > errorValues[errorProfile.NextHighestIndex] && index != errorProfile.HighestIndex)
{
errorProfile.NextHighestIndex = index;
}
}
return errorProfile;
}
/// <summary>
/// Construct an initial simplex, given starting guesses for the constants, and
/// initial step sizes for each dimension
/// </summary>
/// <param name="simplexConstants"></param>
/// <returns></returns>
private static Vector<double>[] InitializeVertices(SimplexConstant[] simplexConstants)
{
int numDimensions = simplexConstants.Length;
Vector<double>[] vertices = new Vector<double>[numDimensions + 1];
// define one point of the simplex as the given initial guesses
var p0 = new MathNet.Numerics.LinearAlgebra.Double.DenseVector(numDimensions);
for (int i = 0; i < numDimensions; i++)
{
p0[i] = simplexConstants[i].Value;
}
// now fill in the vertices, creating the additional points as:
// P(i) = P(0) + Scale(i) * UnitVector(i)
vertices[0] = p0;
for (int i = 0; i < numDimensions; i++)
{
double scale = simplexConstants[i].InitialPerturbation;
Vector<double> unitVector = new MathNet.Numerics.LinearAlgebra.Double.DenseVector(numDimensions);
unitVector[i] = 1;
vertices[i + 1] = p0.Add(unitVector.Multiply(scale));
}
return vertices;
}
/// <summary>
/// Test a scaling operation of the high point, and replace it if it is an improvement
/// </summary>
/// <param name="scaleFactor"></param>
/// <param name="errorProfile"></param>
/// <param name="vertices"></param>
/// <param name="errorValues"></param>
/// <param name="objectiveFunction"></param>
/// <returns></returns>
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);
// define the vector from the centroid to the high point
Vector<double> centroidToHighPoint = vertices[errorProfile.HighestIndex].Subtract(centroid);
// scale and position the vector to determine the new trial point
Vector<double> newPoint = centroidToHighPoint.Multiply(scaleFactor).Add(centroid);
// evaluate the new point
objectiveFunction.EvaluateAt(newPoint);
double newErrorValue = objectiveFunction.Value;
// if it's better, replace the old high point
if (newErrorValue < errorValues[errorProfile.HighestIndex])
{
vertices[errorProfile.HighestIndex] = newPoint;
errorValues[errorProfile.HighestIndex] = newErrorValue;
}
return newErrorValue;
}
/// <summary>
/// Contract the simplex uniformly around the lowest point
/// </summary>
/// <param name="errorProfile"></param>
/// <param name="vertices"></param>
/// <param name="errorValues"></param>
/// <param name="objectiveFunction"></param>
private static void ShrinkSimplex(ErrorProfile errorProfile, Vector<double>[] vertices, double[] errorValues,
IObjectiveFunction objectiveFunction)
{
Vector<double> lowestVertex = vertices[errorProfile.LowestIndex];
for (int i = 0; i < vertices.Length; i++)
{
if (i != errorProfile.LowestIndex)
{
vertices[i] = (vertices[i].Add(lowestVertex)).Multiply(0.5);
objectiveFunction.EvaluateAt(vertices[i]);
errorValues[i] = objectiveFunction.Value;
}
}
}
/// <summary>
/// Compute the centroid of all points except the worst
/// </summary>
/// <param name="vertices"></param>
/// <param name="errorProfile"></param>
/// <returns></returns>
private static Vector<double> ComputeCentroid(Vector<double>[] vertices, ErrorProfile errorProfile)
{
int numVertices = vertices.Length;
// find the centroid of all points except the worst one
Vector<double> centroid = new MathNet.Numerics.LinearAlgebra.Double.DenseVector(numVertices - 1);
for (int i = 0; i < numVertices; i++)
{
if (i != errorProfile.HighestIndex)
{
centroid = centroid.Add(vertices[i]);
}
}
return centroid.Multiply(1.0d / (numVertices - 1));
}
private sealed class SimplexConstant
{
private double _value;
private double _initialPerturbation;
public SimplexConstant(double value, double initialPerturbation)
{
_value = value;
_initialPerturbation = initialPerturbation;
}
/// <summary>
/// The value of the constant
/// </summary>
public double Value
{
get { return _value; }
set { _value = value; }
}
// The size of the initial perturbation
public double InitialPerturbation
{
get { return _initialPerturbation; }
set { _initialPerturbation = value; }
}
public static SimplexConstant[] CreateSimplexConstantsFromVectors(Vector<double> initialGuess, Vector<double> initialPertubation)
{
var constants = new SimplexConstant[initialGuess.Count];
for (int i = 0; i < constants.Length;i++ )
{
constants[i] = new SimplexConstant(initialGuess[i], initialPertubation[i]);
}
return constants;
}
}
private 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; }
}
}
}
}

107
src/UnitTests/OptimizationTests/NelderMeadSimplexTests.cs

@ -0,0 +1,107 @@
// <copyright file="NelderMeadSimplexTests.cs" company="Math.NET">
// Math.NET Numerics, part of the Math.NET Project
// http://numerics.mathdotnet.com
// http://github.com/mathnet/mathnet-numerics
// http://mathnetnumerics.codeplex.com
//
// Copyright (c) 2009-2015 Math.NET
//
// Permission is hereby granted, free of charge, to any person
// obtaining a copy of this software and associated documentation
// files (the "Software"), to deal in the Software without
// restriction, including without limitation the rights to use,
// copy, modify, merge, publish, distribute, sublicense, and/or sell
// copies of the Software, and to permit persons to whom the
// Software is furnished to do so, subject to the following
// conditions:
//
// The above copyright notice and this permission notice shall be
// included in all copies or substantial portions of the Software.
//
// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
// EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES
// OF MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND
// NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT
// HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY,
// WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING
// FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR
// OTHER DEALINGS IN THE SOFTWARE.
// </copyright>
using NUnit.Framework;
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text;
using System.Threading.Tasks;
using MathNet.Numerics.Optimization;
using MathNet.Numerics.LinearAlgebra.Double;
namespace MathNet.Numerics.UnitTests.OptimizationTests
{
[TestFixture]
public class NelderMeadSimplexTests
{
/// <summary>
/// Test that finds the constants of a parable, function adds noise and return the mean square error
/// Copied from the test in https://code.google.com/p/nelder-mead-simplex/
/// </summary>
[Test]
public void FindParableConstantsThatMinimizesErrors()
{
var nms = new NelderMeadSimplex(1e-6, 1000);
double a = 5;
double b = 10;
IObjectiveFunction objFun = ObjectiveFunction.Value((constants)=>
{
double ssq = 0;
System.Random r = new System.Random();
for (double x = -10; x < 10; x += .1)
{
double yTrue = a * x * x + b * x + r.NextDouble();
double yRegress = constants[0] * x * x + constants[1] * x;
ssq += Math.Pow((yTrue - yRegress), 2);
}
return ssq;
});
var initialGuess = new DenseVector(2);
initialGuess[0] = 3;
initialGuess[1] = 5;
var result = nms.FindMinimum(objFun, initialGuess);
Assert.NotNull(result);
Assert.NotNull(result.MinimizingPoint);
Assert.NotNull(result.FunctionInfoAtMinimum);
Assert.That(Math.Abs(result.MinimizingPoint[0] - a), Is.LessThan(1e-2));
Assert.That(Math.Abs(result.MinimizingPoint[1] - b), Is.LessThan(1e-2));
}
[Test]
public void NMS_FindMinimum_Rosenbrock_Easy()
{
var obj = ObjectiveFunction.Value(RosenbrockFunction.Value);
var solver = new NelderMeadSimplex(1e-5, maximumIterations: 1000);
var initialGuess = new DenseVector(new[] { 1.2, 1.2 });
var result = solver.FindMinimum(obj, initialGuess);
Assert.That(Math.Abs(result.MinimizingPoint[0] - 1.0), Is.LessThan(1e-3));
Assert.That(Math.Abs(result.MinimizingPoint[1] - 1.0), Is.LessThan(1e-3));
}
[Test]
public void NMS_FindMinimum_Rosenbrock_Hard()
{
var obj = ObjectiveFunction.Value(RosenbrockFunction.Value);
var solver = new NelderMeadSimplex(1e-5, maximumIterations: 1000);
var initialGuess = new DenseVector(new[] { -1.2, 1.0 });
var result = solver.FindMinimum(obj,initialGuess);
Assert.That(Math.Abs(result.MinimizingPoint[0] - 1.0), Is.LessThan(1e-3));
Assert.That(Math.Abs(result.MinimizingPoint[1] - 1.0), Is.LessThan(1e-3));
}
}
}

1
src/UnitTests/UnitTests.csproj

@ -345,6 +345,7 @@
<Compile Include="EuclidTests\IntegerTheoryTest.cs" />
<Compile Include="LinearAlgebraTests\MatrixStorageCombinatorsTests.cs" />
<Compile Include="LinearAlgebraTests\VectorStorageCombinatorsTests.cs" />
<Compile Include="OptimizationTests\NelderMeadSimplexTests.cs" />
<Compile Include="OptimizationTests\TestGoldenSectionMinimizer.cs" />
<Compile Include="Random\SystemRandomSourceTests.cs" />
<Compile Include="OptimizationTests\BfgsTest.cs" />

Loading…
Cancel
Save