From 5aece4e4c0fecf7352ce19023e3233c2f12982b5 Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Wed, 13 Nov 2013 16:42:51 +0000 Subject: [PATCH] Add test for LogisticRegressionFunction::Evaluate(). --- src/mlpack/tests/CMakeLists.txt | 1 + src/mlpack/tests/logistic_regression_test.cpp | 41 +++++++++++++++++++ 2 files changed, 42 insertions(+) create mode 100644 src/mlpack/tests/logistic_regression_test.cpp diff --git a/src/mlpack/tests/CMakeLists.txt b/src/mlpack/tests/CMakeLists.txt index 7242ceceb7..75493995a0 100644 --- a/src/mlpack/tests/CMakeLists.txt +++ b/src/mlpack/tests/CMakeLists.txt @@ -24,6 +24,7 @@ add_executable(mlpack_test linear_regression_test.cpp load_save_test.cpp local_coordinate_coding_test.cpp + logistic_regression_test.cpp lrsdp_test.cpp lsh_test.cpp math_test.cpp diff --git a/src/mlpack/tests/logistic_regression_test.cpp b/src/mlpack/tests/logistic_regression_test.cpp new file mode 100644 index 0000000000..2b9dd8323e --- /dev/null +++ b/src/mlpack/tests/logistic_regression_test.cpp @@ -0,0 +1,41 @@ +/** + * @file logistic_regression_test.cpp + * @author Ryan Curtin + * + * Test for LogisticFunction and LogisticRegression. + */ +#include +#include + +#include +#include "old_boost_test_definitions.hpp" + +using namespace mlpack; +using namespace mlpack::regression; + +BOOST_AUTO_TEST_SUITE(LogisticRegressionTest); + +/** + * Test the LogisticFunction on a simple set of points. + */ +BOOST_AUTO_TEST_CASE(LogisticRegressionFunctionEvaluate) +{ + // Very simple fake dataset. + arma::mat data("1 1 1;" // Fake row for intercept. + "1 2 3;" + "1 2 3"); + arma::vec responses("1 1 0"); + + // Create a LogisticRegressionFunction. + LogisticRegressionFunction lrf(data, responses, 0.0 /* no regularization */); + + // These were hand-calculated using Octave. + BOOST_REQUIRE_CLOSE(lrf.Evaluate(arma::vec("1 1 1")), 7.0562141665, 1e-5); + BOOST_REQUIRE_CLOSE(lrf.Evaluate(arma::vec("0 0 0")), 2.0794415417, 1e-5); + BOOST_REQUIRE_CLOSE(lrf.Evaluate(arma::vec("-1 -1 -1")), 8.0562141665, 1e-5); + BOOST_REQUIRE_CLOSE(lrf.Evaluate(arma::vec("200 -40 -40")), 0.0, 1e-5); + BOOST_REQUIRE_CLOSE(lrf.Evaluate(arma::vec("200 -80 0")), 0.0, 1e-5); + BOOST_REQUIRE_CLOSE(lrf.Evaluate(arma::vec("200 -100 20")), 0.0, 1e-5); +} + +BOOST_AUTO_TEST_SUITE_END();