diff --git a/CMake/FindStbImage.cmake b/CMake/FindStbImage.cmake new file mode 100644 index 0000000000..7b2f13eab7 --- /dev/null +++ b/CMake/FindStbImage.cmake @@ -0,0 +1,23 @@ +# - Find STB_IMAGE +# Find the STB_IMAGE C++ library +# +# This module sets the following variables: +# STB_IMAGE_FOUND - set to true if the library is found +# STB_IMAGE_INCLUDE_DIR - list of required include directories + +file(GLOB STB_IMAGE_SEARCH_PATHS + ${CMAKE_BINARY_DIR}/deps/stb) +find_path(STB_IMAGE_INCLUDE_DIR + NAMES stb_image.h stb_image_write.h + PATHS ${STB_IMAGE_SEARCH_PATHS}) + +if(STB_IMAGE_INCLUDE_DIR) + set(STB_IMAGE_FOUND YES) +endif () + +# Checks 'REQUIRED'. +include(FindPackageHandleStandardArgs) +find_package_handle_standard_args(STB_IMAGE + REQUIRED_VARS STB_IMAGE_INCLUDE_DIR) + +mark_as_advanced(STB_IMAGE_INCLUDE_DIR) diff --git a/CMakeLists.txt b/CMakeLists.txt index aebd9ccc7c..5cbbb3a3d3 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -13,6 +13,7 @@ option(TEST_VERBOSE "Run test cases with verbose output." OFF) option(BUILD_TESTS "Build tests." ON) option(BUILD_CLI_EXECUTABLES "Build command-line executables." ON) option(DOWNLOAD_ENSMALLEN "If ensmallen is not found, download it." ON) +option(DOWNLOAD_STB_IMAGE "Download stb_image for image loading." ON) # Currently Python bindings aren't known to build successfully on Windows, so # set BUILD_PYTHON_BINDINGS to OFF when the platform is Windows. @@ -75,6 +76,10 @@ if(CMAKE_COMPILER_IS_GNUCC OR "${CMAKE_CXX_COMPILER_ID}" STREQUAL "Clang") set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wall -Wextra -ftemplate-depth=1000") set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -Wall -Wextra") + + # To remove unused functions warnings. + set(CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -Wno-unused-function") + set(CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -Wno-unused-function") endif() # These support libraries are used if we need to link against something @@ -310,6 +315,49 @@ endif () set(MLPACK_INCLUDE_DIRS ${MLPACK_INCLUDE_DIRS} ${ARMADILLO_INCLUDE_DIRS}) set(MLPACK_LIBRARIES ${MLPACK_LIBRARIES} ${ARMADILLO_LIBRARIES}) +# Find stb_image.h and stb_image_write.h. +find_package(StbImage) +# Download stb_image for image loading. +if (NOT STB_IMAGE_FOUND) + if (DOWNLOAD_STB_IMAGE) + set(STB_DIR "stb") + install(DIRECTORY DESTINATION ${CMAKE_BINARY_DIR}/deps/${STB_DIR}) + file(DOWNLOAD http://mlpack.org/files/stb/stb_image.h + "${CMAKE_BINARY_DIR}/deps/${STB_DIR}/stb_image.h" + STATUS STB_IMAGE_DOWNLOAD_STATUS_LIST LOG STB_IMAGE_DOWNLOAD_LOG + SHOW_PROGRESS) + list(GET STB_IMAGE_DOWNLOAD_STATUS_LIST 0 STB_IMAGE_DOWNLOAD_STATUS) + file(DOWNLOAD http://mlpack.org/files/stb/stb_image_write.h + "${CMAKE_BINARY_DIR}/deps/${STB_DIR}/stb_image_write.h" + STATUS STB_IMAGE_WRITE_DOWNLOAD_STATUS_LIST + LOG STB_IMAGE_WRITE_DOWNLOAD_LOG + SHOW_PROGRESS) + list(GET STB_IMAGE_WRITE_DOWNLOAD_STATUS_LIST 0 + STB_IMAGE_WRITE_DOWNLOAD_STATUS) + if (STB_IMAGE_DOWNLOAD_STATUS EQUAL 0 AND + STB_IMAGE_WRITE_DOWNLOAD_STATUS EQUAL 0) + set(MLPACK_INCLUDE_DIRS ${MLPACK_INCLUDE_DIRS} + "${CMAKE_BINARY_DIR}/deps/${STB_DIR}") + message(STATUS + "Successfully downloaded stb into ${CMAKE_BINARY_DIR}/deps/${STB_DIR}/") + # Now we have to also ensure these header files get installed. + install(FILES ${CMAKE_BINARY_DIR}/deps/${STB_DIR}/stb_image.h DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}) + install(FILES ${CMAKE_BINARY_DIR}/deps/${STB_DIR}/stb_image_write.h DESTINATION ${CMAKE_INSTALL_INCLUDEDIR}) + add_definitions(-DHAS_STB) + else () + list(GET STB_IMAGE_DOWNLOAD_STATUS_LIST 1 STB_DOWNLOAD_ERROR) + message(WARNING + "Could not download stb! Error code ${STB_DOWNLOAD_STATUS}: ${STB_DOWNLOAD_ERROR}! Error log: ${STB_DOWNLOAD_LOG}") + endif () + else () + message(WARNING + "stb/stb_image.h is not installed. Image utilites will not be available!") + endif () +else () + # Already has STB installed. + add_definitions(-DHAS_STB) +endif () + # Find ensmallen. # Once ensmallen is readily available in package repos, the automatic downloader # here can be removed. diff --git a/COPYRIGHT.txt b/COPYRIGHT.txt index 29e10ee1f4..73e94ab8e1 100644 --- a/COPYRIGHT.txt +++ b/COPYRIGHT.txt @@ -113,6 +113,7 @@ Copyright: Copyright 2019, Dan Timson Copyright 2019, Miguel Canteras Copyright 2019, Bishwa Karki + Copyright 2019, Mehul Kumar Nirala Copyright 2019, Yashwant Singh Parihar License: BSD-3-clause diff --git a/HISTORY.md b/HISTORY.md index 628a41dd66..e1a3cc6100 100644 --- a/HISTORY.md +++ b/HISTORY.md @@ -4,6 +4,8 @@ * Fix gcc 9 OpenMP compilation issue (#1970). + * Added support for loading and saving of images (#1903). + * Add Multiple Pole Balancing Environment (#1901, #1951). * Added functionality for scaling of data (#1876); see the command-line @@ -633,3 +635,4 @@ * Initial release. See any resolved tickets numbered less than #196 or execute this query: http://www.mlpack.org/trac/query?status=closed&milestone=mlpack+1.0.0 + diff --git a/LICENSE.txt b/LICENSE.txt index d28fe33a7d..febcd581ec 100644 --- a/LICENSE.txt +++ b/LICENSE.txt @@ -12,6 +12,10 @@ licensed under the Boost Software License, version 1.0. This code is found in src/mlpack/core/boost_backport/ and more details on the licensing are available there. +mlpack may contain some usage of the source code of stb, which is licensed +under the MIT License and the Public Domain (www.unlicense.org). This code +is used in src/mlpack/core/data/load_image.hpp. + ---- Copyright (c) 2007-2018, mlpack contributors (see COPYRIGHT.txt) All rights reserved. diff --git a/doc/tutorials/image/image.txt b/doc/tutorials/image/image.txt new file mode 100644 index 0000000000..8a635bf260 --- /dev/null +++ b/doc/tutorials/image/image.txt @@ -0,0 +1,185 @@ +/*! +@file image.txt +@author Mehul Kumar Nirala +@brief Tutorial for how to load and save images in mlpack. + +@page imagetutorial Image Utilities tutorial + +@section intro_imagetut Introduction + +Image datasets are becoming increasingly popular in deep learning. + +mlpack's image saving/loading functionality is based on [stb/](https://github.com/nothings/stb). + +@section toc_imagetut Table of Contents + +This tutorial is split into the following sections: + + - \ref intro_imagetut + - \ref toc_imagetut + - \ref model_api_imagetut + - \ref imageinfo_api_imagetut + - \ref load_api_imagetut + - \ref save_api_imagetut + +@section model_api_imagetut Model API + +Image utilities supports loading and saving of images. + +It supports filetypes "jpg", "png", "tga","bmp", "psd", "gif", "hdr", "pic", "pnm" for loading and "jpg", "png", "tga", "bmp", "hdr" for saving. + +The datatype associated is unsigned char to support RGB values in the range 1-255. To feed data into the network typecast of `arma::Mat` may be required. Images are stored in matrix as (width * height * channels, NumberOfImages). Therefore imageMatrix.col(0) would be the first image if images are loaded in imageMatrix. + +@section imageinfo_api_imagetut ImageInfo + +ImageInfo class contains the metadata of the images. +@code + /** + * Instantiate the ImageInfo object with the image width, height, channels. + * + * @param width Image width. + * @param height Image height. + * @param channels number of channels in the image. + */ + ImageInfo(const size_t width, + const size_t height, + const size_t channels); +@endcode +Other public memebers include: + - flipVertical Flip the image vertical upon loading. + - quality Compression of the image if saved as jpg (0-100). + +@section load_api_imagetut Load + + +Standalone loading of images. +@code + /** + * Load the image file into the given matrix. + * + * @param filename Name of the image file. + * @param matrix Matrix to load the image into. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ + template + bool Load(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose); +@endcode + +Loading a test image. It also fills up the ImageInfo class object. +@code +data::ImageInfo info; +data::Load("test_image.png", matrix, info, false, true); +@endcode + +ImageInfo requires height, width, number of channels of the image. + +@code +size_t height = 64, width = 64, channels = 1; +data::ImageInfo info(width, height, channels); +@endcode + +More than one image can be loaded into the same matrix. + +Loading multiple images: + +@code + /** + * Load the image file into the given matrix. + * + * @param files A vector consisting of filenames. + * @param matrix Matrix to save the image from. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ + template + bool Load(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose); +@endcode + +@code + data::ImageInfo info; + std::vector> files{"test_image1.bmp","test_image2.bmp"}; + data::load(files, matrix, info, false, true); +@endcode + +@section save_api_imagetut Save + +Save images expects a matrix of type unsigned char in the form (width * height * channels, NumberOfImages). +Just like load it can be used to save one image or multiple images. Besides image data it also expects the shape of the image as input (width, height, channels). + +Saving one image: + +@code + /** + * Save the image file from the given matrix. + * + * @param filename Name of the image file. + * @param matrix Matrix to save the image from. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ + template + bool Save(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose); +@endcode + +@code + data::ImageInfo info; + info.width = info.height = 25; + info.channels = 3; + info.quality = 90; + data::Save("test_image.bmp", matrix, info, false, true); +@endcode + +If the matrix contains more than one image, only the first one is saved. + +Saving multiple images: + +@code + /** + * Save the image file from the given matrix. + * + * @param files A vector consisting of filenames. + * @param matrix Matrix to save the image from. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ + template + bool Save(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose); +@endcode + +@code + data::ImageInfo info; + info.width = info.height = 25; + info.channels = 3; + info.quality = 90; + std::vector> files{"test_image1.bmp", "test_image2.bmp"}; + data::Save(files, matrix, info, false, true); +@endcode + +Multiple images are saved according to the vector of filenames specified. + +*/ diff --git a/src/mlpack/core/data/CMakeLists.txt b/src/mlpack/core/data/CMakeLists.txt index d06216a843..5a7e13a352 100644 --- a/src/mlpack/core/data/CMakeLists.txt +++ b/src/mlpack/core/data/CMakeLists.txt @@ -10,6 +10,7 @@ set(SOURCES load_csv.hpp load_csv.cpp load.hpp + load_image_impl.hpp load_model_impl.hpp load_vec_impl.hpp load_impl.hpp diff --git a/src/mlpack/core/data/image_info.hpp b/src/mlpack/core/data/image_info.hpp new file mode 100644 index 0000000000..cd8c86ee8e --- /dev/null +++ b/src/mlpack/core/data/image_info.hpp @@ -0,0 +1,108 @@ +/** + * @file image_info.hpp + * @author Mehul Kumar Nirala + * + * An image information holder. + * + * 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_IMAGE_INFO_HPP +#define MLPACK_CORE_DATA_IMAGE_INFO_HPP + + +#include + +#include "extension.hpp" + +#define STB_IMAGE_STATIC +#define STB_IMAGE_IMPLEMENTATION +#include + +#define STB_IMAGE_WRITE_STATIC +#define STB_IMAGE_WRITE_IMPLEMENTATION +#include + +namespace mlpack { +namespace data { + +#ifdef HAS_STB // Compile this only if stb is present. + +/** + * Checks if the given image filename is supported. + * + * @param filename Name of the image file. + * @return Boolean value indicating success if it is an image. + */ +inline bool ImageFormatSupported(const std::string& fileName, + const bool save = false); + +/** + * Implements meta-data of images required by data::Load and + * data::Save for loading and saving images into arma::Mat. + */ +class ImageInfo +{ + public: + /** + * Instantiate the ImageInfo object with the given image width, height, + * number of channels and quality parameter. + * + * @param width Image width. + * @param height Image height. + * @param channels Number of channels in the image. + * @param quality Compression of the image if saved as jpg (0 - 100). + */ + ImageInfo(const size_t width = 0, + const size_t height = 0, + const size_t channels = 3, + const size_t quality = 90); + + //! Get the image width. + const size_t& Width() const { return width; } + //! Modify the image width. + size_t& Width() { return width; } + //! Get the image height. + + const size_t& Height() const { return height; } + //! Modify the image height. + size_t& Height() { return height; } + + //! Get the image channels. + const size_t& Channels() const { return channels; } + //! Modify the image channels. + size_t& Channels() { return channels; } + + //! Get the image quality. + const size_t& Quality() const { return quality; } + //! Modify the image quality. + size_t& Quality() { return quality; } + + private: + // To store the image width. + size_t width; + + // To store the image height. + size_t height; + + // To store the number of channels in the image. + size_t channels; + + // Compression of the image if saved as jpg (0 - 100). + size_t quality; +}; +#else +class ImageInfo { }; + +#endif // HAS_STB. + +} // namespace data +} // namespace mlpack + +// Include implementation of Image. +#include "image_info_impl.hpp" + +#endif diff --git a/src/mlpack/core/data/image_info_impl.hpp b/src/mlpack/core/data/image_info_impl.hpp new file mode 100644 index 0000000000..dac80d163e --- /dev/null +++ b/src/mlpack/core/data/image_info_impl.hpp @@ -0,0 +1,71 @@ +/** + * @file image_info_impl.hpp + * @author Mehul Kumar Nirala + * + * An image information holder implementation. + * + * 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_IMAGE_INFO_IMPL_HPP +#define MLPACK_CORE_DATA_IMAGE_INFO_IMPL_HPP + +#ifdef HAS_STB // Compile this only if stb is present. + +// In case it hasn't been included yet. +#include "image_info.hpp" + +namespace mlpack { +namespace data { + +static const std::vector loadFileTypes({"jpg", "png", "tga", + "bmp", "psd", "gif", "hdr", "pic", "pnm", "jpeg"}); + +static const std::vector saveFileTypes({"jpg", "png", "tga", + "bmp", "hdr"}); + +inline bool ImageFormatSupported(const std::string& fileName, const bool save) +{ + if (save) + { + // Iterate over all supported file types that can be saved. + for (auto extension : saveFileTypes) + { + if (extension == Extension(fileName)) + return true; + } + } + else + { + // Iterate over all supported file types that can be loaded. + for (auto extension : loadFileTypes) + { + if (extension == Extension(fileName)) + return true; + } + } + + return false; +} + +inline ImageInfo::ImageInfo(const size_t width, + const size_t height, + const size_t channels, + const size_t quality) : + width(width), + height(height), + channels(channels), + quality(quality) +{ + // Do nothing. +} + +} // namespace data +} // namespace mlpack + +#endif // HAS_STB. + +#endif diff --git a/src/mlpack/core/data/load.hpp b/src/mlpack/core/data/load.hpp index 99a91807e7..7974227409 100644 --- a/src/mlpack/core/data/load.hpp +++ b/src/mlpack/core/data/load.hpp @@ -20,6 +20,7 @@ #include "format.hpp" #include "dataset_mapper.hpp" +#include "image_info.hpp" namespace mlpack { namespace data /** Functions to load and save matrices and models. */ { @@ -287,6 +288,47 @@ bool Load(const std::string& filename, const bool fatal = false, format f = format::autodetect); +/** + * Image load/save interfaces. + */ +#ifdef HAS_STB + +/** + * Load the image file into the given matrix. + * + * @param filename Name of the image file. + * @param matrix Matrix to load the image into. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ +template +bool Load(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true); + +/** + * Load the image file into the given matrix. + * + * @param files A vector consisting of filenames. + * @param matrix Matrix to save the image from. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ +template +bool Load(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true); + +#endif // HAS_STB. + } // namespace data } // namespace mlpack @@ -294,5 +336,7 @@ bool Load(const std::string& filename, #include "load_model_impl.hpp" // Include implementation of Load() for vectors. #include "load_vec_impl.hpp" +// Include implementation of Load() for images. +#include "load_image_impl.hpp" #endif diff --git a/src/mlpack/core/data/load_image_impl.hpp b/src/mlpack/core/data/load_image_impl.hpp new file mode 100644 index 0000000000..e4fa8230af --- /dev/null +++ b/src/mlpack/core/data/load_image_impl.hpp @@ -0,0 +1,149 @@ +/** + * @file load_image_impl.hpp + * @author Mehul Kumar Nirala + * + * An image loading utility implementation. + * + * 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_LOAD_IMAGE_IMPL_HPP +#define MLPACK_CORE_DATA_LOAD_IMAGE_IMPL_HPP + + +// In case it hasn't been included yet. +#include "load.hpp" + +namespace mlpack { +namespace data { + +#ifdef HAS_STB // Compile this only if stb is present. + +// Image loading API. +template +bool Load(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose) +{ + Timer::Start("loading_image"); + unsigned char* image; + + if (!ImageFormatSupported(filename)) + { + std::ostringstream oss; + oss << "File type " << Extension(filename) << " not supported.\n"; + oss << "Currently it supports "; + for (auto extension : loadFileTypes) + oss << " " << extension; + oss << std::endl; + throw std::runtime_error(oss.str()); + return false; + } + + stbi_set_flip_vertically_on_load(transpose); + + // Temporary variables needed as stb_image.h supports int parameters. + int tempWidth, tempHeight, tempChannels; + + // For grayscale images. + if (info.Channels() == 1) + { + image = stbi_load(filename.c_str(), &tempWidth, &tempHeight, &tempChannels, + STBI_grey); + } + else + { + image = stbi_load(filename.c_str(), &tempWidth, &tempHeight, &tempChannels, + STBI_rgb); + } + + if (tempWidth <= 0 || tempHeight <= 0) + { + std::ostringstream oss; + oss << "Image '" << filename << "' not found." << std::endl; + free(image); + throw std::runtime_error(oss.str()); + + return false; + } + + info.Width() = tempWidth; + info.Height() = tempHeight; + info.Channels() = tempChannels; + + // Copy image into armadillo Mat. + matrix = arma::Mat(image, info.Width() * info.Height() * + info.Channels(), 1, true, true); + + // Free the image pointer. + free(image); + Timer::Stop("loading_image"); + return true; +} + +// Image loading API for multiple files. +template +bool Load(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose) +{ + if (files.size() == 0) + { + std::ostringstream oss; + oss << "Files vector is empty." << std::endl; + + throw std::runtime_error(oss.str()); + return false; + } + + arma::Mat img; + bool status = Load(files[0], img, info, fatal, transpose); + + // Decide matrix dimension using the image height and width. + matrix.set_size(info.Width() * info.Height() * info.Channels(), files.size()); + matrix.col(0) = img; + + for (size_t i = 1; i < files.size() ; i++) + { + arma::Mat colImg(matrix.colptr(i), matrix.n_rows, 1, + false, true); + status &= Load(files[i], colImg, info, fatal, transpose); + } + return status; +} + +#else // No STB. +template +bool Load(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true) +{ + throw std::runtime_error("Load(): HAS_STB is not defined, " + "so STB is not available and images cannot be loaded!"); +} + +template +bool Load(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true) +{ + throw std::runtime_error("Load(): HAS_STB is not defined, " + "so STB is not available and images cannot be loaded!"); +} +#endif // HAS_STB. + +} // namespace data +} // namespace mlpack + +#endif diff --git a/src/mlpack/core/data/save.hpp b/src/mlpack/core/data/save.hpp index 52c1b57d13..121265ac81 100644 --- a/src/mlpack/core/data/save.hpp +++ b/src/mlpack/core/data/save.hpp @@ -19,6 +19,7 @@ #include #include "format.hpp" +#include "image_info.hpp" namespace mlpack { namespace data /** Functions to load and save matrices. */ { @@ -90,6 +91,44 @@ bool Save(const std::string& filename, const bool fatal = false, format f = format::autodetect); +#ifdef HAS_STB + +/** + * Save the image file from the given matrix. + * + * @param filename Name of the image file. + * @param matrix Matrix to save the image from. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ +template +bool Save(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true); + +/** + * Save the image file from the given matrix. + * + * @param files A vector consisting of filenames. + * @param matrix Matrix to save the image from. + * @param info An object of ImageInfo class. + * @param fatal If an error should be reported as fatal (default false). + * @param transpose If true, transpose the matrix after loading. + * @return Boolean value indicating success or failure of load. + */ +template +bool Save(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true); + +#endif // HAS_STB. + } // namespace data } // namespace mlpack diff --git a/src/mlpack/core/data/save_impl.hpp b/src/mlpack/core/data/save_impl.hpp index 94d3547ac0..18a61530a5 100644 --- a/src/mlpack/core/data/save_impl.hpp +++ b/src/mlpack/core/data/save_impl.hpp @@ -283,6 +283,147 @@ bool Save(const std::string& filename, } } +#ifdef HAS_STB +// Image saving API. +template +bool Save(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose) +{ + Timer::Start("saving_image"); + // We transpose by default. So, un-transpose if necessary. + if (!transpose) + matrix = arma::trans(matrix); + + int tempWidth, tempHeight, tempChannels, tempQuality; + + tempWidth = info.Width(); + tempHeight = info.Height(); + tempChannels = info.Channels(); + tempQuality = info.Quality(); + + if (!ImageFormatSupported(filename, true)) + { + std::ostringstream oss; + oss << "File type " << Extension(filename) << " not supported.\n"; + oss << "Currently it supports "; + for (auto extension : saveFileTypes) + oss << ", " << extension; + oss << std::endl; + throw std::runtime_error(oss.str()); + return false; + } + if (matrix.n_cols > 1) + { + std::cout << "Input Matrix contains more than 1 image." << std::endl; + std::cout << "Only the firstimage will be saved!" << std::endl; + } + stbi_flip_vertically_on_write(transpose); + + bool status = false; + try + { + unsigned char* image = matrix.memptr(); + + if ("png" == Extension(filename)) + { + status = stbi_write_png(filename.c_str(), tempWidth, tempHeight, + tempChannels, image, tempWidth * tempChannels); + } + else if ("bmp" == Extension(filename)) + { + status = stbi_write_bmp(filename.c_str(), tempWidth, tempHeight, + tempChannels, image); + } + else if ("tga" == Extension(filename)) + { + status = stbi_write_tga(filename.c_str(), tempWidth, tempHeight, + tempChannels, image); + } + else if ("hdr" == Extension(filename)) + { + status = stbi_write_hdr(filename.c_str(), tempWidth, tempHeight, + tempChannels, reinterpret_cast(image)); + } + else if ("jpg" == Extension(filename)) + { + status = stbi_write_jpg(filename.c_str(), tempWidth, tempHeight, + tempChannels, image, tempQuality); + } + } + catch (std::exception& e) + { + Timer::Stop("saving_image"); + if (fatal) + Log::Fatal << e.what() << std::endl; + Log::Warn << e.what() << std::endl; + return false; + } + Timer::Stop("saving_image"); + return status; +} + +// Image saving API for multiple files. +template +bool Save(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal, + const bool transpose) +{ + if (files.size() == 0) + { + std::ostringstream oss; + oss << "Files vector is empty." << std::endl; + + throw std::runtime_error(oss.str()); + return false; + } + // We transpose by default. So, un-transpose if necessary. + if (!transpose) + matrix = arma::trans(matrix); + + arma::Mat img; + bool status = Save(files[0], img, info, fatal, transpose); + + // Decide matrix dimension using the image height and width. + matrix.set_size(info.Width() * info.Height() * info.Channels(), files.size()); + matrix.col(0) = img; + + for (size_t i = 1; i < files.size() ; i++) + { + arma::Mat colImg(matrix.colptr(i), matrix.n_rows, 1, + false, true); + status &= Save(files[i], colImg, info, fatal, transpose); + } + return status; +} +#else +template +bool Save(const std::string& filename, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true) +{ + throw std::runtime_error("Save(): HAS_STB is not defined, " + "so STB is not available and images cannot be saved!"); +} + +template +bool Save(const std::vector& files, + arma::Mat& matrix, + ImageInfo& info, + const bool fatal = false, + const bool transpose = true) +{ + throw std::runtime_error("Save(): HAS_STB is not defined, " + "so STB is not available and images cannot be saved!"); +} +#endif // HAS_STB. + } // namespace data } // namespace mlpack diff --git a/src/mlpack/tests/CMakeLists.txt b/src/mlpack/tests/CMakeLists.txt index 80c1c231e1..b1d80eed5a 100644 --- a/src/mlpack/tests/CMakeLists.txt +++ b/src/mlpack/tests/CMakeLists.txt @@ -37,6 +37,7 @@ add_executable(mlpack_test hoeffding_tree_test.cpp hpt_test.cpp hyperplane_test.cpp + image_load_test.cpp imputation_test.cpp init_rules_test.cpp kde_test.cpp diff --git a/src/mlpack/tests/data/test_image.png b/src/mlpack/tests/data/test_image.png new file mode 100644 index 0000000000..5a29595a1e Binary files /dev/null and b/src/mlpack/tests/data/test_image.png differ diff --git a/src/mlpack/tests/image_load_test.cpp b/src/mlpack/tests/image_load_test.cpp new file mode 100644 index 0000000000..0e2bbb0459 --- /dev/null +++ b/src/mlpack/tests/image_load_test.cpp @@ -0,0 +1,72 @@ +/** + * @file image_load_test.cpp + * @author Mehul Kumar Nirala + * + * Tests for loading and saving images. + * + * 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. + */ + +#include +#include + +using namespace mlpack; +using namespace mlpack::data; +using namespace std; + +#ifdef HAS_STB // Compile this only if stb is present. + +BOOST_AUTO_TEST_SUITE(ImageLoadTest); + +/** + * Test if an image with an unsupported extension throws an expected + * exception. + */ +BOOST_AUTO_TEST_CASE(LoadInvalidExtensionFile) +{ + arma::Mat matrix; + data::ImageInfo info; + BOOST_REQUIRE_THROW(data::Load("invalidExtendion.p4ng", matrix, info, + false, true), std::runtime_error); +} + +/** + * Test that the image is loaded correctly into the matrix using the API. + */ +BOOST_AUTO_TEST_CASE(LoadImageAPITest) +{ + arma::Mat matrix; + data::ImageInfo info; + BOOST_REQUIRE(data::Load("test_image.png", matrix, info, false, + true) == true); + BOOST_REQUIRE_EQUAL(matrix.n_rows, 50 * 50 * 3); // width * height * channels. + BOOST_REQUIRE_EQUAL(matrix.n_cols, 1); +} + +/** + * Test if the image is saved correctly using API. + */ +BOOST_AUTO_TEST_CASE(SaveImageAPITest) +{ + data::ImageInfo info(5, 5, 3, 90); + + arma::Mat im1; + size_t dimension = info.Width() * info.Height() * info.Channels(); + im1 = arma::randi>(dimension, 1); + BOOST_REQUIRE(data::Save("APITest.bmp", im1, info, false, true) == true); + + arma::Mat im2; + BOOST_REQUIRE(data::Load("APITest.bmp", im2, info, false, true) == true); + + BOOST_REQUIRE_EQUAL(im1.n_cols, im2.n_cols); + BOOST_REQUIRE_EQUAL(im1.n_rows, im2.n_rows); + for (size_t i = 10; i < im1.n_elem; ++i) + BOOST_REQUIRE_EQUAL(im1[i], im2[i]); +} + +BOOST_AUTO_TEST_SUITE_END(); + +#endif // HAS_STB.