This commit is contained in:
Ayush
2019-04-07 22:05:07 -04:00
committed by Ryan Curtin
parent d247beaeb2
commit b097abd3b7
3 changed files with 19 additions and 21 deletions
+6 -2
View File
@@ -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.
@@ -38,14 +38,18 @@ LinearSVM<MatType>::LinearSVM(
template <typename MatType>
LinearSVM<MatType>::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<MatType>::InitializeWeights(
parameters, inputSize, numClasses, fitIntercept);
}
template <typename MatType>
+6 -16
View File
@@ -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<size_t> labels = "1 0 1";
// Create a linear svm object using L-BFGS optimizer.
LinearSVM<arma::mat> lsvm(dataset, labels, numClasses, lambda,
delta, false, ens::L_BFGS());
LinearSVM<arma::mat> 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<arma::mat>(3, 3));
@@ -539,8 +536,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMLBFGSTwoClasses)
}
// Create a linear svm object using L-BFGS optimizer.
LinearSVM<arma::mat> lsvm(data, labels, numClasses, lambda,
delta, false, ens::L_BFGS());
LinearSVM<arma::mat> 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<arma::mat> 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<arma::mat>(5, 5);
@@ -855,8 +850,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMLBFGSMultipleClasses)
}
// Train linear svm object using L-BFGS optimizer.
LinearSVM<arma::mat> lsvm(data, labels, numClasses, lambda,
delta, false, ens::L_BFGS());
LinearSVM<arma::mat> 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<arma::mat>(5, 5);
@@ -975,8 +968,7 @@ BOOST_AUTO_TEST_CASE(LinearSVMClassifySinglePointTest)
}
// Train linear svm object.
LinearSVM<arma::mat> lsvm(data, labels, numClasses, lambda,
delta, false, ens::L_BFGS());
LinearSVM<arma::mat> 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<arma::mat>(5, 5);
@@ -1063,8 +1054,7 @@ BOOST_AUTO_TEST_CASE(SinglePointClassifyTest)
}
// Train linear svm object.
LinearSVM<arma::mat> lsvm(data, labels, numClasses, lambda,
delta, false, ens::L_BFGS());
LinearSVM<arma::mat> lsvm(data, labels, numClasses, lambda);
// Create test dataset.
for (size_t i = 0; i < points / 5; i++)