main/mean_shift_test.cpp from boost to catch2
This commit is contained in:
@@ -97,7 +97,6 @@ add_executable(mlpack_test
|
||||
main_tests/local_coordinate_coding_test.cpp
|
||||
main_tests/logistic_regression_test.cpp
|
||||
main_tests/lsh_test.cpp
|
||||
main_tests/mean_shift_test.cpp
|
||||
main_tests/nbc_test.cpp
|
||||
main_tests/nmf_test.cpp
|
||||
main_tests/pca_test.cpp
|
||||
@@ -171,6 +170,7 @@ add_executable(mlpack_catch_test
|
||||
main_tests/kmeans_test.cpp
|
||||
main_tests/knn_test.cpp
|
||||
main_tests/linear_regression_test.cpp
|
||||
main_tests/mean_shift_test.cpp
|
||||
main_tests/nca_test.cpp
|
||||
main_tests/preprocess_binarize_test.cpp
|
||||
main_tests/preprocess_imputer_test.cpp
|
||||
|
||||
@@ -12,15 +12,16 @@
|
||||
#include <string>
|
||||
|
||||
#define BINDING_TYPE BINDING_TYPE_TEST
|
||||
static const std::string testName = "MeanShift";
|
||||
|
||||
#include <mlpack/core.hpp>
|
||||
static const std::string testName = "MeanShift";
|
||||
|
||||
#include <mlpack/core/util/mlpack_main.hpp>
|
||||
#include <mlpack/methods/mean_shift/mean_shift_main.cpp>
|
||||
#include "test_helper.hpp"
|
||||
|
||||
#include <boost/test/unit_test.hpp>
|
||||
#include "../test_tools.hpp"
|
||||
#include "test_helper.hpp"
|
||||
#include "../test_catch_tools.hpp"
|
||||
#include "../catch.hpp"
|
||||
|
||||
using namespace mlpack;
|
||||
|
||||
@@ -48,13 +49,13 @@ static void ResetSettings()
|
||||
IO::RestoreSettings(testName);
|
||||
}
|
||||
|
||||
BOOST_FIXTURE_TEST_SUITE(MeanShiftMainTest, MeanShiftTestFixture);
|
||||
|
||||
/**
|
||||
* Ensure that the output has 1 extra row for the labels and
|
||||
* check the number of points for output remain the same.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(MeanShiftOutputDimensionTest)
|
||||
TEST_CASE_METHOD(
|
||||
MeanShiftTestFixture, "MeanShiftOutputDimensionTest",
|
||||
"[MeanShiftMainTest][BindingTests]")
|
||||
{
|
||||
arma::mat x;
|
||||
x.randu(3, 100); // 100 points in 3 dimension
|
||||
@@ -65,16 +66,18 @@ BOOST_AUTO_TEST_CASE(MeanShiftOutputDimensionTest)
|
||||
mlpackMain();
|
||||
|
||||
// Now check that the output has 1 extra row for labels.
|
||||
BOOST_REQUIRE_EQUAL(IO::GetParam<arma::mat>("output").n_rows, 3 + 1);
|
||||
REQUIRE(IO::GetParam<arma::mat>("output").n_rows == 3 + 1);
|
||||
// Check number of output points are the same.
|
||||
BOOST_REQUIRE_EQUAL(IO::GetParam<arma::mat>("output").n_cols, 100);
|
||||
REQUIRE(IO::GetParam<arma::mat>("output").n_cols == 100);
|
||||
}
|
||||
|
||||
/**
|
||||
* Ensure that if we ask for labels_only, output has 1 row and
|
||||
* same number of columns for each point's label.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(MeanShiftLabelOnlyOutputDimensionTest)
|
||||
TEST_CASE_METHOD(
|
||||
MeanShiftTestFixture, "MeanShiftLabelOnlyOutputDimensionTest",
|
||||
"[MeanShiftMainTest][BindingTests]")
|
||||
{
|
||||
arma::mat x;
|
||||
x.randu(3, 100); // 100 points in 3 dimension
|
||||
@@ -86,9 +89,9 @@ BOOST_AUTO_TEST_CASE(MeanShiftLabelOnlyOutputDimensionTest)
|
||||
mlpackMain();
|
||||
|
||||
// Check that there is only 1 row containing all the labels.
|
||||
BOOST_REQUIRE_EQUAL(IO::GetParam<arma::mat>("output").n_rows, 1);
|
||||
REQUIRE(IO::GetParam<arma::mat>("output").n_rows == 1);
|
||||
// Check number of output points are the same.
|
||||
BOOST_REQUIRE_EQUAL(IO::GetParam<arma::mat>("output").n_cols, 100);
|
||||
REQUIRE(IO::GetParam<arma::mat>("output").n_cols == 100);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -96,7 +99,9 @@ BOOST_AUTO_TEST_CASE(MeanShiftLabelOnlyOutputDimensionTest)
|
||||
* and check the number of points remain the same if the --in_place
|
||||
* flag is set.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(MeanShiftInPlaceTest)
|
||||
TEST_CASE_METHOD(
|
||||
MeanShiftTestFixture, "MeanShiftInPlaceTest",
|
||||
"[MeanShiftMainTest][BindingTests]")
|
||||
{
|
||||
arma::mat x;
|
||||
if (!data::Load("iris_test.csv", x))
|
||||
@@ -113,16 +118,18 @@ BOOST_AUTO_TEST_CASE(MeanShiftInPlaceTest)
|
||||
mlpackMain();
|
||||
|
||||
// Now check that the output has 1 extra row for labels.
|
||||
BOOST_REQUIRE_EQUAL(IO::GetParam<arma::mat>("output").n_rows, numRows + 1);
|
||||
REQUIRE(IO::GetParam<arma::mat>("output").n_rows == numRows + 1);
|
||||
// Check number of output points are the same.
|
||||
BOOST_REQUIRE_EQUAL(IO::GetParam<arma::mat>("output").n_cols, numCols);
|
||||
REQUIRE(IO::GetParam<arma::mat>("output").n_cols == numCols);
|
||||
}
|
||||
|
||||
/**
|
||||
* Ensure that force_convergence is used by testing that the
|
||||
* force_convergence flag makes a difference in the program.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(MeanShiftForceConvergenceTest)
|
||||
TEST_CASE_METHOD(
|
||||
MeanShiftTestFixture, "MeanShiftForceConvergenceTest",
|
||||
"[MeanShiftMainTest][BindingTests]")
|
||||
{
|
||||
arma::mat x;
|
||||
if (!data::Load("iris_test.csv", x))
|
||||
@@ -150,14 +157,16 @@ BOOST_AUTO_TEST_CASE(MeanShiftForceConvergenceTest)
|
||||
|
||||
const int numCentroids2 = IO::GetParam<arma::mat>("centroid").n_cols;
|
||||
// Resulting number of centroids should be different.
|
||||
BOOST_REQUIRE_NE(numCentroids1, numCentroids2);
|
||||
REQUIRE(numCentroids1 != numCentroids2);
|
||||
}
|
||||
|
||||
/**
|
||||
* Ensure that radius is used by testing that the radius
|
||||
* makes a difference in the program.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(MeanShiftRadiusTest)
|
||||
TEST_CASE_METHOD(
|
||||
MeanShiftTestFixture, "MeanShiftRadiusTest",
|
||||
"[MeanShiftMainTest][BindingTests]")
|
||||
{
|
||||
arma::mat x;
|
||||
if (!data::Load("iris_test.csv", x))
|
||||
@@ -183,14 +192,16 @@ BOOST_AUTO_TEST_CASE(MeanShiftRadiusTest)
|
||||
|
||||
const int numCentroids2 = IO::GetParam<arma::mat>("centroid").n_cols;
|
||||
// Resulting number of centroids should be different.
|
||||
BOOST_REQUIRE_NE(numCentroids1, numCentroids2);
|
||||
REQUIRE(numCentroids1 != numCentroids2);
|
||||
}
|
||||
|
||||
/**
|
||||
* Ensure that max_iterations is used by testing that the
|
||||
* max_iteration makes a difference in the program.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(MeanShiftMaxIterationsTest)
|
||||
TEST_CASE_METHOD(
|
||||
MeanShiftTestFixture, "MeanShiftMaxIterationsTest",
|
||||
"[MeanShiftMainTest][BindingTests]")
|
||||
{
|
||||
arma::mat x;
|
||||
if (!data::Load("iris_test.csv", x))
|
||||
@@ -216,13 +227,15 @@ BOOST_AUTO_TEST_CASE(MeanShiftMaxIterationsTest)
|
||||
|
||||
const int numCentroids2 = IO::GetParam<arma::mat>("centroid").n_cols;
|
||||
// Resulting number of centroids should be different.
|
||||
BOOST_REQUIRE_NE(numCentroids1, numCentroids2);
|
||||
REQUIRE(numCentroids1 != numCentroids2);
|
||||
}
|
||||
|
||||
/**
|
||||
* Ensure that we can't specify an invalid max number of iterations.
|
||||
*/
|
||||
BOOST_AUTO_TEST_CASE(MeanShiftInvalidMaxIterationsTest)
|
||||
TEST_CASE_METHOD(
|
||||
MeanShiftTestFixture, "MeanShiftInvalidMaxIterationsTest",
|
||||
"[MeanShiftMainTest][BindingTests]")
|
||||
{
|
||||
arma::mat x;
|
||||
x.randu(3, 100); // 100 points in 3 dimension
|
||||
@@ -233,8 +246,6 @@ BOOST_AUTO_TEST_CASE(MeanShiftInvalidMaxIterationsTest)
|
||||
SetInputParam("max_iterations", (int) -1);
|
||||
|
||||
Log::Fatal.ignoreInput = true;
|
||||
BOOST_REQUIRE_THROW(mlpackMain(), std::runtime_error);
|
||||
REQUIRE_THROW_AS(mlpackMain(), std::runtime_error);
|
||||
Log::Fatal.ignoreInput = false;
|
||||
}
|
||||
|
||||
BOOST_AUTO_TEST_SUITE_END();
|
||||
|
||||
Reference in New Issue
Block a user