diff --git a/.ci/linux-steps.yaml b/.ci/linux-steps.yaml index ad5bef434e..a69396bc76 100644 --- a/.ci/linux-steps.yaml +++ b/.ci/linux-steps.yaml @@ -24,8 +24,9 @@ steps: sudo apt-get install -y --allow-unauthenticated libopenblas-dev liblapack-dev g++ libboost1.70-dev libarmadillo-dev xz-utils libcereal-dev if [ "$(binding)" == "python" ]; then - /usr/bin/python3 -m pip install --upgrade pip - /usr/bin/python3 -m pip install --upgrade --ignore-installed setuptools cython pandas + export PYBIN=$(which python) + $PYBIN -m pip install --upgrade pip + $PYBIN -m pip install --upgrade --ignore-installed setuptools cython pandas fi if [ "a$(julia.version)" != "a" ]; then @@ -55,7 +56,7 @@ steps: export GOPATH=$PWD/src/mlpack/bindings/go go get -u -t gonum.org/v1/gonum/... fi - cmake $(CMakeArgs) .. + cmake $(CMakeArgs) -DPYTHON_EXECUTABLE=`which python` .. displayName: 'CMake' # Build mlpack diff --git a/CMakeLists.txt b/CMakeLists.txt index d6f369bdb2..2b0e9990b4 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -418,6 +418,8 @@ find_package(cereal "${CEREAL_VERSION}" REQUIRED) # Unfortunately this configuration variable is necessary and will need to be # updated as time goes on and new versions are released. set(Boost_ADDITIONAL_VERSIONS + "1.74.0" "1.74" + "17.3.0" "17.3" "1.72.0" "1.72" "1.71.0" "1.71" "1.70.0" "1.70" diff --git a/Doxyfile b/Doxyfile index 23597e8a5f..8c5b7ec5b4 100644 --- a/Doxyfile +++ b/Doxyfile @@ -4,7 +4,7 @@ # Project related configuration options #--------------------------------------------------------------------------- PROJECT_NAME = mlpack -PROJECT_NUMBER = 3.3.2 +PROJECT_NUMBER = 3.4.1 OUTPUT_DIRECTORY = ./doc CREATE_SUBDIRS = NO OUTPUT_LANGUAGE = English diff --git a/HISTORY.md b/HISTORY.md index 104089bd05..4d3807e380 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -1,5 +1,21 @@ ### mlpack ?.?.? ###### ????-??-?? + +### mlpack 3.4.1 +###### 2020-09-07 + * Fix incorrect parsing of required matrix/model parameters for command-line + bindings (#2600). + + * Add manual type specification support to `data::Load()` and `data::Save()` + (#2084, #2135, #2602). + + * Remove use of internal Armadillo functionality (#2596, #2601, #2602). + +### mlpack 3.4.0 +###### 2020-09-01 + + * Issue warnings when metrics produce NaNs in KFoldCV (#2595). + * Added bindings for _R_ during Google Summer of Code (#2556). * Added common striptype function for all bindings (#2556). diff --git a/README.md b/README.md index 62c4bed0c4..99d7fb429d 100644 --- a/README.md +++ b/README.md @@ -23,7 +23,7 @@ src="https://cdn.rawgit.com/mlpack/mlpack.org/e7d36ed8/mlpack-black.svg" style="

Download: - current stable version (3.2.2) + current stable version (3.4.1)

@@ -153,7 +153,7 @@ on Ubuntu, you can install mlpack with the following command: Note: Older Ubuntu versions may not have the most recent version of mlpack available---for instance, at the time of this writing, Ubuntu 16.04 only has -mlpack 3.2.2 available. Options include upgrading your Ubuntu version, finding +mlpack 3.4.1 available. Options include upgrading your Ubuntu version, finding a PPA or other non-official sources, or installing with a manual build. There are some useful pages to consult in addition to this section: diff --git a/doc/examples/sample-ml-app/sample-ml-app/sample-ml-app.vcxproj b/doc/examples/sample-ml-app/sample-ml-app/sample-ml-app.vcxproj index 2c427cfc3a..92c3d1ae0a 100644 --- a/doc/examples/sample-ml-app/sample-ml-app/sample-ml-app.vcxproj +++ b/doc/examples/sample-ml-app/sample-ml-app/sample-ml-app.vcxproj @@ -104,16 +104,16 @@ true _DEBUG;_CONSOLE;%(PreprocessorDefinitions) false - C:\boost\boost_1_66_0;C:\mlpack\armadillo-8.500.1\include;C:\mlpack\mlpack-3.2.2\build\include;%(AdditionalIncludeDirectories) + C:\boost\boost_1_66_0;C:\mlpack\armadillo-8.500.1\include;C:\mlpack\mlpack-3.4.1\build\include;%(AdditionalIncludeDirectories) Console true - C:\mlpack\mlpack-3.2.2\build\Debug\mlpack.lib;C:\boost\boost_1_66_0\lib64-msvc-14.1\libboost_serialization-vc141-mt-gd-x64-1_66.lib;%(AdditionalDependencies) + C:\mlpack\mlpack-3.4.1\build\Debug\mlpack.lib;C:\boost\boost_1_66_0\lib64-msvc-14.1\libboost_serialization-vc141-mt-gd-x64-1_66.lib;%(AdditionalDependencies) - xcopy /y "C:\mlpack\mlpack-3.2.2\build\Debug\mlpack.dll" $(OutDir) -xcopy /y "C:\mlpack\mlpack-3.2.2\packages\OpenBLAS.0.2.14.1\lib\native\bin\x64\*.dll" $(OutDir) + xcopy /y "C:\mlpack\mlpack-3.4.1\build\Debug\mlpack.dll" $(OutDir) +xcopy /y "C:\mlpack\mlpack-3.4.1\packages\OpenBLAS.0.2.14.1\lib\native\bin\x64\*.dll" $(OutDir) xcopy /y "$(ProjectDir)..\..\..\..\src\mlpack\tests\data\german.csv" "$(ProjectDir)data\german.csv*" diff --git a/doc/guide/build.hpp b/doc/guide/build.hpp index 1e8f4be536..33f564475b 100644 --- a/doc/guide/build.hpp +++ b/doc/guide/build.hpp @@ -30,7 +30,7 @@ to build mlpack on Windows, see \ref build_windows (alternatively, you can read is based on older versions). You can download the latest mlpack release from here: -mlpack-3.2.2 +mlpack-3.4.1 @section build_simple Simple Linux build instructions @@ -38,9 +38,9 @@ Assuming all dependencies are installed in the system, you can run the commands below directly to build and install mlpack. @code -$ wget https://www.mlpack.org/files/mlpack-3.2.2.tar.gz -$ tar -xvzpf mlpack-3.2.2.tar.gz -$ mkdir mlpack-3.2.2/build && cd mlpack-3.2.2/build +$ wget https://www.mlpack.org/files/mlpack-3.4.1.tar.gz +$ tar -xvzpf mlpack-3.4.1.tar.gz +$ mkdir mlpack-3.4.1/build && cd mlpack-3.4.1/build $ cmake ../ $ make -j4 # The -j is the number of cores you want to use for a build. $ sudo make install @@ -65,8 +65,8 @@ configure mlpack. First we should unpack the mlpack source and create a build directory. @code -$ tar -xvzpf mlpack-3.2.2.tar.gz -$ cd mlpack-3.2.2 +$ tar -xvzpf mlpack-3.4.1.tar.gz +$ cd mlpack-3.4.1 $ mkdir build @endcode diff --git a/doc/guide/python_quickstart.hpp b/doc/guide/python_quickstart.hpp index 981c1449f3..1f442fc39e 100644 --- a/doc/guide/python_quickstart.hpp +++ b/doc/guide/python_quickstart.hpp @@ -32,9 +32,9 @@ build and install mlpack. You can copy-paste the commands into your shell. @code{.sh} sudo apt-get install libboost-all-dev g++ cmake libarmadillo-dev python-pip wget sudo pip install cython setuptools distutils numpy pandas -wget https://www.mlpack.org/files/mlpack-3.2.2.tar.gz -tar -xvzpf mlpack-3.2.2.tar.gz -mkdir -p mlpack-3.2.2/build/ && cd mlpack-3.2.2/build/ +wget https://www.mlpack.org/files/mlpack-3.4.1.tar.gz +tar -xvzpf mlpack-3.4.1.tar.gz +mkdir -p mlpack-3.4.1/build/ && cd mlpack-3.4.1/build/ cmake ../ && make -j4 && sudo make install @endcode diff --git a/doc/guide/sample_ml_app.hpp b/doc/guide/sample_ml_app.hpp index 6a1b1aa30a..d2a350c85d 100644 --- a/doc/guide/sample_ml_app.hpp +++ b/doc/guide/sample_ml_app.hpp @@ -29,17 +29,17 @@ mlpack and dependencies in Release Mode). @code - C:\boost\boost_1_71_0\lib\native\include - C:\mlpack\armadillo-9.800.3\include - - C:\mlpack\mlpack-3.2.2\build\include + - C:\mlpack\mlpack-3.4.1\build\include @endcode - Under Linker > Input > Additional Dependencies add: @code - - C:\mlpack\mlpack-3.2.2\build\Debug\mlpack.lib + - C:\mlpack\mlpack-3.4.1\build\Debug\mlpack.lib - C:\boost\boost_1_71_0\lib64-msvc-14.2\libboost_serialization-vc142-mt-gd-x64-1_71.lib @endcode - Under Build Events > Post-Build Event > Command Line add: @code - - xcopy /y "C:\mlpack\mlpack-3.2.2\build\Debug\mlpack.dll" $(OutDir) - - xcopy /y "C:\mlpack\mlpack-3.2.2\packages\OpenBLAS.0.2.14.1\lib\native\bin\x64\*.dll" $(OutDir) + - xcopy /y "C:\mlpack\mlpack-3.4.1\build\Debug\mlpack.dll" $(OutDir) + - xcopy /y "C:\mlpack\mlpack-3.4.1\packages\OpenBLAS.0.2.14.1\lib\native\bin\x64\*.dll" $(OutDir) @endcode @note Recent versions of Visual Studio set "Conformance Mode" enabled by default. This causes some issues with diff --git a/src/mlpack/CMakeLists.txt b/src/mlpack/CMakeLists.txt index b284d60dfc..12ff343a34 100644 --- a/src/mlpack/CMakeLists.txt +++ b/src/mlpack/CMakeLists.txt @@ -54,7 +54,7 @@ target_link_libraries(mlpack ${MLPACK_LIBRARIES}) set_target_properties(mlpack PROPERTIES - VERSION 3.3 + VERSION 3.4 SOVERSION 3 ) diff --git a/src/mlpack/bindings/cli/parse_command_line.hpp b/src/mlpack/bindings/cli/parse_command_line.hpp index a611e178fb..381d749bc1 100644 --- a/src/mlpack/bindings/cli/parse_command_line.hpp +++ b/src/mlpack/bindings/cli/parse_command_line.hpp @@ -47,8 +47,8 @@ void ParseCommandLine(int argc, char** argv) { // Add the parameter to desc. util::ParamData& d = it->second; - IO::GetSingleton().functionMap[d.tname]["AddToCLI11"] - (d, NULL, (void*) &app); + IO::GetSingleton().functionMap[d.tname]["AddToCLI11"](d, NULL, (void*) + &app); } // Mark that we did parsing. @@ -136,13 +136,15 @@ void ParseCommandLine(int argc, char** argv) util::ParamData d = iter->second; if (d.required) { - const std::string cliName; + // CLI11 expects the parameter name to have "--" prepended. + std::string cliName; IO::GetSingleton().functionMap[d.tname]["MapParameterName"](d, NULL, (void*) &cliName); + cliName = "--" + cliName; if (!app.count(cliName)) { - Log::Fatal << "Required option --" << cliName << " is undefined." + Log::Fatal << "Required option " << cliName << " is undefined." << std::endl; } } diff --git a/src/mlpack/core/cv/k_fold_cv_impl.hpp b/src/mlpack/core/cv/k_fold_cv_impl.hpp index 46b0a9bf89..8dd9672053 100644 --- a/src/mlpack/core/cv/k_fold_cv_impl.hpp +++ b/src/mlpack/core/cv/k_fold_cv_impl.hpp @@ -246,17 +246,32 @@ double KFoldCV +#include +#include + +namespace mlpack { +namespace data { + +/** + * Given a file type, return a logical name corresponding to that file type. + * + * @param type Type to get the logical name of. + */ +std::string GetStringType(const arma::file_type& type) +{ + switch (type) + { + case arma::csv_ascii: return "CSV data"; + case arma::raw_ascii: return "raw ASCII formatted data"; + case arma::raw_binary: return "raw binary formatted data"; + case arma::arma_ascii: return "Armadillo ASCII formatted data"; + case arma::arma_binary: return "Armadillo binary formatted data"; + case arma::pgm_binary: return "PGM data"; + case arma::hdf5_binary: return "HDF5 data"; + default: return ""; + } +} + +/** + * Given an istream, attempt to guess the file type. This is taken originally + * from Armadillo's function guess_file_type_internal(), but we avoid using + * internal Armadillo functionality. + * + * @param f Opened istream to look into to guess the file type. + */ +arma::file_type GuessFileType(std::istream& f) +{ + f.clear(); + const std::fstream::pos_type pos1 = f.tellg(); + + f.clear(); + f.seekg(0, std::ios::end); + + f.clear(); + // Get the length of the stream. + const std::fstream::pos_type pos2 = f.tellg(); + + // Compute length of the stream. + const arma::uword nMax = ((pos1 >= 0) && (pos2 >= 0) && (pos2 > pos1)) ? + arma::uword(pos2 - pos1) : arma::uword(0); + + f.clear(); + f.seekg(pos1); + + // Handle empty files. + if (nMax == 0) + return arma::file_type_unknown; + + const arma::uword nUse = std::min(nMax, arma::uword(4096)); + + unsigned char* dataMem = new unsigned char[nUse]; + memset(dataMem, 0, nUse); + + f.clear(); + f.read(reinterpret_cast(dataMem), std::streamsize(nUse)); + + const bool loadOkay = f.good(); + + f.clear(); + f.seekg(pos1); + + if (!loadOkay) + { + delete[] dataMem; + return arma::file_type_unknown; + } + + bool hasBinary = false; + bool hasBracket = false; + bool hasComma = false; + + for (arma::uword i = 0; i < nUse; ++i) + { + const unsigned char val = dataMem[i]; + if ((val <= 8) || (val >= 123)) + { + hasBinary = true; + break; + } // The range checking can be made more elaborate. + + if ((val == '(') || (val == ')')) + { + hasBracket = true; + } + if (val == ',') + { + hasComma = true; + } + } + + delete[] dataMem; + + if (hasBinary) + return arma::raw_binary; + + if (hasComma && (hasBracket == false)) + return arma::csv_ascii; + + return arma::raw_ascii; +} + +/** + * Attempt to auto-detect the type of a file given its extension, and by + * inspecting the parts of the file to disambiguate between types when + * necessary. (For instance, a .csv file could be delimited by spaces, commas, + * or tabs.) This is meant to be used during loading. + * + * @param stream Opened file stream to look into for autodetection. + * @param filename Name of the file. + * @return The detected file type. + */ +arma::file_type AutoDetect(std::fstream& stream, + const std::string& filename) +{ + // Get the extension. + std::string extension = Extension(filename); + arma::file_type detectedLoadType = arma::file_type_unknown; + + if (extension == "csv" || extension == "tsv") + { + detectedLoadType = GuessFileType(stream); + if (detectedLoadType == arma::csv_ascii) + { + if (extension == "tsv") + Log::Warn << "'" << filename << "' is comma-separated, not " + "tab-separated!" << std::endl; + } + else if (detectedLoadType == arma::raw_ascii) // .csv file can be tsv. + { + if (extension == "csv") + { + // We should issue a warning, but we don't want to issue the warning if + // there is only one column in the CSV (since there will be no commas + // anyway, and it will be detected as arma::raw_ascii). + const std::streampos pos = stream.tellg(); + std::string line; + std::getline(stream, line, '\n'); + boost::trim(line); + + // Reset stream position. + stream.seekg(pos); + + // If there are no spaces or whitespace in the line, then we shouldn't + // print the warning. + if ((line.find(' ') != std::string::npos) || + (line.find('\t') != std::string::npos)) + { + Log::Warn << "'" << filename << "' is not a standard csv file." + << std::endl; + } + } + } + else + { + detectedLoadType = arma::file_type_unknown; + } + } + else if (extension == "txt") + { + // This could be raw ASCII or Armadillo ASCII (ASCII with size header). + // We'll let Armadillo do its guessing (although we have to check if it is + // arma_ascii ourselves) and see what we come up with. + + // This is adapted from load_auto_detect() in diskio_meat.hpp. + const std::string ARMA_MAT_TXT = "ARMA_MAT_TXT"; + std::string rawHeader(ARMA_MAT_TXT.length(), '\0'); + std::streampos pos = stream.tellg(); + + stream.read(&rawHeader[0], std::streamsize(ARMA_MAT_TXT.length())); + stream.clear(); + stream.seekg(pos); // Reset stream position after peeking. + + if (rawHeader == ARMA_MAT_TXT) + { + detectedLoadType = arma::arma_ascii; + } + else // It's not arma_ascii. Now we let Armadillo guess. + { + detectedLoadType = GuessFileType(stream); + + if (detectedLoadType != arma::raw_ascii && + detectedLoadType != arma::csv_ascii) + detectedLoadType = arma::file_type_unknown; + } + } + else if (extension == "bin") + { + // This could be raw binary or Armadillo binary (binary with header). We + // will check to see if it is Armadillo binary. + const std::string ARMA_MAT_BIN = "ARMA_MAT_BIN"; + std::string rawHeader(ARMA_MAT_BIN.length(), '\0'); + + std::streampos pos = stream.tellg(); + + stream.read(&rawHeader[0], std::streamsize(ARMA_MAT_BIN.length())); + stream.clear(); + stream.seekg(pos); // Reset stream position after peeking. + + if (rawHeader == ARMA_MAT_BIN) + { + detectedLoadType = arma::arma_binary; + } + else // We can only assume it's raw binary. + { + detectedLoadType = arma::raw_binary; + } + } + else if (extension == "pgm") + { + detectedLoadType = arma::pgm_binary; + } + else if (extension == "h5" || extension == "hdf5" || extension == "hdf" || + extension == "he5") + { + detectedLoadType = arma::hdf5_binary; + } + else // Unknown extension... + { + detectedLoadType = arma::file_type_unknown; + } + + return detectedLoadType; +} + +/** + * Return the type based only on the extension. + * + * @param filename Name of the file whose type we should detect. + * @return Detected type of file. + */ +arma::file_type DetectFromExtension(const std::string& filename) +{ + const std::string extension = Extension(filename); + + if (extension == "csv") + { + return arma::csv_ascii; + } + else if (extension == "txt") + { + return arma::raw_ascii; + } + else if (extension == "bin") + { + return arma::arma_binary; + } + else if (extension == "pgm") + { + return arma::pgm_binary; + } + else if (extension == "h5" || extension == "hdf5" || extension == "hdf" || + extension == "he5") + { + return arma::hdf5_binary; + } + else + { + return arma::file_type_unknown; + } +} + +} // namespace data +} // namespace mlpack diff --git a/src/mlpack/core/data/detect_file_type.hpp b/src/mlpack/core/data/detect_file_type.hpp new file mode 100644 index 0000000000..ab387ad0ba --- /dev/null +++ b/src/mlpack/core/data/detect_file_type.hpp @@ -0,0 +1,61 @@ +/** + * @file core/data/detect_file_type.hpp + * @author Conrad Sanderson + * @author Ryan Curtin + * + * Functionality to guess the type of a file by inspecting it. Parts of the + * implementation are adapted from the Armadillo sources and relicensed to be a + * part of mlpack with permission from Conrad. + * + * mlpack is free software; you may redistribute it and/or modify it under the + * terms of the 3-clause BSD license. You should have received a copy of the + * 3-clause BSD license along with mlpack. If not, see + * http://www.opensource.org/licenses/BSD-3-Clause for more information. + */ +#ifndef MLPACK_CORE_DATA_DETECT_FILE_TYPE_HPP +#define MLPACK_CORE_DATA_DETECT_FILE_TYPE_HPP + +namespace mlpack { +namespace data { + +/** + * Given a file type, return a logical name corresponding to that file type. + * + * @param type Type to get the logical name of. + */ +std::string GetStringType(const arma::file_type& type); + +/** + * Given an istream, attempt to guess the file type. This is taken originally + * from Armadillo's function guess_file_type_internal(), but we avoid using + * internal Armadillo functionality. + * + * @param f Opened istream to look into to guess the file type. + */ +arma::file_type GuessFileType(std::istream& f); + +/** + * Attempt to auto-detect the type of a file given its extension, and by + * inspecting the parts of the file to disambiguate between types when + * necessary. (For instance, a .csv file could be delimited by spaces, commas, + * or tabs.) This is meant to be used during loading. + * + * @param stream Opened file stream to look into for autodetection. + * @param filename Name of the file. + * @return The detected file type. arma::file_type_unknown if unknown. + */ +arma::file_type AutoDetect(std::fstream& stream, + const std::string& filename); + +/** + * Return the type based only on the extension. + * + * @param filename Name of the file whose type we should detect. + * @return Detected type of file. arma::file_type_unknown if unknown. + */ +arma::file_type DetectFromExtension(const std::string& filename); + +} // namespace data +} // namespace mlpack + +#endif diff --git a/src/mlpack/core/data/load.cpp b/src/mlpack/core/data/load.cpp index 8532dbed77..71ad35cb64 100644 --- a/src/mlpack/core/data/load.cpp +++ b/src/mlpack/core/data/load.cpp @@ -18,32 +18,38 @@ namespace data /** Functions to load and save matrices and models. */ { template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); template bool Load(const std::string&, arma::SpMat&, diff --git a/src/mlpack/core/data/load.hpp b/src/mlpack/core/data/load.hpp index 74c7fa4f34..d6681d5dce 100644 --- a/src/mlpack/core/data/load.hpp +++ b/src/mlpack/core/data/load.hpp @@ -28,23 +28,25 @@ namespace data /** Functions to load and save matrices and models. */ { /** * Loads a matrix from file, guessing the filetype from the extension. This * will transpose the matrix at load time (unless the transpose parameter is set - * to false). If the filetype cannot be determined, an error will be given. + * to false). * * The supported types of files are the same as found in Armadillo: * - * - CSV (csv_ascii), denoted by .csv, or optionally .txt - * - TSV (raw_ascii), denoted by .tsv, .csv, or .txt - * - ASCII (raw_ascii), denoted by .json - * - Armadillo ASCII (arma_ascii), also denoted by .txt - * - PGM (pgm_binary), denoted by .pgm - * - PPM (ppm_binary), denoted by .ppm - * - Raw binary (raw_binary), denoted by .bin - * - Armadillo binary (arma_binary), denoted by .bin - * - HDF5, denoted by .hdf, .hdf5, .h5, or .he5 + * - CSV (arma::csv_ascii), denoted by .csv, or optionally .txt + * - TSV (arma::raw_ascii), denoted by .tsv, .csv, or .txt + * - ASCII (arma::raw_ascii), denoted by .json + * - Armadillo ASCII (arma::arma_ascii), also denoted by .txt + * - PGM (arma::pgm_binary), denoted by .pgm + * - PPM (arma::ppm_binary), denoted by .ppm + * - Raw binary (arma::raw_binary), denoted by .bin + * - Armadillo binary (arma::arma_binary), denoted by .bin + * - HDF5 (arma::hdf5_binary), denoted by .hdf, .hdf5, .h5, or .he5 * - * If the file extension is not one of those types, an error will be given. - * This is preferable to Armadillo's default behavior of loading an unknown - * filetype as raw_binary, which can have very confusing effects. + * By default, this function will try to automatically determine the type of + * file to load based on its extension and by inspecting the file. If you know + * the file type and want to specify it manually, override the default + * `inputLoadType` parameter with the correct type above (e.g. + * `arma::csv_ascii`.) * * If the parameter 'fatal' is set to true, a std::runtime_error exception will * be thrown if the matrix does not load successfully. The parameter @@ -57,13 +59,15 @@ namespace data /** Functions to load and save matrices and models. */ { * @param matrix Matrix to load contents of file into. * @param fatal If an error should be reported as fatal (default false). * @param transpose If true, transpose the matrix after loading (default true). + * @param inputLoadType Used to determine the type of file to load (default arma::auto_detect). * @return Boolean value indicating success or failure of load. */ template bool Load(const std::string& filename, arma::Mat& matrix, const bool fatal = false, - const bool transpose = true); + const bool transpose = true, + const arma::file_type inputLoadType = arma::auto_detect); /** * Loads a sparse matrix from file, using arma::coord_ascii format. This @@ -110,38 +114,45 @@ bool Load(const std::string& filename, extern template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); // size_t and uword should be one of these three typedefs. extern template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); extern template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); extern template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); extern template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); extern template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); extern template bool Load(const std::string&, arma::Mat&, const bool, - const bool); + const bool, + const arma::file_type); extern template bool Load(const std::string&, arma::SpMat&, diff --git a/src/mlpack/core/data/load_impl.hpp b/src/mlpack/core/data/load_impl.hpp index ca15655e14..344a8b6c13 100644 --- a/src/mlpack/core/data/load_impl.hpp +++ b/src/mlpack/core/data/load_impl.hpp @@ -21,6 +21,7 @@ #include "load_csv.hpp" #include "load.hpp" #include "extension.hpp" +#include "detect_file_type.hpp" #include #include @@ -88,13 +89,11 @@ template bool Load(const std::string& filename, arma::Mat& matrix, const bool fatal, - const bool transpose) + const bool transpose, + const arma::file_type inputLoadType) { Timer::Start("loading_data"); - // Get the extension. - std::string extension = Extension(filename); - // Catch nonexistent files by opening the stream ourselves. std::fstream stream; #ifdef _WIN32 // Always open in binary mode on Windows. @@ -114,129 +113,33 @@ bool Load(const std::string& filename, return false; } - bool unknownType = false; - arma::file_type loadType; + arma::file_type loadType = inputLoadType; std::string stringType; - - if (extension == "csv" || extension == "tsv") + if (inputLoadType == arma::auto_detect) { -#if (ARMA_VERSION_MAJOR == 9 && ARMA_VERSION_MINOR >= 800) - loadType = arma::diskio::guess_file_type_internal(stream); -#else - loadType = arma::diskio::guess_file_type(stream); -#endif - if (loadType == arma::csv_ascii) + // Attempt to auto-detect the type from the given file. + loadType = AutoDetect(stream, filename); + // Provide error if we don't know the type. + if (loadType == arma::file_type_unknown) { - if (extension == "tsv") - Log::Warn << "'" << filename << "' is comma-separated, not " - "tab-separated!" << std::endl; - stringType = "CSV data"; - } - else if (loadType == arma::raw_ascii) // .csv file can be tsv. - { - if (extension == "csv") - { - // We should issue a warning, but we don't want to issue the warning if - // there is only one column in the CSV (since there will be no commas - // anyway, and it will be detected as arma::raw_ascii). - const std::streampos pos = stream.tellg(); - std::string line; - std::getline(stream, line, '\n'); - boost::trim(line); + Timer::Stop("loading_data"); + if (fatal) + Log::Fatal << "Unable to detect type of '" << filename << "'; " + << "incorrect extension?" << std::endl; + else + Log::Warn << "Unable to detect type of '" << filename << "'; load " + << " failed. Incorrect extension?" << std::endl; - // Reset stream position. - stream.seekg(pos); - - // If there are no spaces or whitespace in the line, then we shouldn't - // print the warning. - if ((line.find(' ') != std::string::npos) || - (line.find('\t') != std::string::npos)) - { - Log::Warn << "'" << filename << "' is not a standard csv file." - << std::endl; - } - } - stringType = "raw ASCII formatted data"; - } - else - { - unknownType = true; - loadType = arma::raw_binary; // Won't be used; prevent a warning. - stringType = ""; + return false; } } - else if (extension == "txt") + + stringType = GetStringType(loadType); + +#ifndef ARMA_USE_HDF5 + if (inputLoadType == arma::hdf5_binary) { - // This could be raw ASCII or Armadillo ASCII (ASCII with size header). - // We'll let Armadillo do its guessing (although we have to check if it is - // arma_ascii ourselves) and see what we come up with. - - // This is taken from load_auto_detect() in diskio_meat.hpp - const std::string ARMA_MAT_TXT = "ARMA_MAT_TXT"; - // char* rawHeader = new char[ARMA_MAT_TXT.length() + 1]; - std::string rawHeader(ARMA_MAT_TXT.length(), '\0'); - std::streampos pos = stream.tellg(); - - stream.read(&rawHeader[0], std::streamsize(ARMA_MAT_TXT.length())); - stream.clear(); - stream.seekg(pos); // Reset stream position after peeking. - - if (rawHeader == ARMA_MAT_TXT) - { - loadType = arma::arma_ascii; - stringType = "Armadillo ASCII formatted data"; - } - else // It's not arma_ascii. Now we let Armadillo guess. - { -#if (ARMA_VERSION_MAJOR == 9 && ARMA_VERSION_MINOR >= 800) - loadType = arma::diskio::guess_file_type_internal(stream); -#else - loadType = arma::diskio::guess_file_type(stream); -#endif - if (loadType == arma::raw_ascii) // Raw ASCII (space-separated). - stringType = "raw ASCII formatted data"; - else if (loadType == arma::csv_ascii) // CSV can be .txt too. - stringType = "CSV data"; - else // Unknown .txt... we will throw an error. - unknownType = true; - } - } - else if (extension == "bin") - { - // This could be raw binary or Armadillo binary (binary with header). We - // will check to see if it is Armadillo binary. - const std::string ARMA_MAT_BIN = "ARMA_MAT_BIN"; - std::string rawHeader(ARMA_MAT_BIN.length(), '\0'); - - std::streampos pos = stream.tellg(); - - stream.read(&rawHeader[0], std::streamsize(ARMA_MAT_BIN.length())); - stream.clear(); - stream.seekg(pos); // Reset stream position after peeking. - - if (rawHeader == ARMA_MAT_BIN) - { - stringType = "Armadillo binary formatted data"; - loadType = arma::arma_binary; - } - else // We can only assume it's raw binary. - { - stringType = "raw binary formatted data"; - loadType = arma::raw_binary; - } - } - else if (extension == "pgm") - { - loadType = arma::pgm_binary; - stringType = "PGM data"; - } - else if (extension == "h5" || extension == "hdf5" || extension == "hdf" || - extension == "he5") - { -#ifdef ARMA_USE_HDF5 - loadType = arma::hdf5_binary; - stringType = "HDF5 data"; -#else + // Ensure that HDF5 is supported. Timer::Stop("loading_data"); if (fatal) Log::Fatal << "Attempted to load '" << filename << "' as HDF5 data, but " @@ -248,28 +151,8 @@ bool Load(const std::string& filename, << std::endl; return false; + } #endif - } - else // Unknown extension... - { - unknownType = true; - loadType = arma::raw_binary; // Won't be used; prevent a warning. - stringType = ""; - } - - // Provide error if we don't know the type. - if (unknownType) - { - Timer::Stop("loading_data"); - if (fatal) - Log::Fatal << "Unable to detect type of '" << filename << "'; " - << "incorrect extension?" << std::endl; - else - Log::Warn << "Unable to detect type of '" << filename << "'; load failed." - << " Incorrect extension?" << std::endl; - - return false; - } // Try to load the file; but if it's raw_binary, it could be a problem. if (loadType == arma::raw_binary) diff --git a/src/mlpack/core/data/save.hpp b/src/mlpack/core/data/save.hpp index bcfc83c558..6f8889cf8e 100644 --- a/src/mlpack/core/data/save.hpp +++ b/src/mlpack/core/data/save.hpp @@ -31,33 +31,40 @@ namespace data /** Functions to load and save matrices. */ { * * The supported types of files are the same as found in Armadillo: * - * - CSV (csv_ascii), denoted by .csv, or optionally .txt - * - ASCII (raw_ascii), denoted by .json - * - Armadillo ASCII (arma_ascii), also denoted by .txt - * - PGM (pgm_binary), denoted by .pgm - * - PPM (ppm_binary), denoted by .ppm - * - Raw binary (raw_binary), denoted by .bin - * - Armadillo binary (arma_binary), denoted by .bin - * - HDF5 (hdf5_binary), denoted by .hdf5, .hdf, .h5, or .he5 + * - CSV (arma::csv_ascii), denoted by .csv, or optionally .txt + * - ASCII (arma::raw_ascii), denoted by .json + * - Armadillo ASCII (arma::arma_ascii), also denoted by .txt + * - PGM (arma::pgm_binary), denoted by .pgm + * - PPM (arma::ppm_binary), denoted by .ppm + * - Raw binary (arma::raw_binary), denoted by .bin + * - Armadillo binary (arma::arma_binary), denoted by .bin + * - HDF5 (arma::hdf5_binary), denoted by .hdf5, .hdf, .h5, or .he5 * - * If the file extension is not one of those types, an error will be given. If - * the 'fatal' parameter is set to true, a std::runtime_error exception will be - * thrown upon failure. If the 'transpose' parameter is set to true, the matrix - * will be transposed before saving. Generally, because mlpack stores matrices - * in a column-major format and most datasets are stored on disk as row-major, - * this parameter should be left at its default value of 'true'. + * By default, this function will try to automatically determine the format to + * save with based only on the filename's extension. If you would prefer to + * specify a file type manually, override the default + * `inputSaveType` parameter with the correct type above (e.g. + * `arma::csv_ascii`.) + * + * If the 'fatal' parameter is set to true, a std::runtime_error exception will + * be thrown upon failure. If the 'transpose' parameter is set to true, the + * matrix will be transposed before saving. Generally, because mlpack stores + * matrices in a column-major format and most datasets are stored on disk as + * row-major, this parameter should be left at its default value of 'true'. * * @param filename Name of file to save to. * @param matrix Matrix to save into file. * @param fatal If an error should be reported as fatal (default false). * @param transpose If true, transpose the matrix before saving (default true). + * @param inputSaveType File type to save to (defaults to arma::auto_detect). * @return Boolean value indicating success or failure of save. */ template bool Save(const std::string& filename, const arma::Mat& matrix, const bool fatal = false, - bool transpose = true); + bool transpose = true, + arma::file_type inputSaveType = arma::auto_detect); /** * Saves a sparse matrix to file, guessing the filetype from the diff --git a/src/mlpack/core/data/save_impl.hpp b/src/mlpack/core/data/save_impl.hpp index 2fc936aa61..644e1a005d 100644 --- a/src/mlpack/core/data/save_impl.hpp +++ b/src/mlpack/core/data/save_impl.hpp @@ -15,6 +15,7 @@ // In case it hasn't already been included. #include "save.hpp" #include "extension.hpp" +#include "detect_file_type.hpp" #include #include @@ -26,43 +27,53 @@ namespace data { template bool Save(const std::string& filename, const arma::Col& vec, - const bool fatal) + const bool fatal, + arma::file_type inputSaveType) { // Don't transpose: one observation per line (for CSVs at least). - return Save(filename, vec, fatal, false); + return Save(filename, vec, fatal, false, inputSaveType); } template bool Save(const std::string& filename, const arma::Row& rowvec, - const bool fatal) + const bool fatal, + arma::file_type inputSaveType) { - return Save(filename, rowvec, fatal, true); + return Save(filename, rowvec, fatal, true, inputSaveType); } template bool Save(const std::string& filename, const arma::Mat& matrix, const bool fatal, - bool transpose) + bool transpose, + arma::file_type inputSaveType) { Timer::Start("saving_data"); - // First we will try to discriminate by file extension. - std::string extension = Extension(filename); - if (extension == "") - { - Timer::Stop("saving_data"); - if (fatal) - Log::Fatal << "No extension given with filename '" << filename << "'; " - << "type unknown. Save failed." << std::endl; - else - Log::Warn << "No extension given with filename '" << filename << "'; " - << "type unknown. Save failed." << std::endl; + arma::file_type saveType = inputSaveType; + std::string stringType = ""; - return false; + if (inputSaveType == arma::auto_detect) + { + // Detect the file type using only the extension. + saveType = DetectFromExtension(filename); + if (saveType == arma::file_type_unknown) + { + if (fatal) + Log::Fatal << "Could not detect type of file '" << filename << "' for " + << "writing. Save failed." << std::endl; + else + Log::Warn << "Could not detect type of file '" << filename << "' for " + << "writing. Save failed." << std::endl; + + return false; + } } + stringType = GetStringType(saveType); + // Catch errors opening the file. std::fstream stream; #ifdef _WIN32 // Always open in binary mode on Windows. diff --git a/src/mlpack/core/util/version.hpp b/src/mlpack/core/util/version.hpp index 09e1042691..2ba6d122eb 100644 --- a/src/mlpack/core/util/version.hpp +++ b/src/mlpack/core/util/version.hpp @@ -17,8 +17,8 @@ // The version of mlpack. If this is a git repository, this will be a version // with higher number than the most recent release. #define MLPACK_VERSION_MAJOR 3 -#define MLPACK_VERSION_MINOR 2 -#define MLPACK_VERSION_PATCH 3 +#define MLPACK_VERSION_MINOR 4 +#define MLPACK_VERSION_PATCH 2 // The name of the version (for use by --version). namespace mlpack { diff --git a/src/mlpack/methods/bayesian_linear_regression/bayesian_linear_regression.hpp b/src/mlpack/methods/bayesian_linear_regression/bayesian_linear_regression.hpp index 9bb70ad5e7..4486145c37 100644 --- a/src/mlpack/methods/bayesian_linear_regression/bayesian_linear_regression.hpp +++ b/src/mlpack/methods/bayesian_linear_regression/bayesian_linear_regression.hpp @@ -5,6 +5,11 @@ * Definition of the BayesianRidge class, which performs the * bayesian linear regression. According to the armadillo standards, * all the functions consider data in column-major format. + * + * mlpack is free software; you may redistribute it and/or modify it under the + * terms of the 3-clause BSD license. You should have received a copy of the + * 3-clause BSD license along with mlpack. If not, see + * http://www.opensource.org/licenses/BSD-3-Clause for more information. **/ #ifndef MLPACK_METHODS_BAYESIAN_LINEAR_REGRESSION_HPP diff --git a/src/mlpack/tests/CMakeLists.txt b/src/mlpack/tests/CMakeLists.txt index 4cdd426a70..1de96f845f 100644 --- a/src/mlpack/tests/CMakeLists.txt +++ b/src/mlpack/tests/CMakeLists.txt @@ -16,7 +16,6 @@ add_executable(mlpack_test emst_test.cpp fastmks_test.cpp facilities_test.cpp - feedforward_network_test.cpp gan_test.cpp gmm_test.cpp hmm_test.cpp @@ -25,7 +24,6 @@ add_executable(mlpack_test hyperplane_test.cpp init_rules_test.cpp kde_test.cpp - kmeans_test.cpp krann_search_test.cpp ksinit_test.cpp lars_test.cpp @@ -59,9 +57,7 @@ add_executable(mlpack_test random_forest_test.cpp random_test.cpp range_search_test.cpp - rbm_network_test.cpp rectangle_tree_test.cpp - recurrent_network_test.cpp reward_clipping_test.cpp rl_components_test.cpp serialization.cpp @@ -97,7 +93,6 @@ add_executable(mlpack_test main_tests/hmm_viterbi_test.cpp main_tests/hoeffding_tree_test.cpp main_tests/kde_test.cpp - main_tests/kmeans_test.cpp main_tests/krann_test.cpp main_tests/linear_svm_test.cpp main_tests/lmnn_test.cpp @@ -135,12 +130,14 @@ add_executable(mlpack_catch_test cv_test.cpp decision_stump_test.cpp decision_tree_test.cpp + feedforward_network_test.cpp image_load_test.cpp imputation_test.cpp kernel_pca_test.cpp kernel_test.cpp kernel_traits_test.cpp kfn_test.cpp + kmeans_test.cpp knn_test.cpp linear_regression_test.cpp load_save_test.cpp @@ -149,6 +146,8 @@ add_executable(mlpack_catch_test one_hot_encoding_test.cpp quic_svd_test.cpp randomized_svd_test.cpp + rbm_network_test.cpp + recurrent_network_test.cpp regularized_svd_test.cpp scaling_test.cpp serialization_catch.cpp @@ -169,6 +168,7 @@ add_executable(mlpack_catch_test main_tests/image_converter_test.cpp main_tests/kernel_pca_test.cpp main_tests/kfn_test.cpp + main_tests/kmeans_test.cpp main_tests/knn_test.cpp main_tests/linear_regression_test.cpp main_tests/nca_test.cpp @@ -229,8 +229,6 @@ add_custom_command(TARGET mlpack_test set(parallel_tests "AsyncLearningTest;" "LocalCoordinateCodingTest;" - "FeedForwardNetworkTest;" - "RecurrentNetworkTest;" "GMMTest;" "CFTest;" "HMMTest;" diff --git a/src/mlpack/tests/cv_test.cpp b/src/mlpack/tests/cv_test.cpp index 635afb761f..48bb9a6481 100644 --- a/src/mlpack/tests/cv_test.cpp +++ b/src/mlpack/tests/cv_test.cpp @@ -380,6 +380,26 @@ TEST_CASE("SimpleCVMSETest", "[CVTest]") REQUIRE(std::abs(weightedCV2.Evaluate() - expectedMSE) > 1e-5); } +/** + * Test that scores of -nan are filtered out. + */ +TEST_CASE("FilterNANCVTest", "[CVTest]") +{ + // Create a dataset with only one positive label, so it will not be in every + // fold. + arma::mat data(3, 10, arma::fill::randu); + arma::Row labels(10, arma::fill::zeros); + labels[0] = 1; + + const size_t numClasses = 2; + KFoldCV, F1> kfoldcv(2, data, labels, + numClasses); + + const double result = kfoldcv.Evaluate(); + REQUIRE(!std::isnan(result)); + REQUIRE(!std::isinf(result)); +} + template arma::Row PredictLabelsWithDT(const arma::mat& data, const DTArgs&... args) diff --git a/src/mlpack/tests/feedforward_network_test.cpp b/src/mlpack/tests/feedforward_network_test.cpp index 18c3772cee..c7bf7ce896 100644 --- a/src/mlpack/tests/feedforward_network_test.cpp +++ b/src/mlpack/tests/feedforward_network_test.cpp @@ -19,17 +19,14 @@ #include -#include -#include "test_tools.hpp" -#include "serialization.hpp" +#include "catch.hpp" +#include "serialization_catch.hpp" #include "custom_layer.hpp" using namespace mlpack; using namespace mlpack::ann; using namespace mlpack::kmeans; -BOOST_AUTO_TEST_SUITE(FeedForwardNetworkTest); - /** * Train and evaluate a model with the specified structure. */ @@ -57,13 +54,13 @@ void TestNetwork(ModelType& model, size_t correct = arma::accu(prediction == testLabels); double classificationError = 1 - double(correct) / testData.n_cols; - BOOST_REQUIRE_LE(classificationError, classificationErrorThreshold); + REQUIRE(classificationError <= classificationErrorThreshold); } /** * Train the vanilla network on a larger dataset. */ -BOOST_AUTO_TEST_CASE(VanillaNetworkTest) +TEST_CASE("FFVanillaNetworkTest", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -131,7 +128,7 @@ BOOST_AUTO_TEST_CASE(VanillaNetworkTest) TestNetwork<>(model1, dataset, labels, dataset, labels, 10, 0.2); } -BOOST_AUTO_TEST_CASE(ForwardBackwardTest) +TEST_CASE("ForwardBackwardTest", "[FeedForwardNetworkTest]") { arma::mat dataset; dataset.load("mnist_first250_training_4s_and_9s.arm"); @@ -204,13 +201,13 @@ BOOST_AUTO_TEST_CASE(ForwardBackwardTest) } } - BOOST_REQUIRE(converged); + REQUIRE(converged); } /** * Train the dropout network on a larger dataset. */ -BOOST_AUTO_TEST_CASE(DropoutNetworkTest) +TEST_CASE("DropoutNetworkTest", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -284,7 +281,7 @@ BOOST_AUTO_TEST_CASE(DropoutNetworkTest) /** * Train the highway network on a larger dataset. */ -BOOST_AUTO_TEST_CASE(HighwayNetworkTest) +TEST_CASE("HighwayNetworkTest", "[FeedForwardNetworkTest]") { arma::mat dataset; dataset.load("mnist_first250_training_4s_and_9s.arm"); @@ -311,7 +308,7 @@ BOOST_AUTO_TEST_CASE(HighwayNetworkTest) /** * Train the DropConnect network on a larger dataset. */ -BOOST_AUTO_TEST_CASE(DropConnectNetworkTest) +TEST_CASE("DropConnectNetworkTest", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -385,7 +382,7 @@ BOOST_AUTO_TEST_CASE(DropConnectNetworkTest) * Test miscellaneous things of FFN, * e.g. copy/move constructor, assignment operator. */ -BOOST_AUTO_TEST_CASE(FFNMiscTest) +TEST_CASE("FFNMiscTest", "[FeedForwardNetworkTest]") { FFN> model; model.Add>(2, 3); @@ -400,7 +397,7 @@ BOOST_AUTO_TEST_CASE(FFNMiscTest) /** * Test that serialization works ok. */ -BOOST_AUTO_TEST_CASE(SerializationTest) +TEST_CASE("FFSerializationTest", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -449,7 +446,7 @@ BOOST_AUTO_TEST_CASE(SerializationTest) * Test if the custom layers work. The target is to see if the code compiles * when the Train and Prediction are called. */ -BOOST_AUTO_TEST_CASE(CustomLayerTest) +TEST_CASE("CustomLayerTest", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -481,7 +478,7 @@ BOOST_AUTO_TEST_CASE(CustomLayerTest) /** * Test the overload of Forward function which allows partial forward pass. */ -BOOST_AUTO_TEST_CASE(PartialForwardTest) +TEST_CASE("PartialForwardTest", "[FeedForwardNetworkTest]") { FFN, RandomInitialization> model; model.Add >(5, 10); @@ -528,7 +525,7 @@ BOOST_AUTO_TEST_CASE(PartialForwardTest) /** * Test that FFN::Train() returns finite objective value. */ -BOOST_AUTO_TEST_CASE(FFNTrainReturnObjective) +TEST_CASE("FFNTrainReturnObjective", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -557,13 +554,13 @@ BOOST_AUTO_TEST_CASE(FFNTrainReturnObjective) double objVal = model.Train(trainData, trainLabels, opt); - BOOST_REQUIRE_EQUAL(std::isfinite(objVal), true); + REQUIRE(std::isfinite(objVal) == true); } /** * Test that FFN::Model() allows us to access the instantiated network. */ -BOOST_AUTO_TEST_CASE(FFNReturnModel) +TEST_CASE("FFNReturnModel", "[FeedForwardNetworkTest]") { // Create dummy network. FFN > model; @@ -598,7 +595,7 @@ BOOST_AUTO_TEST_CASE(FFNReturnModel) * Test to see if the FFN code compiles when the Optimizer * doesn't have the MaxIterations() method. */ -BOOST_AUTO_TEST_CASE(OptimizerTest) +TEST_CASE("OptimizerTest", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -626,7 +623,7 @@ BOOST_AUTO_TEST_CASE(OptimizerTest) /** * Train the RBF network on a larger dataset. */ -BOOST_AUTO_TEST_CASE(RBFNetworkTest) +TEST_CASE("RBFNetworkTest", "[FeedForwardNetworkTest]") { // Load the dataset. arma::mat trainData; @@ -703,5 +700,3 @@ BOOST_AUTO_TEST_CASE(RBFNetworkTest) // RBFN neural net with MeanSquaredError. TestNetwork<>(model1, dataset, labels1, dataset, labels, 10, 0.1); } - -BOOST_AUTO_TEST_SUITE_END(); diff --git a/src/mlpack/tests/io_test.cpp b/src/mlpack/tests/io_test.cpp index c650b6d0b5..d136963518 100644 --- a/src/mlpack/tests/io_test.cpp +++ b/src/mlpack/tests/io_test.cpp @@ -598,6 +598,105 @@ BOOST_AUTO_TEST_CASE(InputMatrixParamTest) BOOST_REQUIRE_CLOSE(dataset[i], dataset2[i], 1e-10); } +// Make sure we can correctly load required matrix parameters. +BOOST_AUTO_TEST_CASE(RequiredInputMatrixParamTest) +{ + AddRequiredCLIOptions(); + + // --matrix is an input parameter; it won't be transposed. + PARAM_MATRIX_IN_REQ("matrix", "Test matrix", "m"); + + // Set some fake arguments. + const char* argv[3]; + argv[0] = "./test"; + argv[1] = "--matrix_file"; + argv[2] = "test_data_3_1000.csv"; + + int argc = 3; + + // The const-cast is a little hacky but should be fine... + ParseCommandLine(argc, const_cast(argv)); + + // The --matrix parameter should exist. + BOOST_REQUIRE(IO::HasParam("matrix")); + // The --matrix_file parameter should not exist (it should be transparent from + // inside the program). + Log::Fatal.ignoreInput = true; + BOOST_REQUIRE_THROW(IO::HasParam("matrix_file"), runtime_error); + Log::Fatal.ignoreInput = false; + + arma::mat dataset = IO::GetParam("matrix"); + arma::mat dataset2 = IO::GetParam("matrix"); + + BOOST_REQUIRE_EQUAL(dataset.n_rows, 3); + BOOST_REQUIRE_EQUAL(dataset.n_cols, 1000); + BOOST_REQUIRE_EQUAL(dataset2.n_rows, 3); + BOOST_REQUIRE_EQUAL(dataset2.n_cols, 1000); + + for (size_t i = 0; i < dataset.n_elem; ++i) + BOOST_REQUIRE_CLOSE(dataset[i], dataset2[i], 1e-10); +} + +// Make sure loading required matrix options by alias succeeds. +BOOST_AUTO_TEST_CASE(RequiredInputMatrixParamAliasTest) +{ + AddRequiredCLIOptions(); + + // --matrix is an input parameter; it won't be transposed. + PARAM_MATRIX_IN_REQ("matrix", "Test matrix", "m"); + + // Set some fake arguments. + const char* argv[3]; + argv[0] = "./test"; + argv[1] = "-m"; + argv[2] = "test_data_3_1000.csv"; + + int argc = 3; + + // The const-cast is a little hacky but should be fine... + ParseCommandLine(argc, const_cast(argv)); + + // The --matrix parameter should exist. + BOOST_REQUIRE(IO::HasParam("matrix")); + // The --matrix_file parameter should not exist (it should be transparent from + // inside the program). + Log::Fatal.ignoreInput = true; + BOOST_REQUIRE_THROW(IO::HasParam("matrix_file"), runtime_error); + Log::Fatal.ignoreInput = false; + + arma::mat dataset = IO::GetParam("matrix"); + arma::mat dataset2 = IO::GetParam("matrix"); + + BOOST_REQUIRE_EQUAL(dataset.n_rows, 3); + BOOST_REQUIRE_EQUAL(dataset.n_cols, 1000); + BOOST_REQUIRE_EQUAL(dataset2.n_rows, 3); + BOOST_REQUIRE_EQUAL(dataset2.n_cols, 1000); + + for (size_t i = 0; i < dataset.n_elem; ++i) + BOOST_REQUIRE_CLOSE(dataset[i], dataset2[i], 1e-10); +} + +// Make sure that when we don't pass a required matrix, parsing fails. +BOOST_AUTO_TEST_CASE(RequiredUnspecifiedInputMatrixParamTest) +{ + AddRequiredCLIOptions(); + + // --matrix is an input parameter; it won't be transposed. + PARAM_MATRIX_IN_REQ("matrix", "Test matrix", "m"); + + // Set some fake arguments. + const char* argv[1]; + argv[0] = "./test"; + + int argc = 1; + + // The const-cast is a little hacky but should be fine... + Log::Fatal.ignoreInput = true; + BOOST_REQUIRE_THROW(ParseCommandLine(argc, const_cast(argv)), + std::exception); + Log::Fatal.ignoreInput = false; +} + BOOST_AUTO_TEST_CASE(InputMatrixNoTransposeParamTest) { AddRequiredCLIOptions(); diff --git a/src/mlpack/tests/kmeans_test.cpp b/src/mlpack/tests/kmeans_test.cpp index df879a048f..5d4bae2abf 100644 --- a/src/mlpack/tests/kmeans_test.cpp +++ b/src/mlpack/tests/kmeans_test.cpp @@ -22,9 +22,8 @@ #include #include -#include +#include "catch.hpp" #include -#include "test_tools.hpp" using namespace mlpack; using namespace mlpack::kmeans; @@ -32,8 +31,6 @@ using namespace mlpack::metric; using namespace mlpack::tree; using namespace mlpack::neighbor; -BOOST_AUTO_TEST_SUITE(KMeansTest); - // Generate dataset; written transposed because it's easier to read. arma::mat kMeansData(" 0.0 0.0;" // Class 1. " 0.3 0.4;" @@ -69,7 +66,7 @@ arma::mat kMeansData(" 0.0 0.0;" // Class 1. /** * 30-point 3-class test case for K-Means. */ -BOOST_AUTO_TEST_CASE(KMeansSimpleTest) +TEST_CASE("KMeansSimpleTest", "[KMeansTest]") { // This test was originally written to use RandomPartition, and is left that // way because RandomPartition gives better initializations here. @@ -83,30 +80,30 @@ BOOST_AUTO_TEST_CASE(KMeansSimpleTest) size_t firstClass = assignments(0); for (size_t i = 1; i < 13; ++i) - BOOST_REQUIRE_EQUAL(assignments(i), firstClass); + REQUIRE(assignments(i) == firstClass); size_t secondClass = assignments(13); // To ensure that class 1 != class 2. - BOOST_REQUIRE_NE(firstClass, secondClass); + REQUIRE(firstClass != secondClass); for (size_t i = 13; i < 20; ++i) - BOOST_REQUIRE_EQUAL(assignments(i), secondClass); + REQUIRE(assignments(i) == secondClass); size_t thirdClass = assignments(20); // To ensure that this is the third class which we haven't seen yet. - BOOST_REQUIRE_NE(firstClass, thirdClass); - BOOST_REQUIRE_NE(secondClass, thirdClass); + REQUIRE(firstClass != thirdClass); + REQUIRE(secondClass != thirdClass); for (size_t i = 20; i < 30; ++i) - BOOST_REQUIRE_EQUAL(assignments(i), thirdClass); + REQUIRE(assignments(i) == thirdClass); } /** * Make sure the empty cluster policy class does nothing. */ -BOOST_AUTO_TEST_CASE(AllowEmptyClusterTest) +TEST_CASE("AllowEmptyClusterTest", "[KMeansTest]") { arma::Row assignments; assignments.randu(30); @@ -129,17 +126,17 @@ BOOST_AUTO_TEST_CASE(AllowEmptyClusterTest) // Make sure no assignments were changed. for (size_t i = 0; i < assignments.n_elem; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], assignmentsOld[i]); + REQUIRE(assignments[i] == assignmentsOld[i]); // Make sure no counts were changed. for (size_t i = 0; i < 3; ++i) - BOOST_REQUIRE_EQUAL(counts[i], countsOld[i]); + REQUIRE(counts[i] == countsOld[i]); } /** * Make sure kill empty cluster policy removes the empty cluster. */ -BOOST_AUTO_TEST_CASE(KillEmptyClusterTest) +TEST_CASE("KillEmptyClusterTest", "[KMeansTest]") { arma::Row assignments; assignments.randu(30); @@ -162,20 +159,20 @@ BOOST_AUTO_TEST_CASE(KillEmptyClusterTest) // Make sure no assignments were changed. for (size_t i = 0; i < assignments.n_elem; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], assignmentsOld[i]); + REQUIRE(assignments[i] == assignmentsOld[i]); // Make sure no counts were changed for clusters that are not empty. for (size_t i = 0; i < 2; ++i) - BOOST_REQUIRE_EQUAL(counts[i], countsOld[i]); + REQUIRE(counts[i] == countsOld[i]); // Make sure that counts contain one less element than old counts. - BOOST_REQUIRE_GT(countsOld.n_elem, counts.n_elem); + REQUIRE(countsOld.n_elem > counts.n_elem); } /** * Make sure the max variance method finds the correct point. */ -BOOST_AUTO_TEST_CASE(MaxVarianceNewClusterTest) +TEST_CASE("MaxVarianceNewClusterTest", "[KMeansTest]") { // Five points. arma::mat data("0.4 1.0 5.0 -2.0 -2.5;" @@ -220,22 +217,22 @@ BOOST_AUTO_TEST_CASE(MaxVarianceNewClusterTest) assignments[i] = closestCluster; } - BOOST_REQUIRE_EQUAL(assignments[0], 0); - BOOST_REQUIRE_EQUAL(assignments[1], 0); - BOOST_REQUIRE_EQUAL(assignments[2], 2); - BOOST_REQUIRE_EQUAL(assignments[3], 1); - BOOST_REQUIRE_EQUAL(assignments[4], 1); + REQUIRE(assignments[0] == 0); + REQUIRE(assignments[1] == 0); + REQUIRE(assignments[2] == 2); + REQUIRE(assignments[3] == 1); + REQUIRE(assignments[4] == 1); // Ensure that the counts are right. - BOOST_REQUIRE_EQUAL(counts[0], 2); - BOOST_REQUIRE_EQUAL(counts[1], 2); - BOOST_REQUIRE_EQUAL(counts[2], 1); + REQUIRE(counts[0] == 2); + REQUIRE(counts[1] == 2); + REQUIRE(counts[2] == 1); } /** * Make sure the random partitioner seems to return valid results. */ -BOOST_AUTO_TEST_CASE(RandomPartitionTest) +TEST_CASE("RandomPartitionTest", "[KMeansTest]") { arma::mat data; data.randu(2, 1000); // One thousand points. @@ -246,17 +243,17 @@ BOOST_AUTO_TEST_CASE(RandomPartitionTest) RandomPartition::Cluster(data, 18, assignments); // Ensure that the right number of assignments were given. - BOOST_REQUIRE_EQUAL(assignments.n_elem, 1000); + REQUIRE(assignments.n_elem == 1000); // Ensure that no value is greater than 17 (the maximum valid cluster). for (size_t i = 0; i < 1000; ++i) - BOOST_REQUIRE_LT(assignments[i], 18); + REQUIRE(assignments[i] < 18); } /** * Make sure that random initialization fails for a corner case dataset. */ -BOOST_AUTO_TEST_CASE(RandomInitialAssignmentFailureTest) +TEST_CASE("RandomInitialAssignmentFailureTest", "[KMeansTest]") { // This is a very synthetic dataset. It is one Gaussian with a huge number of // points combined with one faraway Gaussian with very few points. Normally, @@ -292,14 +289,14 @@ BOOST_AUTO_TEST_CASE(RandomInitialAssignmentFailureTest) // Only one success allowed. The probability of two successes should be // infinitesimal. - BOOST_REQUIRE_LT(successes, 2); + REQUIRE(successes < 2); } /** * Make sure that specifying initial assignments is successful for a corner case * dataset which doesn't usually converge otherwise. */ -BOOST_AUTO_TEST_CASE(InitialAssignmentTest) +TEST_CASE("InitialAssignmentTest", "[KMeansTest]") { // For a better description of this dataset, see // RandomInitialAssignmentFailureTest. @@ -321,9 +318,9 @@ BOOST_AUTO_TEST_CASE(InitialAssignmentTest) // Check results. for (size_t i = 0; i < 10000; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 0); + REQUIRE(assignments[i] == 0); for (size_t i = 10000; i < 10002; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 1); + REQUIRE(assignments[i] == 1); // Now, slightly harder. Give it one incorrect assignment in each cluster. // The wrong assignment should be quickly fixed. @@ -334,16 +331,16 @@ BOOST_AUTO_TEST_CASE(InitialAssignmentTest) // Check results. for (size_t i = 0; i < 10000; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 0); + REQUIRE(assignments[i] == 0); for (size_t i = 10000; i < 10002; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 1); + REQUIRE(assignments[i] == 1); } /** * Make sure specifying initial centroids is successful for a corner case which * doesn't usually converge otherwise. */ -BOOST_AUTO_TEST_CASE(InitialCentroidTest) +TEST_CASE("InitialCentroidTest", "[KMeansTest]") { // For a better description of this dataset, see // RandomInitialAssignmentFailureTest. @@ -365,9 +362,9 @@ BOOST_AUTO_TEST_CASE(InitialCentroidTest) // Check results. for (size_t i = 0; i < 10000; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 0); + REQUIRE(assignments[i] == 0); for (size_t i = 10000; i < 10002; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 1); + REQUIRE(assignments[i] == 1); // Now add a little noise to the initial centroids. centroids.col(0) = arma::vec("3 4"); @@ -377,15 +374,15 @@ BOOST_AUTO_TEST_CASE(InitialCentroidTest) // Check results. for (size_t i = 0; i < 10000; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 0); + REQUIRE(assignments[i] == 0); for (size_t i = 10000; i < 10002; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 1); + REQUIRE(assignments[i] == 1); } /** * Ensure that initial assignments override initial centroids. */ -BOOST_AUTO_TEST_CASE(InitialAssignmentOverrideTest) +TEST_CASE("InitialAssignmentOverrideTest", "[KMeansTest]") { // For a better description of this dataset, see // RandomInitialAssignmentFailureTest. @@ -412,22 +409,22 @@ BOOST_AUTO_TEST_CASE(InitialAssignmentOverrideTest) // Because the initial assignments guess should take priority, we should get // those same results back. for (size_t i = 0; i < 10000; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 0); + REQUIRE(assignments[i] == 0); for (size_t i = 10000; i < 10002; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], 1); + REQUIRE(assignments[i] == 1); // Make sure the centroids are about right too. - BOOST_REQUIRE_LT(centroids(0, 0), 10.0); - BOOST_REQUIRE_LT(centroids(1, 0), 10.0); - BOOST_REQUIRE_GT(centroids(0, 1), 40.0); - BOOST_REQUIRE_GT(centroids(1, 1), 40.0); + REQUIRE(centroids(0, 0) < 10.0); + REQUIRE(centroids(1, 0) < 10.0); + REQUIRE(centroids(0, 1) > 40.0); + REQUIRE(centroids(1, 1) > 40.0); } /** * Test that the refined starting policy returns decent initial cluster * estimates. */ -BOOST_AUTO_TEST_CASE(RefinedStartTest) +TEST_CASE("RefinedStartTest", "[KMeansTest]") { // Our dataset will be five Gaussians of largely varying numbers of points and // we expect that the refined starting policy should return good guesses at @@ -486,14 +483,14 @@ BOOST_AUTO_TEST_CASE(RefinedStartTest) // figure is a corner case which actually does not give good clusters), and // random initial starts give distortion around 22000. So we'll require that // our distortion is less than 14000. - BOOST_REQUIRE_LT(distortion, 14000.0); + REQUIRE(distortion < 14000.0); } #ifdef ARMA_HAS_SPMAT /** * Make sure sparse k-means works okay. */ -BOOST_AUTO_TEST_CASE(SparseKMeansTest) +TEST_CASE("SparseKMeansTest", "[KMeansTest]") { // Huge dimensionality, few points. arma::SpMat data(5000, 12); @@ -520,23 +517,23 @@ BOOST_AUTO_TEST_CASE(SparseKMeansTest) size_t clusterOne = assignments[0]; size_t clusterTwo = assignments[6]; - BOOST_REQUIRE_EQUAL(assignments[0], clusterOne); - BOOST_REQUIRE_EQUAL(assignments[1], clusterOne); - BOOST_REQUIRE_EQUAL(assignments[2], clusterOne); - BOOST_REQUIRE_EQUAL(assignments[3], clusterOne); - BOOST_REQUIRE_EQUAL(assignments[4], clusterOne); - BOOST_REQUIRE_EQUAL(assignments[5], clusterOne); - BOOST_REQUIRE_EQUAL(assignments[6], clusterTwo); - BOOST_REQUIRE_EQUAL(assignments[7], clusterTwo); - BOOST_REQUIRE_EQUAL(assignments[8], clusterTwo); - BOOST_REQUIRE_EQUAL(assignments[9], clusterTwo); - BOOST_REQUIRE_EQUAL(assignments[10], clusterTwo); - BOOST_REQUIRE_EQUAL(assignments[11], clusterTwo); + REQUIRE(assignments[0] == clusterOne); + REQUIRE(assignments[1] == clusterOne); + REQUIRE(assignments[2] == clusterOne); + REQUIRE(assignments[3] == clusterOne); + REQUIRE(assignments[4] == clusterOne); + REQUIRE(assignments[5] == clusterOne); + REQUIRE(assignments[6] == clusterTwo); + REQUIRE(assignments[7] == clusterTwo); + REQUIRE(assignments[8] == clusterTwo); + REQUIRE(assignments[9] == clusterTwo); + REQUIRE(assignments[10] == clusterTwo); + REQUIRE(assignments[11] == clusterTwo); } #endif // ARMA_HAS_SPMAT -BOOST_AUTO_TEST_CASE(ElkanTest) +TEST_CASE("ElkanTest", "[KMeansTest]") { const size_t trials = 5; @@ -563,14 +560,14 @@ BOOST_AUTO_TEST_CASE(ElkanTest) elkan.Cluster(dataset, k, elkanAssignments, elkanCentroids, false, true); for (size_t i = 0; i < dataset.n_cols; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], elkanAssignments[i]); + REQUIRE(assignments[i] == elkanAssignments[i]); for (size_t i = 0; i < centroids.n_elem; ++i) - BOOST_REQUIRE_CLOSE(naiveCentroids[i], elkanCentroids[i], 1e-5); + REQUIRE(naiveCentroids[i] == Approx(elkanCentroids[i]).epsilon(1e-7)); } } -BOOST_AUTO_TEST_CASE(HamerlyTest) +TEST_CASE("HamerlyTest", "[KMeansTest]") { const size_t trials = 5; @@ -598,14 +595,14 @@ BOOST_AUTO_TEST_CASE(HamerlyTest) true); for (size_t i = 0; i < dataset.n_cols; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], hamerlyAssignments[i]); + REQUIRE(assignments[i] == hamerlyAssignments[i]); for (size_t i = 0; i < centroids.n_elem; ++i) - BOOST_REQUIRE_CLOSE(naiveCentroids[i], hamerlyCentroids[i], 1e-5); + REQUIRE(naiveCentroids[i] == Approx(hamerlyCentroids[i]).epsilon(1e-7)); } } -BOOST_AUTO_TEST_CASE(PellegMooreTest) +TEST_CASE("PellegMooreTest", "[KMeansTest]") { const size_t trials = 5; @@ -632,14 +629,14 @@ BOOST_AUTO_TEST_CASE(PellegMooreTest) pellegMoore.Cluster(dataset, k, pmAssignments, pmCentroids, false, true); for (size_t i = 0; i < dataset.n_cols; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], pmAssignments[i]); + REQUIRE(assignments[i] == pmAssignments[i]); for (size_t i = 0; i < centroids.n_elem; ++i) - BOOST_REQUIRE_CLOSE(naiveCentroids[i], pmCentroids[i], 1e-5); + REQUIRE(naiveCentroids[i] == Approx(pmCentroids[i]).epsilon(1e-7)); } } -BOOST_AUTO_TEST_CASE(DTNNTest) +TEST_CASE("DTNNTest", "[KMeansTest]") { const size_t trials = 5; @@ -664,14 +661,14 @@ BOOST_AUTO_TEST_CASE(DTNNTest) dtnn.Cluster(dataset, k, dtnnAssignments, dtnnCentroids, false, true); for (size_t i = 0; i < dataset.n_cols; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], dtnnAssignments[i]); + REQUIRE(assignments[i] == dtnnAssignments[i]); for (size_t i = 0; i < centroids.n_elem; ++i) - BOOST_REQUIRE_CLOSE(naiveCentroids[i], dtnnCentroids[i], 1e-5); + REQUIRE(naiveCentroids[i] == Approx(dtnnCentroids[i]).epsilon(1e-7)); } } -BOOST_AUTO_TEST_CASE(DTNNCoverTreeTest) +TEST_CASE("DTNNCoverTreeTest", "[KMeansTest]") { const size_t trials = 5; @@ -696,10 +693,10 @@ BOOST_AUTO_TEST_CASE(DTNNCoverTreeTest) dtnn.Cluster(dataset, k, dtnnAssignments, dtnnCentroids, false, true); for (size_t i = 0; i < dataset.n_cols; ++i) - BOOST_REQUIRE_EQUAL(assignments[i], dtnnAssignments[i]); + REQUIRE(assignments[i] == dtnnAssignments[i]); for (size_t i = 0; i < centroids.n_elem; ++i) - BOOST_REQUIRE_CLOSE(naiveCentroids[i], dtnnCentroids[i], 1e-5); + REQUIRE(naiveCentroids[i] == Approx(dtnnCentroids[i]).epsilon(1e-7)); } } @@ -707,7 +704,7 @@ BOOST_AUTO_TEST_CASE(DTNNCoverTreeTest) * Make sure that the sample initialization strategy successfully samples points * from the dataset. */ -BOOST_AUTO_TEST_CASE(SampleInitializationTest) +TEST_CASE("SampleInitializationTest", "[KMeansTest]") { arma::mat dataset = arma::randu(5, 100); const size_t clusters = 10; @@ -716,8 +713,8 @@ BOOST_AUTO_TEST_CASE(SampleInitializationTest) SampleInitialization::Cluster(dataset, clusters, centroids); // Check that the size of the matrix is correct. - BOOST_REQUIRE_EQUAL(centroids.n_cols, 10); - BOOST_REQUIRE_EQUAL(centroids.n_rows, 5); + REQUIRE(centroids.n_cols == 10); + REQUIRE(centroids.n_rows == 5); // Check that each entry in the matrix is some sample from the dataset. for (size_t i = 0; i < clusters; ++i) @@ -733,8 +730,6 @@ BOOST_AUTO_TEST_CASE(SampleInitializationTest) break; } - BOOST_REQUIRE_LT(j, dataset.n_cols); + REQUIRE(j < dataset.n_cols); } } - -BOOST_AUTO_TEST_SUITE_END(); diff --git a/src/mlpack/tests/load_save_test.cpp b/src/mlpack/tests/load_save_test.cpp index f82284280e..28f3523933 100644 --- a/src/mlpack/tests/load_save_test.cpp +++ b/src/mlpack/tests/load_save_test.cpp @@ -48,6 +48,56 @@ TEST_CASE("NotExistLoad", "[LoadSaveTest]") REQUIRE(data::Load("nonexistentfile_______________.csv", out) == false); } +/** + * Make sure load fails if the file extension is wrong in automatic detection mode. + */ +TEST_CASE("WrongExtensionWrongLoad", "[LoadSaveTest]") +{ + // Try to load arma::arma_binary file with ".csv" extension + arma::mat test = "1 5;" + "2 6;" + "3 7;" + "4 8;"; + + arma::mat testTrans = trans(test); + REQUIRE(testTrans.quiet_save("test_file.csv", arma::arma_binary) == true); + + // Now reload through our interface. + REQUIRE(data::Load("test_file.csv", test) == false); + + // Remove the file. + remove("test_file.csv"); +} + +/** + * Make sure load is successful even if the file extension is wrong when file type is specified. + */ +TEST_CASE("WrongExtensionCorrectLoad", "[LoadSaveTest]") +{ + // Try to load arma::arma_binary file with ".csv" extension + arma::mat test = "1 5;" + "2 6;" + "3 7;" + "4 8;"; + + arma::mat testTrans = trans(test); + REQUIRE(testTrans.quiet_save("test_file.csv", arma::arma_binary) == true); + + // Now reload through our interface. + REQUIRE( + data::Load("test_file.csv", test, false, true, arma::arma_binary) + == true); + + REQUIRE(test.n_rows == 4); + REQUIRE(test.n_cols == 2); + + for (size_t i = 0; i < 8; i++) + REQUIRE(test[i] == Approx((double) (i + 1)).epsilon(1e-3)); + + // Remove the file. + remove("test_file.csv"); +} + /** * Make sure a CSV is loaded correctly. */ @@ -205,6 +255,32 @@ TEST_CASE("LoadTSVExtensionTest", "[LoadSaveTest]") remove("test_file.tsv"); } +/** + * Test that we can manually specify the format for loading. + */ +TEST_CASE("LoadAnyExtensionFileTest", "[LoadSaveTest]") +{ + fstream f; + f.open("test_file.blah", fstream::out); + + f << "1\t2\t3\t4" << endl; + f << "5\t6\t7\t8" << endl; + + f.close(); + + arma::mat test; + REQUIRE(data::Load("test_file.blah", test, false, true, arma::raw_ascii)); + + REQUIRE(test.n_rows == 4); + REQUIRE(test.n_cols == 2); + + for (size_t i = 0; i < 8; ++i) + REQUIRE(test[i] == Approx((double) (i + 1)).epsilon(1e-7)); + + // Remove the file. + remove("test_file.blah"); +} + /** * Make sure a CSV is saved correctly. */ @@ -892,6 +968,32 @@ TEST_CASE("SaveArmaBinaryTest", "[LoadSaveTest]") remove("test_file.bin"); } +/** + * Make sure that we can manually specify the format. + */ +TEST_CASE("SaveArmaBinaryArbitraryExtensionTest", "[LoadSaveTest]") +{ + arma::mat test = "1 5;" + "2 6;" + "3 7;" + "4 8;"; + + REQUIRE(data::Save("test_file.blerp.blah", test, false, true, + arma::arma_binary) == true); + + REQUIRE(data::Load("test_file.blerp.blah", test, false, true, + arma::arma_binary) == true); + + REQUIRE(test.n_rows == 4); + REQUIRE(test.n_cols == 2); + + for (size_t i = 0; i < 8; ++i) + REQUIRE(test[i] == Approx((double) (i + 1)).epsilon(1e-7)); + + // Remove the file. + remove("test_file.blerp.blah"); +} + /** * Make sure raw_binary is loaded correctly. */ diff --git a/src/mlpack/tests/main_tests/kmeans_test.cpp b/src/mlpack/tests/main_tests/kmeans_test.cpp index 9de3bc1e09..6db26d1fcb 100644 --- a/src/mlpack/tests/main_tests/kmeans_test.cpp +++ b/src/mlpack/tests/main_tests/kmeans_test.cpp @@ -19,8 +19,8 @@ static const std::string testName = "Kmeans"; #include "test_helper.hpp" #include -#include -#include "../test_tools.hpp" +#include "../catch.hpp" +#include "../test_catch_tools.hpp" using namespace mlpack; @@ -46,22 +46,21 @@ void ResetKmSettings() IO::RestoreSettings(testName); } -BOOST_FIXTURE_TEST_SUITE(KmeansMainTest, KmTestFixture); - /** * Checking that number of Clusters are non negative */ -BOOST_AUTO_TEST_CASE(NonNegativeClustersTest) +TEST_CASE_METHOD(KmTestFixture, "NonNegativeClustersTest", + "[KmeansMainTest][BindingTests]") { arma::mat inputData; if (!data::Load("vc2.csv", inputData)) - BOOST_FAIL("Unable to load train dataset vc2.csv!"); + FAIL("Unable to load train dataset vc2.csv!"); SetInputParam("input", std::move(inputData)); SetInputParam("clusters", (int) -1); // Invalid Log::Fatal.ignoreInput = true; - BOOST_REQUIRE_THROW(mlpackMain(), std::runtime_error); + REQUIRE_THROWS_AS(mlpackMain(), std::runtime_error); Log::Fatal.ignoreInput = false; } @@ -69,7 +68,8 @@ BOOST_AUTO_TEST_CASE(NonNegativeClustersTest) /** * Checking that initial centroids are provided if clusters are to be auto detected */ -BOOST_AUTO_TEST_CASE(AutoDetectClusterTest) +TEST_CASE_METHOD(KmTestFixture, "AutoDetectClusterTest", + "[KmeansMainTest][BindingTests]") { constexpr int N = 10; constexpr int D = 4; @@ -80,7 +80,7 @@ BOOST_AUTO_TEST_CASE(AutoDetectClusterTest) SetInputParam("clusters", (int) 0); // Invalid Log::Fatal.ignoreInput = true; - BOOST_REQUIRE_THROW(mlpackMain(), std::runtime_error); + REQUIRE_THROWS_AS(mlpackMain(), std::runtime_error); Log::Fatal.ignoreInput = false; } @@ -88,13 +88,14 @@ BOOST_AUTO_TEST_CASE(AutoDetectClusterTest) /** * Checking that percentage is between 0 and 1 when --refined_start is specified */ -BOOST_AUTO_TEST_CASE(RefinedStartPercentageTest) +TEST_CASE_METHOD(KmTestFixture, "RefinedStartPercentageTest", + "[KmeansMainTest][BindingTests]") { int c = 2; double P = 2.0; arma::mat inputData; if (!data::Load("vc2.csv", inputData)) - BOOST_FAIL("Unable to load train dataset vc2.csv!"); + FAIL("Unable to load train dataset vc2.csv!"); SetInputParam("input", std::move(inputData)); SetInputParam("refined_start", true); @@ -102,7 +103,7 @@ BOOST_AUTO_TEST_CASE(RefinedStartPercentageTest) SetInputParam("percentage", std::move(P)); // Invalid Log::Fatal.ignoreInput = true; - BOOST_REQUIRE_THROW(mlpackMain(), std::runtime_error); + REQUIRE_THROWS_AS(mlpackMain(), std::runtime_error); Log::Fatal.ignoreInput = false; } @@ -110,13 +111,14 @@ BOOST_AUTO_TEST_CASE(RefinedStartPercentageTest) /** * Checking percentage is non-negative when --refined_start is specified */ -BOOST_AUTO_TEST_CASE(NonNegativePercentageTest) +TEST_CASE_METHOD(KmTestFixture, "NonNegativePercentageTest", + "[KmeansMainTest][BindingTests]") { int c = 2; double P = -1.0; arma::mat inputData; if (!data::Load("vc2.csv", inputData)) - BOOST_FAIL("Unable to load train dataset vc2.csv!"); + FAIL("Unable to load train dataset vc2.csv!"); SetInputParam("input", std::move(inputData)); SetInputParam("refined_start", true); @@ -124,7 +126,7 @@ BOOST_AUTO_TEST_CASE(NonNegativePercentageTest) SetInputParam("percentage", P); // Invalid Log::Fatal.ignoreInput = true; - BOOST_REQUIRE_THROW(mlpackMain(), std::runtime_error); + REQUIRE_THROWS_AS(mlpackMain(), std::runtime_error); Log::Fatal.ignoreInput = false; } @@ -132,12 +134,13 @@ BOOST_AUTO_TEST_CASE(NonNegativePercentageTest) /** * Checking that size and dimensionality of prediction is correct. */ -BOOST_AUTO_TEST_CASE(KmClusteringSizeCheck) +TEST_CASE_METHOD(KmTestFixture, "KmClusteringSizeCheck", + "[KmeansMainTest][BindingTests]") { int c = 2; arma::mat inputData; if (!data::Load("vc2.csv", inputData)) - BOOST_FAIL("Unable to load train dataset vc2.csv!"); + FAIL("Unable to load train dataset vc2.csv!"); size_t col = inputData.n_cols; size_t row = inputData.n_rows; @@ -147,22 +150,23 @@ BOOST_AUTO_TEST_CASE(KmClusteringSizeCheck) mlpackMain(); - BOOST_REQUIRE_EQUAL(IO::GetParam("output").n_rows, row+1); - BOOST_REQUIRE_EQUAL(IO::GetParam("output").n_cols, col); - BOOST_REQUIRE_EQUAL(IO::GetParam("centroid").n_rows, row); - BOOST_REQUIRE_EQUAL(IO::GetParam("centroid").n_cols, c); + REQUIRE(IO::GetParam("output").n_rows == row+1); + REQUIRE(IO::GetParam("output").n_cols == col); + REQUIRE(IO::GetParam("centroid").n_rows == row); + REQUIRE(IO::GetParam("centroid").n_cols == c); } /** * Checking that size and dimensionality of prediction is correct when --labels_only is specified */ -BOOST_AUTO_TEST_CASE(KmClusteringSizeCheckLabelOnly) +TEST_CASE_METHOD(KmTestFixture, "KmClusteringSizeCheckLabelOnly", + "[KmeansMainTest][BindingTests]") { int c = 2; arma::mat inputData; if (!data::Load("vc2.csv", inputData)) - BOOST_FAIL("Unable to load train dataset vc2.csv!"); + FAIL("Unable to load train dataset vc2.csv!"); size_t col = inputData.n_cols; size_t row = inputData.n_rows; @@ -172,24 +176,25 @@ BOOST_AUTO_TEST_CASE(KmClusteringSizeCheckLabelOnly) mlpackMain(); - BOOST_REQUIRE_EQUAL(IO::GetParam("output").n_rows, 1); - BOOST_REQUIRE_EQUAL(IO::GetParam("output").n_cols, col); - BOOST_REQUIRE_EQUAL(IO::GetParam("centroid").n_rows, row); - BOOST_REQUIRE_EQUAL(IO::GetParam("centroid").n_cols, c); + REQUIRE(IO::GetParam("output").n_rows == 1); + REQUIRE(IO::GetParam("output").n_cols == col); + REQUIRE(IO::GetParam("centroid").n_rows == row); + REQUIRE(IO::GetParam("centroid").n_cols == c); } /** * Checking that predictions are not same when --allow_empty_clusters or kill_empty_clusters are specified */ -BOOST_AUTO_TEST_CASE(KmClusteringEmptyClustersCheck) +TEST_CASE_METHOD(KmTestFixture, "KmClusteringEmptyClustersCheck", + "[KmeansMainTest][BindingTests]") { int c = 400; int iterations = 100; arma::mat inputData; if (!data::Load("test_data_3_1000.csv", inputData)) - BOOST_FAIL("Unable to load train dataset test_data_3_1000.csv!"); + FAIL("Unable to load train dataset test_data_3_1000.csv!"); arma::mat initCentroid = arma::randu(inputData.n_rows, c); SetInputParam("input", inputData); @@ -235,23 +240,24 @@ BOOST_AUTO_TEST_CASE(KmClusteringEmptyClustersCheck) if (killEmptyOutput.n_elem == allowEmptyOutput.n_elem) { - BOOST_REQUIRE_GT(arma::accu(killEmptyOutput != allowEmptyOutput), 1); - BOOST_REQUIRE_GT(arma::accu(killEmptyOutput != normalOutput), 1); + REQUIRE(arma::accu(killEmptyOutput != allowEmptyOutput) > 1); + REQUIRE(arma::accu(killEmptyOutput != normalOutput) > 1); } - BOOST_REQUIRE_GT(arma::accu(normalOutput != allowEmptyOutput), 1); + REQUIRE(arma::accu(normalOutput != allowEmptyOutput) > 1); } /** * Checking that that size and dimensionality of Final Input File is correct * when flag --in_place is specified */ -BOOST_AUTO_TEST_CASE(KmClusteringResultSizeCheck) +TEST_CASE_METHOD(KmTestFixture, "KmClusteringResultSizeCheck", + "[KmeansMainTest][BindingTests]") { int c = 2; arma::mat inputData; if (!data::Load("vc2.csv", inputData)) - BOOST_FAIL("Unable to load train dataset vc2.csv!"); + FAIL("Unable to load train dataset vc2.csv!"); size_t row = inputData.n_rows; size_t col = inputData.n_cols; @@ -265,30 +271,32 @@ BOOST_AUTO_TEST_CASE(KmClusteringResultSizeCheck) // here input is actually accessed through output // due to a little trick in kmeans_main - BOOST_REQUIRE_EQUAL(processedInput.n_cols, col); - BOOST_REQUIRE_EQUAL(processedInput.n_rows, row+1); + REQUIRE(processedInput.n_cols == col); + REQUIRE(processedInput.n_rows == row+1); } /** * Ensuring that absence of Number of Clusters is checked. */ -BOOST_AUTO_TEST_CASE(KmClustersNotDefined) +TEST_CASE_METHOD(KmTestFixture, "KmClustersNotDefined", + "[KmeansMainTest][BindingTests]") { arma::mat inputData; if (!data::Load("vc2.csv", inputData)) - BOOST_FAIL("Unable to load train dataset vc2.csv!"); + FAIL("Unable to load train dataset vc2.csv!"); SetInputParam("input", std::move(inputData)); Log::Fatal.ignoreInput = true; - BOOST_REQUIRE_THROW(mlpackMain(), std::runtime_error); + REQUIRE_THROWS_AS(mlpackMain(), std::runtime_error); Log::Fatal.ignoreInput = false; } /** * Checking that all the algorithms yield same results */ -BOOST_AUTO_TEST_CASE(AlgorithmsSimilarTest) +TEST_CASE_METHOD(KmTestFixture, "AlgorithmsSimilarTest", + "[KmeansMainTest][BindingTests]") { int c = 5; arma::mat inputData(10, 1000); @@ -393,5 +401,3 @@ BOOST_AUTO_TEST_CASE(AlgorithmsSimilarTest) CheckMatrices(naiveCentroid, dualTreeCentroid); CheckMatrices(naiveCentroid, dualCoverTreeCentroid); } - -BOOST_AUTO_TEST_SUITE_END(); diff --git a/src/mlpack/tests/rbm_network_test.cpp b/src/mlpack/tests/rbm_network_test.cpp index 820f9616f5..f14e2726bd 100644 --- a/src/mlpack/tests/rbm_network_test.cpp +++ b/src/mlpack/tests/rbm_network_test.cpp @@ -26,20 +26,17 @@ #include #include -#include -#include "test_tools.hpp" +#include "catch.hpp" using namespace mlpack; using namespace mlpack::ann; using namespace ens; using namespace mlpack::regression; -BOOST_AUTO_TEST_SUITE(RBMNetworkTest); - /* * Tests the BinaryRBM implementation on the Digits dataset. */ -BOOST_AUTO_TEST_CASE(BinaryRBMClassificationTest) +TEST_CASE("BinaryRBMClassificationTest", "[RBMNetworkTest]") { // Normalised dataset. int hiddenLayerSize = 100; @@ -84,7 +81,7 @@ BOOST_AUTO_TEST_CASE(BinaryRBMClassificationTest) double objVal = model.Train(msgd); // Test that objective value returned by RBM::Train() is finite. - BOOST_REQUIRE_EQUAL(std::isfinite(objVal), true); + REQUIRE(std::isfinite(objVal) == true); for (size_t i = 0; i < trainData.n_cols; ++i) { @@ -117,13 +114,13 @@ BOOST_AUTO_TEST_CASE(BinaryRBMClassificationTest) // We allow a 6% tolerance because the RBM may not reconstruct samples as // well. (Typically it does, but we have no guarantee.) - BOOST_REQUIRE_GE(rbmClassificationAccuracy, classificationAccuracy - 6.0); + REQUIRE(rbmClassificationAccuracy >= classificationAccuracy - 6.0); } /* * Tests the SpikeSlabRBM implementation on the Digits dataset. */ -BOOST_AUTO_TEST_CASE(ssRBMClassificationTest) +TEST_CASE("ssRBMClassificationTest", "[RBMNetworkTest]") { size_t batchSize = 10; size_t numEpoches = 3; @@ -184,7 +181,7 @@ BOOST_AUTO_TEST_CASE(ssRBMClassificationTest) double objVal = modelssRBM.Train(msgd); // Test that objective value returned by RBM::Train() is finite. - BOOST_REQUIRE_EQUAL(std::isfinite(objVal), true); + REQUIRE(std::isfinite(objVal) == true); for (size_t i = 0; i < trainData.n_cols; ++i) { @@ -211,7 +208,7 @@ BOOST_AUTO_TEST_CASE(ssRBMClassificationTest) // omitted here for speed. We add a margin of 3% since ssRBM isn't guaranteed // to give us better results (we just generally expect it to be about as good // or better). - BOOST_REQUIRE_GE(ssRbmClassificationAccuracy, 76.18 - 3.0); + REQUIRE(ssRbmClassificationAccuracy >= 76.18 - 3.0); } template @@ -239,13 +236,13 @@ void BuildVanillaNetwork(MatType& trainData, } for (size_t i = 0; i < freeEnergy.n_elem; ++i) - BOOST_REQUIRE_CLOSE(calculatedFreeEnergy(i), freeEnergy(i), 1e-3); + REQUIRE(calculatedFreeEnergy(i) == Approx(freeEnergy(i)).epsilon(1e-5)); } /* * Train and evaluate a Vanilla network with the specified structure. */ -BOOST_AUTO_TEST_CASE(MiscTest) +TEST_CASE("MiscTest", "[RBMNetworkTest]") { arma::Mat X = arma::Mat("0.0, 0.0, 0.0;" "0.0, 1.0, 1.0;" @@ -254,5 +251,3 @@ BOOST_AUTO_TEST_CASE(MiscTest) X = X.t(); BuildVanillaNetwork>(X, 2); } - -BOOST_AUTO_TEST_SUITE_END(); diff --git a/src/mlpack/tests/recurrent_network_test.cpp b/src/mlpack/tests/recurrent_network_test.cpp index 0cba70a301..eb687190ae 100644 --- a/src/mlpack/tests/recurrent_network_test.cpp +++ b/src/mlpack/tests/recurrent_network_test.cpp @@ -19,9 +19,8 @@ #include #include -#include -#include "test_tools.hpp" -#include "serialization.hpp" +#include "catch.hpp" +#include "serialization_catch.hpp" #include "custom_layer.hpp" using namespace mlpack; @@ -29,8 +28,6 @@ using namespace mlpack::ann; using namespace ens; using namespace mlpack::math; -BOOST_AUTO_TEST_SUITE(RecurrentNetworkTest); - /** * Construct a 2-class dataset out of noisy sines. * @@ -75,7 +72,7 @@ void GenerateNoisySines(arma::cube& data, /** * Train the BRNN on a larger dataset. */ -BOOST_AUTO_TEST_CASE(SequenceClassificationBRNNTest) +TEST_CASE("SequenceClassificationBRNNTest", "[RecurrentNetworkTest]") { // Using same test for RNN below. size_t successes = 0; @@ -111,10 +108,10 @@ BOOST_AUTO_TEST_CASE(SequenceClassificationBRNNTest) StandardSGD opt(0.1, 1, 500 * input.n_cols, -100); model.Train(input, labels, opt); - BOOST_TEST_CHECKPOINT("Training over"); + INFO("Training over"); arma::cube prediction; model.Predict(input, prediction); - BOOST_TEST_CHECKPOINT("Prediction over"); + INFO("Prediction over"); size_t error = 0; for (size_t i = 0; i < prediction.n_cols; ++i) @@ -133,7 +130,7 @@ BOOST_AUTO_TEST_CASE(SequenceClassificationBRNNTest) } double classificationError = 1 - double(error) / prediction.n_cols; - BOOST_TEST_CHECKPOINT(classificationError); + INFO(classificationError); if (classificationError <= 0.2) { ++successes; @@ -141,13 +138,13 @@ BOOST_AUTO_TEST_CASE(SequenceClassificationBRNNTest) } } - BOOST_REQUIRE_GE(successes, 1); + REQUIRE(successes >= 1); } /** * Train the vanilla network on a larger dataset. */ -BOOST_AUTO_TEST_CASE(SequenceClassificationTest) +TEST_CASE("SequenceClassificationTest", "[RecurrentNetworkTest]") { // It isn't guaranteed that the recurrent network will converge in the // specified number of iterations using random weights. If this works 1 of 6 @@ -231,7 +228,7 @@ BOOST_AUTO_TEST_CASE(SequenceClassificationTest) } } - BOOST_REQUIRE_GE(successes, 1); + REQUIRE(successes >= 1); } /** @@ -645,13 +642,13 @@ void ReberGrammarTestNetwork(ModelType& model, offset += 3; } - BOOST_REQUIRE_GE(successes, 1); + REQUIRE(successes >= 1); } /** * Train the specified networks on an embedded Reber grammar dataset. */ -BOOST_AUTO_TEST_CASE(LSTMReberGrammarTest) +TEST_CASE("LSTMReberGrammarTest", "[RecurrentNetworkTest]") { RNN > model(5); model.Add >(7, 10); @@ -664,7 +661,7 @@ BOOST_AUTO_TEST_CASE(LSTMReberGrammarTest) /** * Train the specified networks on an embedded Reber grammar dataset. */ -BOOST_AUTO_TEST_CASE(FastLSTMReberGrammarTest) +TEST_CASE("FastLSTMReberGrammarTest", "[RecurrentNetworkTest]") { RNN > model(5); model.Add >(7, 8); @@ -677,7 +674,7 @@ BOOST_AUTO_TEST_CASE(FastLSTMReberGrammarTest) /** * Train the specified networks on an embedded Reber grammar dataset. */ -BOOST_AUTO_TEST_CASE(GRURecursiveReberGrammarTest) +TEST_CASE("GRURecursiveReberGrammarTest", "[RecurrentNetworkTest]") { RNN > model(5); model.Add >(7, 16); @@ -690,7 +687,7 @@ BOOST_AUTO_TEST_CASE(GRURecursiveReberGrammarTest) /** * Train BLSTM on an embedded Reber grammar dataset. */ -BOOST_AUTO_TEST_CASE(BRNNReberGrammarTest) +TEST_CASE("BRNNReberGrammarTest", "[RecurrentNetworkTest]") { BRNN, AddMerge<>, SigmoidLayer<> > model(5); model.Add >(7, 10); @@ -869,14 +866,14 @@ void DistractedSequenceRecallTestNetwork( offset += 2; } - BOOST_REQUIRE_GE(successes, 1); + REQUIRE(successes >= 1); } /** * Train the specified networks on the Derek D. Monner's distracted sequence * recall task. */ -BOOST_AUTO_TEST_CASE(LSTMDistractedSequenceRecallTest) +TEST_CASE("LSTMDistractedSequenceRecallTest", "[RecurrentNetworkTest]") { DistractedSequenceRecallTestNetwork >(4, 8); } @@ -885,7 +882,7 @@ BOOST_AUTO_TEST_CASE(LSTMDistractedSequenceRecallTest) * Train the specified networks on the Derek D. Monner's distracted sequence * recall task. */ -BOOST_AUTO_TEST_CASE(FastLSTMDistractedSequenceRecallTest) +TEST_CASE("FastLSTMDistractedSequenceRecallTest", "[RecurrentNetworkTest]") { DistractedSequenceRecallTestNetwork >(4, 8); } @@ -894,7 +891,7 @@ BOOST_AUTO_TEST_CASE(FastLSTMDistractedSequenceRecallTest) * Train the specified networks on the Derek D. Monner's distracted sequence * recall task. */ -BOOST_AUTO_TEST_CASE(GRUDistractedSequenceRecallTest) +TEST_CASE("GRUDistractedSequenceRecallTest", "[RecurrentNetworkTest]") { DistractedSequenceRecallTestNetwork >(4, 8); } @@ -956,7 +953,7 @@ void BatchSizeTest() /** * Ensure LSTMs work with larger batch sizes. */ -BOOST_AUTO_TEST_CASE(LSTMBatchSizeTest) +TEST_CASE("LSTMBatchSizeTest", "[RecurrentNetworkTest]") { BatchSizeTest>(); } @@ -964,7 +961,7 @@ BOOST_AUTO_TEST_CASE(LSTMBatchSizeTest) /** * Ensure fast LSTMs work with larger batch sizes. */ -BOOST_AUTO_TEST_CASE(FastLSTMBatchSizeTest) +TEST_CASE("FastLSTMBatchSizeTest", "[RecurrentNetworkTest]") { BatchSizeTest>(); } @@ -972,7 +969,7 @@ BOOST_AUTO_TEST_CASE(FastLSTMBatchSizeTest) /** * Ensure GRUs work with larger batch sizes. */ -BOOST_AUTO_TEST_CASE(GRUBatchSizeTest) +TEST_CASE("GRUBatchSizeTest", "[RecurrentNetworkTest]") { BatchSizeTest>(); } @@ -980,7 +977,7 @@ BOOST_AUTO_TEST_CASE(GRUBatchSizeTest) /** * Make sure the RNN can be properly serialized. */ -BOOST_AUTO_TEST_CASE(SerializationTest) +TEST_CASE("RNNSerializationTest", "[RecurrentNetworkTest]") { const size_t rho = 10; @@ -1172,13 +1169,13 @@ void ReberGrammarTestCustomNetwork(const size_t hiddenSize = 4, offset += 3; } - BOOST_REQUIRE_GE(successes, 1); + REQUIRE(successes >= 1); } /** * Train the specified networks on an embedded Reber grammar dataset. */ -BOOST_AUTO_TEST_CASE(CustomRecursiveReberGrammarTest) +TEST_CASE("CustomRecursiveReberGrammarTest", "[RecurrentNetworkTest]") { ReberGrammarTestCustomNetwork(16, true); } @@ -1312,16 +1309,16 @@ double RNNSineTest(size_t hiddenUnits, size_t rho, size_t numEpochs = 100) /** * Test RNN using multiple timestep input and single output. */ -BOOST_AUTO_TEST_CASE(MultiTimestepTest) +TEST_CASE("MultiTimestepTest", "[RecurrentNetworkTest]") { double err = RNNSineTest(4, 10, 20); - BOOST_REQUIRE_LE(err, 0.025); + REQUIRE(err <= 0.025); } /** * Test that RNN::Train() returns finite objective value. */ -BOOST_AUTO_TEST_CASE(RNNTrainReturnObjective) +TEST_CASE("RNNTrainReturnObjective", "[RecurrentNetworkTest]") { const size_t rho = 10; @@ -1371,13 +1368,13 @@ BOOST_AUTO_TEST_CASE(RNNTrainReturnObjective) StandardSGD opt(0.1, 1, input.n_cols /* 1 epoch */, -100); double objVal = model.Train(input, labels, opt); - BOOST_REQUIRE_EQUAL(std::isfinite(objVal), true); + REQUIRE(std::isfinite(objVal) == true); } /** * Test that BRNN::Train() returns finite objective value. */ -BOOST_AUTO_TEST_CASE(BRNNTrainReturnObjective) +TEST_CASE("BRNNTrainReturnObjective", "[RecurrentNetworkTest]") { const size_t rho = 10; @@ -1407,16 +1404,16 @@ BOOST_AUTO_TEST_CASE(BRNNTrainReturnObjective) StandardSGD opt(0.1, 1, 500 * input.n_cols, -100); double objVal = model.Train(input, labels, opt); - BOOST_TEST_CHECKPOINT("Training over"); + INFO("Training over"); // Test that BRNN::Train() returns finite objective value. - BOOST_REQUIRE_EQUAL(std::isfinite(objVal), true); + REQUIRE(std::isfinite(objVal) == true); } /** * Test that RNN::Train() does not give an error for large rho. */ -BOOST_AUTO_TEST_CASE(LargeRhoValueRnnTest) +TEST_CASE("LargeRhoValueRnnTest", "[RecurrentNetworkTest]") { // Setting rho value greater than sequence length which is 17. const size_t rho = 100; @@ -1473,7 +1470,5 @@ BOOST_AUTO_TEST_CASE(LargeRhoValueRnnTest) } ens::SGD<> opt(0.01, 1, 100); model.Train(inputs[0], targets[0], opt); - BOOST_TEST_CHECKPOINT("Training over"); + INFO("Training over"); } - -BOOST_AUTO_TEST_SUITE_END();