From 22704ad8bc3bfd844a4c7457666b7c946572b058 Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Thu, 23 Aug 2018 22:03:33 -0400 Subject: [PATCH 1/2] Replace linear regression implementation with something faster. --- .../linear_regression/linear_regression.cpp | 35 +++++-------------- 1 file changed, 8 insertions(+), 27 deletions(-) diff --git a/src/mlpack/methods/linear_regression/linear_regression.cpp b/src/mlpack/methods/linear_regression/linear_regression.cpp index 153dcc8ede..33e2dee7d3 100644 --- a/src/mlpack/methods/linear_regression/linear_regression.cpp +++ b/src/mlpack/methods/linear_regression/linear_regression.cpp @@ -76,34 +76,15 @@ void LinearRegression::Train(const arma::mat& predictors, r = sqrt(weights) % responses; } - if (lambda != 0.0) - { - // Add the identity matrix to the predictors (this is equivalent to ridge - // regression). See http://math.stackexchange.com/questions/299481/ for - // more information. - p.insert_cols(nCols, predictors.n_rows); - p.submat(p.n_rows - predictors.n_rows, nCols, p.n_rows - 1, nCols + - predictors.n_rows - 1) = sqrt(lambda) * - arma::eye(predictors.n_rows, predictors.n_rows); - } + // Convert to this form: + // a * (X X^T) = y X^T. + // Then we'll use Armadillo to solve it. + // The total runtime of this should be O(d^2 N) + O(d^3) + O(dN). + // (assuming the SVD is used to solve it) + arma::mat cov = p * p.t() + + lambda * arma::eye(predictors.n_rows, predictors.n_rows); - // We compute the QR decomposition of the predictors. - // We transpose the predictors because they are in column major order. - arma::mat Q, R; - arma::qr(Q, R, arma::trans(p)); - - // We compute the parameters, B, like so: - // R * B = Q^T * responses - // B = Q^T * responses * R^-1 - // If lambda > 0, then we must add a bunch of empty responses. - if (lambda == 0.0) - arma::solve(parameters, R, arma::trans(r * Q)); - else - { - // Copy responses into larger vector. - r.insert_cols(nCols, p.n_cols - nCols); - arma::solve(parameters, R, arma::trans(r * Q)); - } + parameters = arma::solve(cov, p * r.t()); } void LinearRegression::Predict(const arma::mat& points, From aab0b3d81bd5b1249eda88260b1b9eb191c283ef Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Sat, 1 Sep 2018 00:00:58 -0400 Subject: [PATCH 2/2] Fix wrong variable usage. --- src/mlpack/methods/linear_regression/linear_regression.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/mlpack/methods/linear_regression/linear_regression.cpp b/src/mlpack/methods/linear_regression/linear_regression.cpp index 33e2dee7d3..67c41c618d 100644 --- a/src/mlpack/methods/linear_regression/linear_regression.cpp +++ b/src/mlpack/methods/linear_regression/linear_regression.cpp @@ -82,7 +82,7 @@ void LinearRegression::Train(const arma::mat& predictors, // The total runtime of this should be O(d^2 N) + O(d^3) + O(dN). // (assuming the SVD is used to solve it) arma::mat cov = p * p.t() + - lambda * arma::eye(predictors.n_rows, predictors.n_rows); + lambda * arma::eye(p.n_rows, p.n_rows); parameters = arma::solve(cov, p * r.t()); }