refactor
This commit is contained in:
@@ -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>
|
||||
|
||||
@@ -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++)
|
||||
|
||||
Reference in New Issue
Block a user