Merge pull request #1168 from haritha1313/ams_grad
Add AMSGrad optimizer implementation.
This commit is contained in:
@@ -3,6 +3,7 @@ set(SOURCES
|
||||
adam_impl.hpp
|
||||
adam_update.hpp
|
||||
adamax_update.hpp
|
||||
amsgrad_update.hpp
|
||||
)
|
||||
|
||||
set(DIR_SRCS)
|
||||
|
||||
@@ -5,10 +5,11 @@
|
||||
* @author Marcus Edel
|
||||
* @author Vivek Pal
|
||||
*
|
||||
* Adam and AdaMax optimizer. Adam is an an algorithm for first-order gradient-
|
||||
* -based optimization of stochastic objective functions, based on adaptive
|
||||
* estimates of lower-order moments. AdaMax is simply a variant of Adam based
|
||||
* on the infinity norm.
|
||||
* Adam, AdaMax and AMSGrad optimizers. Adam is an an algorithm for
|
||||
* first-order gradient-based optimization of stochastic objective
|
||||
* functions, based on adaptive estimates of lower-order moments.
|
||||
* AdaMax is simply a variant of Adam based on the infinity norm.
|
||||
* AMSGrad is another variant of Adam with guaranteed convergence.
|
||||
*
|
||||
* mlpack is free software; you may redistribute it and/or modify it under the
|
||||
* terms of the 3-clause BSD license. You should have received a copy of the
|
||||
@@ -23,6 +24,7 @@
|
||||
#include <mlpack/core/optimizers/sgd/sgd.hpp>
|
||||
#include "adam_update.hpp"
|
||||
#include "adamax_update.hpp"
|
||||
#include "amsgrad_update.hpp"
|
||||
|
||||
namespace mlpack {
|
||||
namespace optimization {
|
||||
@@ -43,11 +45,16 @@ namespace optimization {
|
||||
* year = {2014},
|
||||
* url = {http://arxiv.org/abs/1412.6980}
|
||||
* }
|
||||
* @article{
|
||||
* title = {On the convergence of Adam and beyond},
|
||||
* url = {https://openreview.net/pdf?id=ryQu7f-RZ}
|
||||
* year = {2018}
|
||||
* }
|
||||
* @endcode
|
||||
*
|
||||
*
|
||||
* For Adam and AdaMax to work, a DecomposableFunctionType template parameter
|
||||
* is required. This class must implement the following function:
|
||||
* For Adam, AdaMax and AMSGrad to work, a DecomposableFunctionType template
|
||||
* parameter is required. This class must implement the following function:
|
||||
*
|
||||
* size_t NumFunctions();
|
||||
* double Evaluate(const arma::mat& coordinates,
|
||||
@@ -166,6 +173,8 @@ using Adam = AdamType<AdamUpdate>;
|
||||
|
||||
using AdaMax = AdamType<AdaMaxUpdate>;
|
||||
|
||||
using AMSGrad = AdamType<AMSGradUpdate>;
|
||||
|
||||
} // namespace optimization
|
||||
} // namespace mlpack
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
* @author Marcus Edel
|
||||
* @author Vivek Pal
|
||||
*
|
||||
* Implementation of the Adam and AdaMax optimizer.
|
||||
* Implementation of the Adam, AdaMax and AMSGrad optimizer.
|
||||
*
|
||||
* mlpack is free software; you may redistribute it and/or modify it under the
|
||||
* terms of the 3-clause BSD license. You should have received a copy of the
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
/**
|
||||
* @file amsgrad_update.hpp
|
||||
* @author Haritha Nair
|
||||
*
|
||||
* Implementation of AMSGrad optimizer. AMSGrad is an exponential moving average
|
||||
* optimizer that dynamically adapts over time with guaranteed convergence.
|
||||
*
|
||||
* mlpack is free software; you may redistribute it and/or modify it under the
|
||||
* terms of the 3-clause BSD license. You should have received a copy of the
|
||||
* 3-clause BSD license along with mlpack. If not, see
|
||||
* http://www.opensource.org/licenses/BSD-3-Clause for more information.
|
||||
*/
|
||||
#ifndef MLPACK_CORE_OPTIMIZERS_AMS_GRAD_AMS_GRAD_UPDATE_HPP
|
||||
#define MLPACK_CORE_OPTIMIZERS_AMS_GRAD_AMS_GRAD_UPDATE_HPP
|
||||
|
||||
#include <mlpack/prereqs.hpp>
|
||||
|
||||
namespace mlpack {
|
||||
namespace optimization {
|
||||
|
||||
/**
|
||||
* AMSGrad is an exponential moving average variant which along with having
|
||||
* benefits of optimizers like Adam and RMSProp, also guarantees convergence.
|
||||
* Unlike Adam, it uses maximum of past squared gradients rather than their
|
||||
* exponential average for updation.
|
||||
*
|
||||
* For more information, see the following.
|
||||
*
|
||||
* @code
|
||||
* @article{
|
||||
* title = {On the convergence of Adam and beyond},
|
||||
* url = {https://openreview.net/pdf?id=ryQu7f-RZ}
|
||||
* year = {2018}
|
||||
* }
|
||||
* @endcode
|
||||
*/
|
||||
class AMSGradUpdate
|
||||
{
|
||||
public:
|
||||
/**
|
||||
* Construct the AMSGrad update policy with the given parameters.
|
||||
*
|
||||
* @param epsilon The epsilon value used to initialise the squared gradient
|
||||
* parameter.
|
||||
* @param beta1 The smoothing parameter.
|
||||
* @param beta2 The second moment coefficient.
|
||||
*/
|
||||
AMSGradUpdate(const double epsilon = 1e-8,
|
||||
const double beta1 = 0.9,
|
||||
const double beta2 = 0.999) :
|
||||
epsilon(epsilon),
|
||||
beta1(beta1),
|
||||
beta2(beta2),
|
||||
iteration(0)
|
||||
{
|
||||
// Nothing to do.
|
||||
}
|
||||
|
||||
/**
|
||||
* The Initialize method is called by SGD Optimizer method before the start of
|
||||
* the iteration update process.
|
||||
*
|
||||
* @param rows Number of rows in the gradient matrix.
|
||||
* @param cols Number of columns in the gradient matrix.
|
||||
*/
|
||||
void Initialize(const size_t rows, const size_t cols)
|
||||
{
|
||||
m = arma::zeros<arma::mat>(rows, cols);
|
||||
v = arma::zeros<arma::mat>(rows, cols);
|
||||
vImproved = arma::zeros<arma::mat>(rows, cols);
|
||||
}
|
||||
|
||||
/**
|
||||
* Update step for AMSGrad.
|
||||
*
|
||||
* @param iterate Parameters that minimize the function.
|
||||
* @param stepSize Step size to be used for the given iteration.
|
||||
* @param gradient The gradient matrix.
|
||||
*/
|
||||
void Update(arma::mat& iterate,
|
||||
const double stepSize,
|
||||
const arma::mat& gradient)
|
||||
{
|
||||
// Increment the iteration counter variable.
|
||||
++iteration;
|
||||
|
||||
// And update the iterate.
|
||||
m *= beta1;
|
||||
m += (1 - beta1) * gradient;
|
||||
|
||||
v *= beta2;
|
||||
v += (1 - beta2) * (gradient % gradient);
|
||||
|
||||
const double biasCorrection1 = 1.0 - std::pow(beta1, iteration);
|
||||
const double biasCorrection2 = 1.0 - std::pow(beta2, iteration);
|
||||
|
||||
// Element wise maximum of past and present squared gradients.
|
||||
vImproved = arma::max(vImproved, v);
|
||||
|
||||
iterate -= (stepSize * std::sqrt(biasCorrection2) / biasCorrection1) *
|
||||
m / (arma::sqrt(vImproved) + epsilon);
|
||||
}
|
||||
|
||||
//! Get the value used to initialise the squared gradient parameter.
|
||||
double Epsilon() const { return epsilon; }
|
||||
//! Modify the value used to initialise the squared gradient parameter.
|
||||
double& Epsilon() { return epsilon; }
|
||||
|
||||
//! Get the smoothing parameter.
|
||||
double Beta1() const { return beta1; }
|
||||
//! Modify the smoothing parameter.
|
||||
double& Beta1() { return beta1; }
|
||||
|
||||
//! Get the second moment coefficient.
|
||||
double Beta2() const { return beta2; }
|
||||
//! Modify the second moment coefficient.
|
||||
double& Beta2() { return beta2; }
|
||||
|
||||
private:
|
||||
// The epsilon value used to initialise the squared gradient parameter.
|
||||
double epsilon;
|
||||
|
||||
// The smoothing parameter.
|
||||
double beta1;
|
||||
|
||||
// The second moment coefficient.
|
||||
double beta2;
|
||||
|
||||
// The exponential moving average of gradient values.
|
||||
arma::mat m;
|
||||
|
||||
// The exponential moving average of squared gradient values.
|
||||
arma::mat v;
|
||||
|
||||
// The optimal sqaured gradient value.
|
||||
arma::mat vImproved;
|
||||
|
||||
// The number of iterations.
|
||||
double iteration;
|
||||
};
|
||||
|
||||
} // namespace optimization
|
||||
} // namespace mlpack
|
||||
|
||||
#endif
|
||||
@@ -3,7 +3,7 @@
|
||||
* @author Vasanth Kalingeri
|
||||
* @author Vivek Pal
|
||||
*
|
||||
* Tests the Adam and AdaMax optimizer.
|
||||
* Tests the Adam, AdaMax and AMSGrad optimizer.
|
||||
*
|
||||
* mlpack is free software; you may redistribute it and/or modify it under the
|
||||
* terms of the 3-clause BSD license. You should have received a copy of the
|
||||
@@ -62,6 +62,22 @@ BOOST_AUTO_TEST_CASE(SimpleAdaMaxTestFunction)
|
||||
BOOST_REQUIRE_SMALL(coordinates[2], 0.1);
|
||||
}
|
||||
|
||||
/**
|
||||
* Tests the AMSGrad optimizer using a simple test function.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(SimpleAMSGradTestFunction)
|
||||
{
|
||||
SGDTestFunction f;
|
||||
AMSGrad optimizer(1e-3, 1, 0.9, 0.999, 1e-8, 500000, 1e-11, true);
|
||||
|
||||
arma::mat coordinates = f.GetInitialPoint();
|
||||
optimizer.Optimize(f, coordinates);
|
||||
|
||||
BOOST_REQUIRE_SMALL(coordinates[0], 0.1);
|
||||
BOOST_REQUIRE_SMALL(coordinates[1], 0.1);
|
||||
BOOST_REQUIRE_SMALL(coordinates[2], 0.1);
|
||||
}
|
||||
|
||||
/**
|
||||
* Run Adam on logistic regression and make sure the results are acceptable.
|
||||
*/
|
||||
@@ -178,4 +194,62 @@ BOOST_AUTO_TEST_CASE(AdaMaxLogisticRegressionTest)
|
||||
BOOST_REQUIRE_CLOSE(testAcc, 100.0, 0.6); // 0.6% error tolerance.
|
||||
}
|
||||
|
||||
/**
|
||||
* Run AMSGrad on logistic regression and make sure the results are acceptable.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(AMSGradLogisticRegressionTest)
|
||||
{
|
||||
// Generate a two-Gaussian dataset.
|
||||
GaussianDistribution g1(arma::vec("1.0 1.0 1.0"), arma::eye<arma::mat>(3, 3));
|
||||
GaussianDistribution g2(arma::vec("9.0 9.0 9.0"), arma::eye<arma::mat>(3, 3));
|
||||
|
||||
arma::mat data(3, 1000);
|
||||
arma::Row<size_t> responses(1000);
|
||||
for (size_t i = 0; i < 500; ++i)
|
||||
{
|
||||
data.col(i) = g1.Random();
|
||||
responses[i] = 0;
|
||||
}
|
||||
for (size_t i = 500; i < 1000; ++i)
|
||||
{
|
||||
data.col(i) = g2.Random();
|
||||
responses[i] = 1;
|
||||
}
|
||||
|
||||
// Shuffle the dataset.
|
||||
arma::uvec indices = arma::shuffle(arma::linspace<arma::uvec>(0,
|
||||
data.n_cols - 1, data.n_cols));
|
||||
arma::mat shuffledData(3, 1000);
|
||||
arma::Row<size_t> shuffledResponses(1000);
|
||||
for (size_t i = 0; i < data.n_cols; ++i)
|
||||
{
|
||||
shuffledData.col(i) = data.col(indices[i]);
|
||||
shuffledResponses[i] = responses[indices[i]];
|
||||
}
|
||||
|
||||
// Create a test set.
|
||||
arma::mat testData(3, 1000);
|
||||
arma::Row<size_t> testResponses(1000);
|
||||
for (size_t i = 0; i < 500; ++i)
|
||||
{
|
||||
testData.col(i) = g1.Random();
|
||||
testResponses[i] = 0;
|
||||
}
|
||||
for (size_t i = 500; i < 1000; ++i)
|
||||
{
|
||||
testData.col(i) = g2.Random();
|
||||
testResponses[i] = 1;
|
||||
}
|
||||
|
||||
AMSGrad amsgrad(1e-3, 1, 0.9, 0.999, 1e-8, 500000, 1e-11, true);
|
||||
LogisticRegression<> lr(shuffledData, shuffledResponses, amsgrad, 0.5);
|
||||
|
||||
// Ensure that the error is close to zero.
|
||||
const double acc = lr.ComputeAccuracy(data, responses);
|
||||
BOOST_REQUIRE_CLOSE(acc, 100.0, 0.3); // 0.3% error tolerance.
|
||||
|
||||
const double testAcc = lr.ComputeAccuracy(testData, testResponses);
|
||||
BOOST_REQUIRE_CLOSE(testAcc, 100.0, 0.6); // 0.6% error tolerance.
|
||||
}
|
||||
|
||||
BOOST_AUTO_TEST_SUITE_END();
|
||||
|
||||
Reference in New Issue
Block a user