Browse Source

Optimization: Improve error handling and reporting.

unified_optimization
Scott Stephens 14 years ago
committed by Christoph Ruegg
parent
commit
92f24241f4
  1. 52
      src/Numerics/Optimization/ConjugateGradientMinimizer.cs

52
src/Numerics/Optimization/ConjugateGradientMinimizer.cs

@ -11,18 +11,18 @@ namespace MathNet.Numerics.Optimization
public ConjugateGradientMinimizer(double gradientTolerance, int maximumIterations) public ConjugateGradientMinimizer(double gradientTolerance, int maximumIterations)
{ {
this.GradientTolerance = gradientTolerance; GradientTolerance = gradientTolerance;
this.MaximumIterations = maximumIterations; MaximumIterations = maximumIterations;
} }
public MinimizationResult FindMinimum(IObjectiveFunction objective, Vector<double> initialGuess) public MinimizationResult FindMinimum(IObjectiveFunction objective, Vector<double> initialGuess)
{ {
if (!objective.IsGradientSupported) if (!objective.IsGradientSupported)
throw new Exception("Gradient not supported in objective function, but required for ConjugateGradient minimization."); throw new IncompatibleObjectiveException("Gradient not supported in objective function, but required for ConjugateGradient minimization.");
objective.EvaluateAt(initialGuess); objective.EvaluateAt(initialGuess);
var gradient = objective.Gradient; var gradient = objective.Gradient;
ValidateGradient(gradient, initialGuess); ValidateGradient(objective);
// Check that we're not already done // Check that we're not already done
if (ExitCriteriaSatisfied(initialGuess, gradient)) if (ExitCriteriaSatisfied(initialGuess, gradient))
@ -35,11 +35,22 @@ namespace MathNet.Numerics.Optimization
var steepestDirection = -gradient; var steepestDirection = -gradient;
var searchDirection = steepestDirection; var searchDirection = steepestDirection;
double initialStepSize = 100 * GradientTolerance / (gradient * gradient); double initialStepSize = 100 * GradientTolerance / (gradient * gradient);
var result = lineSearcher.FindConformingStep(objective, searchDirection, initialStepSize);
LineSearchResult result;
try
{
result = lineSearcher.FindConformingStep(objective, searchDirection, initialStepSize);
}
catch (Exception e)
{
throw new InnerOptimizationException("Line search failed.", e);
}
objective = result.FunctionInfoAtMinimum; objective = result.FunctionInfoAtMinimum;
ValidateGradient(objective.Gradient, objective.Point); ValidateGradient(objective);
double stepSize = (objective.Point - initialGuess).Norm(2.0); double stepSize = (objective.Point - initialGuess).Norm(2.0);
// Subsequent steps // Subsequent steps
int iterations = 1; int iterations = 1;
int totalLineSearchSteps = result.Iterations; int totalLineSearchSteps = result.Iterations;
@ -57,8 +68,14 @@ namespace MathNet.Numerics.Optimization
steepestDescentResets += 1; steepestDescentResets += 1;
} }
try
result = lineSearcher.FindConformingStep(objective, searchDirection, stepSize); {
result = lineSearcher.FindConformingStep(objective, searchDirection, stepSize);
}
catch (Exception e)
{
throw new InnerOptimizationException("Line search failed.", e);
}
noLineSearchIterations += result.Iterations == 0 ? 1 : 0; noLineSearchIterations += result.Iterations == 0 ? 1 : 0;
totalLineSearchSteps += result.Iterations; totalLineSearchSteps += result.Iterations;
@ -68,27 +85,32 @@ namespace MathNet.Numerics.Optimization
iterations += 1; iterations += 1;
} }
if (iterations == MaximumIterations)
{
throw new MaximumIterationsException(String.Format("Maximum iterations ({0}) reached.", MaximumIterations));
}
return new MinimizationWithLineSearchResult(objective, iterations, MinimizationResult.ExitCondition.AbsoluteGradient, totalLineSearchSteps, noLineSearchIterations); return new MinimizationWithLineSearchResult(objective, iterations, MinimizationResult.ExitCondition.AbsoluteGradient, totalLineSearchSteps, noLineSearchIterations);
} }
private bool ExitCriteriaSatisfied(Vector<double> candidatePoint, Vector<double> gradient) bool ExitCriteriaSatisfied(Vector<double> candidatePoint, Vector<double> gradient)
{ {
return gradient.Norm(2.0) < GradientTolerance; return gradient.Norm(2.0) < GradientTolerance;
} }
private void ValidateGradient(Vector<double> gradient, Vector<double> input) void ValidateGradient(IObjectiveFunction objective)
{ {
foreach (var x in gradient) foreach (var x in objective.Gradient)
{ {
if (Double.IsNaN(x) || Double.IsInfinity(x)) if (Double.IsNaN(x) || Double.IsInfinity(x))
throw new Exception("Non-finite gradient returned."); throw new EvaluationException("Non-finite gradient returned.", objective);
} }
} }
private void ValidateObjective(double objective, Vector<double> input) void ValidateObjective(IObjectiveFunction objective)
{ {
if (Double.IsNaN(objective) || Double.IsInfinity(objective)) if (Double.IsNaN(objective.Value) || Double.IsInfinity(objective.Value))
throw new Exception("Non-finite objective function returned."); throw new EvaluationException("Non-finite objective function returned.", objective);
} }
} }
} }

Loading…
Cancel
Save