From b097abd3b796fa0442365af3d91f069f9d461dc3 Mon Sep 17 00:00:00 2001 From: Ayush Date: Fri, 8 Mar 2019 23:03:51 +0530 Subject: [PATCH] refactor --- src/mlpack/methods/linear_svm/linear_svm.hpp | 8 +++++-- .../methods/linear_svm/linear_svm_impl.hpp | 10 ++++++--- src/mlpack/tests/linear_svm_test.cpp | 22 +++++-------------- 3 files changed, 19 insertions(+), 21 deletions(-) diff --git a/src/mlpack/methods/linear_svm/linear_svm.hpp b/src/mlpack/methods/linear_svm/linear_svm.hpp index b3ed191107..da164a0b62 100644 --- a/src/mlpack/methods/linear_svm/linear_svm.hpp +++ b/src/mlpack/methods/linear_svm/linear_svm.hpp @@ -107,13 +107,17 @@ class LinearSVM * value of lambda is 0.0001. Be sure to use Train() before calling * Classify() or ComputeAccuracy(), otherwise the results may be meaningless. * + * @param inputSize Size of the input feature vector. * @param numClasses Number of classes for classification. * @param lambda L2-regularization constant. * @paran delta Margin of difference between correct class and other classes. + * @param fitIntercept add intercept term or not. */ - LinearSVM(const size_t numClasses = 0, + LinearSVM(const size_t inputSize, + const size_t numClasses = 0, const double lambda = 0.0001, - const double delta = 1.0); + const double delta = 1.0, + const bool fitIntercept = false); /** * Classify the given points, returning the predicted labels for each point. diff --git a/src/mlpack/methods/linear_svm/linear_svm_impl.hpp b/src/mlpack/methods/linear_svm/linear_svm_impl.hpp index d7b401d8b7..789cba8f1a 100644 --- a/src/mlpack/methods/linear_svm/linear_svm_impl.hpp +++ b/src/mlpack/methods/linear_svm/linear_svm_impl.hpp @@ -38,14 +38,18 @@ LinearSVM::LinearSVM( template LinearSVM::LinearSVM( + const size_t inputSize, const size_t numClasses, const double lambda, - const double delta) : + const double delta, + const bool fitIntercept) : numClasses(numClasses), lambda(lambda), - delta(delta) + delta(delta), + fitIntercept(fitIntercept) { - // No training to do here. + LinearSVMFunction::InitializeWeights( + parameters, inputSize, numClasses, fitIntercept); } template diff --git a/src/mlpack/tests/linear_svm_test.cpp b/src/mlpack/tests/linear_svm_test.cpp index 6470eec6a5..6531df981c 100644 --- a/src/mlpack/tests/linear_svm_test.cpp +++ b/src/mlpack/tests/linear_svm_test.cpp @@ -454,7 +454,6 @@ BOOST_AUTO_TEST_CASE(LinearSVMLGFGSSimpleTest) { const size_t numClasses = 2; const double lambda = 0.0001; - const double delta = 1.0; // A very simple fake dataset arma::mat dataset = "2 0 0;" @@ -467,8 +466,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMLGFGSSimpleTest) arma::Row labels = "1 0 1"; // Create a linear svm object using L-BFGS optimizer. - LinearSVM lsvm(dataset, labels, numClasses, lambda, - delta, false, ens::L_BFGS()); + LinearSVM lsvm(dataset, labels, numClasses, lambda); // Compare training accuracy to 1. const double acc = lsvm.ComputeAccuracy(dataset, labels); @@ -518,7 +516,6 @@ BOOST_AUTO_TEST_CASE(LinearSVMLBFGSTwoClasses) const size_t inputSize = 3; const size_t numClasses = 2; const double lambda = 0.5; - const double delta = 1.0; // Generate two-Gaussian dataset. GaussianDistribution g1(arma::vec("1.0 9.0 1.0"), arma::eye(3, 3)); @@ -539,8 +536,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMLBFGSTwoClasses) } // Create a linear svm object using L-BFGS optimizer. - LinearSVM lsvm(data, labels, numClasses, lambda, - delta, false, ens::L_BFGS()); + LinearSVM lsvm(data, labels, numClasses, lambda); // Compare training accuracy to 1. const double acc = lsvm.ComputeAccuracy(data, labels); @@ -652,7 +648,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMDeltaLBFGSTwoClasses) // Create a linear svm object using L-BFGS optimizer. LinearSVM lsvm(data, labels, numClasses, lambda, - delta, false, ens::L_BFGS()); + delta); // Compare training accuracy to 1. const double acc = lsvm.ComputeAccuracy(data, labels); @@ -815,7 +811,6 @@ BOOST_AUTO_TEST_CASE(LinearSVMLBFGSMultipleClasses) const size_t inputSize = 5; const size_t numClasses = 5; const double lambda = 0.5; - const double delta = 1.0; // Generate five-Gaussian dataset. arma::mat identity = arma::eye(5, 5); @@ -855,8 +850,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMLBFGSMultipleClasses) } // Train linear svm object using L-BFGS optimizer. - LinearSVM lsvm(data, labels, numClasses, lambda, - delta, false, ens::L_BFGS()); + LinearSVM lsvm(data, labels, numClasses, lambda); // Compare training accuracy to 1. const double acc = lsvm.ComputeAccuracy(data, labels); @@ -935,7 +929,6 @@ BOOST_AUTO_TEST_CASE(LinearSVMClassifySinglePointTest) const size_t inputSize = 5; const size_t numClasses = 5; const double lambda = 0.5; - const double delta = 1.0; // Generate five-Gaussian dataset. arma::mat identity = arma::eye(5, 5); @@ -975,8 +968,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMClassifySinglePointTest) } // Train linear svm object. - LinearSVM lsvm(data, labels, numClasses, lambda, - delta, false, ens::L_BFGS()); + LinearSVM lsvm(data, labels, numClasses, lambda); // Create test dataset. for (size_t i = 0; i < points / 5; i++) @@ -1023,7 +1015,6 @@ BOOST_AUTO_TEST_CASE(SinglePointClassifyTest) const size_t inputSize = 5; const size_t numClasses = 5; const double lambda = 0.5; - const double delta = 1.0; // Generate five-Gaussian dataset. arma::mat identity = arma::eye(5, 5); @@ -1063,8 +1054,7 @@ BOOST_AUTO_TEST_CASE(SinglePointClassifyTest) } // Train linear svm object. - LinearSVM lsvm(data, labels, numClasses, lambda, - delta, false, ens::L_BFGS()); + LinearSVM lsvm(data, labels, numClasses, lambda); // Create test dataset. for (size_t i = 0; i < points / 5; i++)