From 2965c9f84eabc91ad9038423c38f6932da461237 Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Thu, 21 Nov 2013 14:23:02 +0000 Subject: [PATCH] Don't hold lambda in LogisticRegression because it isn't necessary. Also make predictors and responses const because we don't need to modify them. --- .../logistic_regression/logistic_regression.hpp | 13 +++++++++---- .../logistic_regression_impl.hpp | 6 ++---- 2 files changed, 11 insertions(+), 8 deletions(-) diff --git a/src/mlpack/methods/logistic_regression/logistic_regression.hpp b/src/mlpack/methods/logistic_regression/logistic_regression.hpp index 49eadcace2..a536e7e7d6 100644 --- a/src/mlpack/methods/logistic_regression/logistic_regression.hpp +++ b/src/mlpack/methods/logistic_regression/logistic_regression.hpp @@ -58,9 +58,9 @@ class LogisticRegression arma::vec& Parameters() { return parameters; } //! Return the lambda value for L2-regularization. - const double& Lambda() const { return lambda; } + const double& Lambda() const { return errorFunction.Lambda(); } //! Modify the lambda value for L2-regularization. - double& Lambda() { return lambda; } + double& Lambda() { return errorFunction().Lambda(); } double LearnModel(); @@ -76,12 +76,17 @@ class LogisticRegression double ComputeError(arma::mat& predictors, const arma::vec& responses); private: - arma::vec parameters; + //! Matrix of predictor points (X). const arma::mat& predictors; + //! Vector of responses (y). const arma::vec& responses; + //! Vector of trained parameters. + arma::vec parameters; + + //! Instantiated error function that will be optimized. LogisticRegressionFunction errorFunction; + //! Instantiated optimizer. OptimizerType optimizer; - double lambda; }; }; // namespace regression diff --git a/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp b/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp index 35821799e7..fb2b7654c4 100644 --- a/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp +++ b/src/mlpack/methods/logistic_regression/logistic_regression_impl.hpp @@ -22,8 +22,7 @@ LogisticRegression::LogisticRegression( predictors(predictors), responses(responses), errorFunction(LogisticRegressionFunction(predictors, responses, lambda)), - optimizer(OptimizerType(errorFunction)), - lambda(lambda) + optimizer(OptimizerType(errorFunction)) { parameters.zeros(predictors.n_rows + 1); } @@ -37,8 +36,7 @@ LogisticRegression::LogisticRegression( predictors(predictors), responses(responses), errorFunction(LogisticRegressionFunction(predictors, responses)), - optimizer(OptimizerType(errorFunction)), - lambda(lambda) + optimizer(OptimizerType(errorFunction)) { parameters.zeros(predictors.n_rows + 1); }