diff --git a/src/mlpack/methods/logistic_regression/logistic_regression.hpp b/src/mlpack/methods/logistic_regression/logistic_regression.hpp index 6214c4030c..11ea3ca41a 100644 --- a/src/mlpack/methods/logistic_regression/logistic_regression.hpp +++ b/src/mlpack/methods/logistic_regression/logistic_regression.hpp @@ -125,9 +125,10 @@ class LogisticRegression * @tparam OptimizerType Type of optimizer to use to train the model. * @param predictors Input training variables. * @param responses Outputs results from input training variables. + * @return The final objective of the trained model (NaN or Inf on error) */ template - void Train(const MatType& predictors, + double Train(const MatType& predictors, const arma::Row& responses); /** @@ -145,9 +146,10 @@ class LogisticRegression * @param predictors Input training variables. * @param responses Outputs results from input training variables. * @param optimizer Instantiated optimizer with instantiated error function. + * @return The final objective of the trained model (NaN or Inf on error) */ template - void Train(const MatType& predictors, + double Train(const MatType& predictors, const arma::Row& responses, OptimizerType& optimizer); diff --git a/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp b/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp index 15d8c3dee8..1973806fad 100644 --- a/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp +++ b/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp @@ -68,16 +68,16 @@ LogisticRegression::LogisticRegression( template template -void LogisticRegression::Train(const MatType& predictors, +double LogisticRegression::Train(const MatType& predictors, const arma::Row& responses) { OptimizerType optimizer; - Train(predictors, responses, optimizer); + return Train(predictors, responses, optimizer); } template template -void LogisticRegression::Train( +double LogisticRegression::Train( const MatType& predictors, const arma::Row& responses, OptimizerType& optimizer) @@ -93,6 +93,8 @@ void LogisticRegression::Train( Log::Info << "LogisticRegression::LogisticRegression(): final objective of " << "trained model is " << out << "." << std::endl; + + return out; } template