From 60b13abcffc67c7bb44a63afe3df7edfa72c3605 Mon Sep 17 00:00:00 2001 From: Vikas Shetty Date: Sat, 19 Oct 2019 22:38:40 +0530 Subject: [PATCH] Adding callback parameters for Logistic Regression --- .../logistic_regression/logistic_regression.hpp | 15 +++++++++++---- .../logistic_regression_impl.hpp | 17 ++++++++++------- 2 files changed, 21 insertions(+), 11 deletions(-) diff --git a/src/mlpack/methods/logistic_regression/logistic_regression.hpp b/src/mlpack/methods/logistic_regression/logistic_regression.hpp index cc23483808..b948752fe4 100644 --- a/src/mlpack/methods/logistic_regression/logistic_regression.hpp +++ b/src/mlpack/methods/logistic_regression/logistic_regression.hpp @@ -123,13 +123,16 @@ class LogisticRegression * parameters vector directly with Parameters() and modify it as desired. * * @tparam OptimizerType Type of optimizer to use to train the model. + * @tparam CallbackTypes Types of Callback Functions. * @param predictors Input training variables. * @param responses Outputs results from input training variables. + * @param callbacks Callback Functions. * @return The final objective of the trained model (NaN or Inf on error) */ - template + template double Train(const MatType& predictors, - const arma::Row& responses); + const arma::Row& responses, + CallbackTypes&&... callbacks); /** * Train the LogisticRegression model with the given instantiated optimizer. @@ -143,15 +146,19 @@ class LogisticRegression * optimizer.Function().GetInitialPoint() to the current parameters vector, * accessible via Parameters(). * + * @tparam OptimizerType Type of optimizer to use to train the model. + * @tparam CallbackTypes Types of Callback Functions. * @param predictors Input training variables. * @param responses Outputs results from input training variables. * @param optimizer Instantiated optimizer with instantiated error function. + * @param callbacks Callback Functions. * @return The final objective of the trained model (NaN or Inf on error) */ - template + template double Train(const MatType& predictors, const arma::Row& responses, - OptimizerType& optimizer); + OptimizerType& optimizer, + CallbackTypes&&... callbacks); //! Return the parameters (the b vector). const arma::rowvec& Parameters() const { return parameters; } diff --git a/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp b/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp index 1973806fad..ef45d96300 100644 --- a/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp +++ b/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp @@ -67,20 +67,23 @@ LogisticRegression::LogisticRegression( } template -template -double LogisticRegression::Train(const MatType& predictors, - const arma::Row& responses) +template +double LogisticRegression::Train( + const MatType& predictors, + const arma::Row& responses, + CallbackTypes&&... callbacks) { OptimizerType optimizer; - return Train(predictors, responses, optimizer); + return Train(predictors, responses, optimizer, callbacks...); } template -template +template double LogisticRegression::Train( const MatType& predictors, const arma::Row& responses, - OptimizerType& optimizer) + OptimizerType& optimizer, + CallbackTypes&&... callbacks) { LogisticRegressionFunction errorFunction(predictors, responses, @@ -88,7 +91,7 @@ double LogisticRegression::Train( errorFunction.InitialPoint() = parameters; Timer::Start("logistic_regression_optimization"); - const double out = optimizer.Optimize(errorFunction, parameters); + const double out = optimizer.Optimize(errorFunction, parameters, callbacks...); Timer::Stop("logistic_regression_optimization"); Log::Info << "LogisticRegression::LogisticRegression(): final objective of "