diff --git a/src/mlpack/methods/linear_svm/linear_svm_main.cpp b/src/mlpack/methods/linear_svm/linear_svm_main.cpp index 19864e6cb8..a095b57874 100644 --- a/src/mlpack/methods/linear_svm/linear_svm_main.cpp +++ b/src/mlpack/methods/linear_svm/linear_svm_main.cpp @@ -108,7 +108,7 @@ PARAM_UROW_IN("labels", "A matrix containing labels (0 or 1) for the points " "in the training set (y).", "l"); // Optimizer parameters. -PARAM_DOUBLE_IN("lambda", "L2-regularization parameter for training.", "L", +PARAM_DOUBLE_IN("lambda", "L2-regularization parameter for training.", "r", 0.0001); PARAM_DOUBLE_IN("delta", "Margin of difference between correct class and other " "classes.", "d", 1.0); @@ -123,12 +123,12 @@ PARAM_DOUBLE_IN("tolerance", "Convergence tolerance for optimizer.", "e", PARAM_INT_IN("max_iterations", "Maximum iterations for optimizer (0 indicates " "no limit).", "n", 10000); PARAM_DOUBLE_IN("step_size", "Step size for parallel SGD optimizer.", - "s", 0.01); + "a", 0.01); PARAM_FLAG("shuffle", "Don't shuffle the order in which data points are " "visited for parallel SGD.", "S"); PARAM_INT_IN("epochs", "Maximum number of full epochs over dataset for " "psgd", "E", 50); -PARAM_INT_IN("seed", "Random seed. If 0, 'std::time(NULL)' is used.", "r", 0); +PARAM_INT_IN("seed", "Random seed. If 0, 'std::time(NULL)' is used.", "s", 0); class LinearSVMModel { @@ -390,9 +390,12 @@ static void mlpackMain() // Calculate accuracy, if desired. if (CLI::HasParam("test_labels")) { - arma::Row testLabels = + arma::Row testLabels; + arma::Row testRawLabels = std::move(CLI::GetParam>("test_labels")); + data::NormalizeLabels(testRawLabels, testLabels, model->mappings); + if (testSet.n_cols != testLabels.n_elem) { Log::Fatal << "Test data given with " << PRINT_PARAM_STRING("test") @@ -403,8 +406,11 @@ static void mlpackMain() numClasses = CLI::GetParam("num_classes") == 0 ? model->mappings.n_elem : CLI::GetParam("num_classes"); - vector correctClassCounts(numClasses, 0); - vector labelSize(numClasses, 0); + arma::Col correctClassCounts; + arma::Col labelSize; + correctClassCounts.zeros(numClasses); + labelSize.zeros(numClasses); + for (arma::uword i = 0; i != predictions.n_elem; ++i) { if (predictions(i) == testLabels(i))