Merge remote-tracking branch 'origin/master' into ann-vtable
This commit is contained in:
+226
-155
@@ -102,15 +102,24 @@ corresponding to a given machine learning algorithm, with the suffix
|
||||
|
||||
These files have roughly two parts:
|
||||
|
||||
- definition of the input and output parameters with @c PARAM macros
|
||||
- implementation of @c mlpackMain(), which is the actual machine learning code
|
||||
- definition of the input and output parameters with @c PARAM macros and
|
||||
documentation with @c BINDING macros
|
||||
- implementation of @c BINDING_FUNCTION(), which is the actual machine learning
|
||||
code
|
||||
|
||||
Here is a simple example file:
|
||||
|
||||
@code
|
||||
// This is a stripped version of mean_shift_main.cpp.
|
||||
#include <mlpack/prereqs.hpp>
|
||||
#include <mlpack/core/util/cli.hpp>
|
||||
#include <mlpack/core/util/io.hpp>
|
||||
|
||||
// Define the name of the binding (as seen by the binding generation system).
|
||||
#ifdef BINDING_NAME
|
||||
#undef BINDING_NAME
|
||||
#endif
|
||||
#define BINDING_NAME mean_shift
|
||||
|
||||
#include <mlpack/core/util/mlpack_main.hpp>
|
||||
|
||||
#include <mlpack/core/kernels/gaussian_kernel.hpp>
|
||||
@@ -130,7 +139,7 @@ using namespace std;
|
||||
// generate documentation for the website.
|
||||
|
||||
// Program Name.
|
||||
BINDING_NAME("Mean Shift Clustering");
|
||||
BINDING_USER_NAME("Mean Shift Clustering");
|
||||
|
||||
// Short description.
|
||||
BINDING_SHORT_DESC(
|
||||
@@ -192,11 +201,11 @@ PARAM_DOUBLE_IN("radius", "If the distance between two centroids is less than "
|
||||
"the given radius, one will be removed. A radius of 0 or less means an "
|
||||
"estimate will be calculated and used for the radius.", "r", 0);
|
||||
|
||||
void mlpackMain()
|
||||
void BINDING_FUNCTION(util::Params& params, util::Timers& timers)
|
||||
{
|
||||
// Process the parameters that the user passed.
|
||||
const double radius = IO::GetParam<double>("radius");
|
||||
const int maxIterations = IO::GetParam<int>("max_iterations");
|
||||
const double radius = params.Get<double>("radius");
|
||||
const int maxIterations = params.Get<int>("max_iterations");
|
||||
|
||||
if (maxIterations < 0)
|
||||
{
|
||||
@@ -205,43 +214,44 @@ void mlpackMain()
|
||||
}
|
||||
|
||||
// Warn, if the user did not specify that they wanted any output.
|
||||
if (!IO::HasParam("output") && !IO::HasParam("centroid"))
|
||||
if (!params.Has("output") && !params.Has("centroid"))
|
||||
{
|
||||
Log::Warn << "--output_file, --in_place, and --centroid_file are not set; "
|
||||
<< "no results will be saved." << endl;
|
||||
}
|
||||
|
||||
arma::mat dataset = std::move(IO::GetParam<arma::mat>("input"));
|
||||
arma::mat dataset = std::move(params.Get<arma::mat>("input"));
|
||||
arma::mat centroids;
|
||||
arma::Col<size_t> assignments;
|
||||
|
||||
// Prepare and run the actual algorithm.
|
||||
MeanShift<> meanShift(radius, maxIterations);
|
||||
|
||||
Timer::Start("clustering");
|
||||
timers.Start("clustering");
|
||||
Log::Info << "Performing mean shift clustering..." << endl;
|
||||
meanShift.Cluster(dataset, assignments, centroids);
|
||||
Timer::Stop("clustering");
|
||||
timers.Stop("clustering");
|
||||
|
||||
Log::Info << "Found " << centroids.n_cols << " centroids." << endl;
|
||||
if (radius <= 0.0)
|
||||
Log::Info << "Estimated radius was " << meanShift.Radius() << ".\n";
|
||||
|
||||
// Should we give the user the output matrix?
|
||||
if (IO::HasParam("output"))
|
||||
IO::GetParam<arma::Col<size_t>>("output") = std::move(assignments);
|
||||
if (params.Has("output"))
|
||||
params.Get<arma::Col<size_t>>("output") = std::move(assignments);
|
||||
|
||||
// Should we give the user the centroid matrix?
|
||||
if (IO::HasParam("centroid"))
|
||||
IO::GetParam<arma::mat>("centroid") = std::move(centroids);
|
||||
if (params.Has("centroid"))
|
||||
params.Get<arma::mat>("centroid") = std::move(centroids);
|
||||
}
|
||||
@endcode
|
||||
|
||||
We can see that we have defined the basic program information in the
|
||||
@c BINDING_NAME(), @c BINDING_SHORT_DESC(), @c BINDING_LONG_DESC(),
|
||||
@c BINDING_EXAMPLE() and @c BINDING_SEE_ALSO() macros. This is, for instance,
|
||||
what is displayed to describe the binding if the user passed the
|
||||
<tt>\--help</tt> option for a command-line program.
|
||||
We can see that we have defined the name of the binding with the @c BINDING_NAME
|
||||
macro, and basic program information in the @c BINDING_USER_NAME(), @c
|
||||
BINDING_SHORT_DESC(), @c BINDING_LONG_DESC(), @c BINDING_EXAMPLE() and @c
|
||||
BINDING_SEE_ALSO() macros. This is, for instance, what is displayed to describe
|
||||
the binding if the user passed the <tt>\--help</tt> option for a command-line
|
||||
program.
|
||||
|
||||
Then, we define five parameters, three input and two output, that define the
|
||||
data and options that the mean shift clustering will function on. These
|
||||
@@ -258,16 +268,22 @@ whether the parameter is input or output. Some examples:
|
||||
Note that each of these macros may have slightly different syntax. See the
|
||||
links above for further documentation.
|
||||
|
||||
In order to write a new binding, then, you simply must write @c BINDING_NAME(),
|
||||
@c BINDING_SHORT_DESC(), @c BINDING_LONG_DESC(), @c BINDING_EXAMPLE() and
|
||||
@c BINDING_SEE_ALSO() definitions of the program with some docuentation, define
|
||||
the input and output parameters as @c PARAM macros, and then write an
|
||||
@c mlpackMain() function that actually performs the functionality of the binding.
|
||||
Inside of @c mlpackMain():
|
||||
In order to write a new binding, then, you simply must define @c BINDING_NAME,
|
||||
then write @c BINDING_USER_NAME(), @c BINDING_SHORT_DESC(), @c
|
||||
BINDING_LONG_DESC(), @c BINDING_EXAMPLE() and @c BINDING_SEE_ALSO() definitions
|
||||
of the program with some docuentation, define the input and output parameters as
|
||||
@c PARAM macros, and then write a @c BINDING_FUNCTION() function that actually
|
||||
performs the functionality of the binding.
|
||||
|
||||
- All input parameters are accessible through @c IO::GetParam<type>("name").
|
||||
Inside of @c BINDING_FUNCTION(util::Params& params, util::Timers& timers):
|
||||
|
||||
- All input parameters are accessible through @c params.Get<type>("name").
|
||||
- All output parameters should be set by the end of the function with the
|
||||
@c IO::GetParam<type>("name") method.
|
||||
@c params.Get<type>("name") method.
|
||||
- The @c params.Has("name") function will return @c true if the parameter
|
||||
@c "name" was specified.
|
||||
- Timers can be started and stopped with @c timers.Start("timer_name") and
|
||||
@c timers.Stop("timer_name").
|
||||
|
||||
Then, assuming that your program is saved in the file @c program_name_main.cpp,
|
||||
generating bindings for other languages is a simple addition to the
|
||||
@@ -285,24 +301,48 @@ the categories in @c src/mlpack/bindings/markdown/MarkdownCategories.cmake.
|
||||
|
||||
@section bindings_general How to write mlpack bindings
|
||||
|
||||
This section describes the general structure of the @c IO code and how one
|
||||
might write a new binding for mlpack. After reading this section it should be
|
||||
relatively clear how one could use the @c IO functionality along with CMake to
|
||||
add a binding for a new mlpack machine learning method. If it is not clear,
|
||||
then the examples in the following sections should clarify.
|
||||
This section describes the general structure of the automatic binding system and
|
||||
how one might write a new binding for mlpack. After reading this section it
|
||||
should be relatively clear how one could use the provided functionality in the
|
||||
@c Params and @c Timers class along with CMake to add a binding for a new mlpack
|
||||
machine learning method. If it is not clear, then the examples in the following
|
||||
sections should clarify.
|
||||
|
||||
@subsection bindings_general_binding_name Providing a name with @c BINDING_NAME
|
||||
|
||||
Every binding must have the macro @c BINDING_NAME defined, specifying a name
|
||||
(without spaces, generally all lowercase) that will be used to represent the
|
||||
binding. It is suggested to @c #undef any previous setting of @c BINDING_NAME
|
||||
just to prevent any strange error messages in case it is already defined.
|
||||
|
||||
Here is an example that can be adapted:
|
||||
|
||||
@code
|
||||
#ifdef BINDING_NAME
|
||||
#undef BINDING_NAME
|
||||
#endif
|
||||
#define BINDING_NAME my_binding_name
|
||||
|
||||
// BINDING_NAME should be defined before including mlpack_main.hpp!
|
||||
#include <mlpack/core/util/mlpack_main.hpp>
|
||||
@endcode
|
||||
|
||||
If this macro is not defined, compilation of the binding will fail in many ways
|
||||
with potentially obscure error messages! (Sorry that they are bad error
|
||||
messages. The preprocessor doesn't give us too much to work with.)
|
||||
|
||||
@subsection bindings_general_program_doc Documenting a program with
|
||||
@c BINDING_NAME(), @c BINDING_SHORT_DESC(), @c BINDING_LONG_DESC(),
|
||||
@c BINDING_USER_NAME(), @c BINDING_SHORT_DESC(), @c BINDING_LONG_DESC(),
|
||||
@c BINDING_EXAMPLE() and @c BINDING_SEE_ALSO().
|
||||
|
||||
Any mlpack program should be documented with the @c BINDING_NAME(),
|
||||
Any mlpack program should be documented with the @c BINDING_USER_NAME(),
|
||||
@c BINDING_SHORT_DESC(), @c BINDING_LONG_DESC() , @c BINDING_EXAMPLE() and
|
||||
@c BINDING_SEE_ALSO() macros, which is available from the
|
||||
@c <mlpack/core/util/mlpack_main.hpp> header. The macros
|
||||
are of the form
|
||||
|
||||
@code
|
||||
BINDING_NAME("program name");
|
||||
BINDING_USER_NAME("program name");
|
||||
BINDING_SHORT_DESC("This is a short, two-sentence description of what the program does.");
|
||||
BINDING_LONG_DESC("This is a long description of what the program does."
|
||||
" It might be many lines long and have lots of details about different options.");
|
||||
@@ -473,7 +513,7 @@ Go binding output (snippet):
|
||||
Input C++ (full program, 'random_numbers_main.cpp'):
|
||||
|
||||
// Program Name.
|
||||
BINDING_NAME("Random Numbers");
|
||||
BINDING_USER_NAME("Random Numbers");
|
||||
|
||||
// Short description.
|
||||
BINDING_SHORT_DESC("An implementation of Random Numbers");
|
||||
@@ -503,9 +543,11 @@ Input C++ (full program, 'random_numbers_main.cpp'):
|
||||
"\n\n" +
|
||||
PRINT_CALL("random_numbers", "num_values", 100, "subtract", 3, "output",
|
||||
"rand", "output_model", "rand_lr"));
|
||||
@endcode
|
||||
|
||||
Command line output:
|
||||
|
||||
@code
|
||||
Random Numbers
|
||||
|
||||
This program generates random numbers with a variety of nonsensical
|
||||
@@ -525,9 +567,11 @@ Command line output:
|
||||
|
||||
$ random_numbers --num_values 100 --subtract 3 --output_file rand.csv
|
||||
--output_model_file rand_lr.bin
|
||||
@endcode
|
||||
|
||||
Python binding output:
|
||||
|
||||
@code
|
||||
Random Numbers
|
||||
|
||||
This program generates random numbers with a variety of nonsensical
|
||||
@@ -548,9 +592,11 @@ Python binding output:
|
||||
>>> output = random_numbers(num_values=100, subtract=3)
|
||||
>>> rand = output['output']
|
||||
>>> rand_lr = output['output_model']
|
||||
@endcode
|
||||
|
||||
Julia binding output:
|
||||
|
||||
@code
|
||||
Random Numbers
|
||||
|
||||
This program generates random numbers with a variety of nonsensical
|
||||
@@ -571,9 +617,11 @@ Julia binding output:
|
||||
```julia
|
||||
julia> rand, rand_lr = random_numbers(num_values=100, subtract=3)
|
||||
```
|
||||
@endcode
|
||||
|
||||
Go binding output:
|
||||
|
||||
@code
|
||||
Random Numbers
|
||||
|
||||
This program generates random numbers with a variety of nonsensical
|
||||
@@ -733,32 +781,31 @@ Note that even the parameter documentation strings must be a little be agnostic
|
||||
to the binding type, because the command-line interface is so different than the
|
||||
Python interface to the user.
|
||||
|
||||
@subsection bindings_general_functions Using IO in an mlpackMain() function
|
||||
@subsection bindings_general_functions Using @c Params in a @c BINDING_FUNCTION() function
|
||||
|
||||
mlpack's @c IO module provides a unified abstract interface for getting input
|
||||
mlpack's @c util::Params class provides a unified abstract interface for getting input
|
||||
from and providing output to users without needing to consider the language
|
||||
(command-line, Python, MATLAB, etc.) that the user is running the program from.
|
||||
This means that after the @c BINDING_LONG_DESC() and @c BINDING_EXAMPLE() macros
|
||||
and the @c PARAM_*() macros have been defined, a language-agnostic
|
||||
@c mlpackMain() function can be written. This function then can perform the
|
||||
actual computation that the entire program is meant to.
|
||||
and the @c PARAM_*() macros have been defined, a language-agnostic <tt>void
|
||||
BINDING_FUNCTION(util::Params& params, util::Timers& timers)</tt> function can
|
||||
be written. This function then can perform the actual computation that the
|
||||
entire program is meant to.
|
||||
|
||||
Inside of an @c mlpackMain() function, the @c mlpack::IO module can be used to
|
||||
access input parameters and set output parameters. There are two main functions
|
||||
for this, plus a utility printing function:
|
||||
|
||||
- @c IO::GetParam<T>() - get a reference to a parameter
|
||||
- @c IO::HasParam() - returns true if the user specified the parameter
|
||||
- @c IO::GetPrintableParam<T>() - returns a string representing the value of
|
||||
the parameter
|
||||
- @c params.Get<T>() - get a reference to a parameter
|
||||
- @c params.Has() - returns true if the user specified the parameter
|
||||
- @c params.GetPrintable<T>() - returns a string representing the value of the
|
||||
parameter
|
||||
|
||||
So, to print "hello" if the user specified the @c print_hello parameter, the
|
||||
following code could be used:
|
||||
|
||||
@code
|
||||
using namespace mlpack;
|
||||
|
||||
if (IO::HasParam("print_hello"))
|
||||
if (params.Has("print_hello"))
|
||||
std::cout << "Hello!" << std::endl;
|
||||
else
|
||||
std::cout << "No greetings for you!" << std::endl;
|
||||
@@ -768,17 +815,13 @@ To access a string that a user passed in to the @c string parameter, the
|
||||
following code could be used:
|
||||
|
||||
@code
|
||||
using namespace mlpack;
|
||||
|
||||
const std::string& str = IO::GetParam<std::string>("string");
|
||||
const std::string& str = params.Has<std::string>("string");
|
||||
@endcode
|
||||
|
||||
Matrix types are accessed in the same way:
|
||||
|
||||
@code
|
||||
using namespace mlpack;
|
||||
|
||||
arma::mat& matrix = IO::GetParam<arma::mat>("matrix");
|
||||
arma::mat& matrix = params.Get<arma::mat>("matrix");
|
||||
@endcode
|
||||
|
||||
Similarly, model types can be accessed. If a @c LinearRegression model was
|
||||
@@ -786,9 +829,7 @@ specified by the user as the parameter @c model, the following code can access
|
||||
the model:
|
||||
|
||||
@code
|
||||
using namespace mlpack;
|
||||
|
||||
LinearRegression& lr = IO::GetParam<LinearRegression>("model");
|
||||
LinearRegression& lr = params.Get<LinearRegression>("model");
|
||||
@endcode
|
||||
|
||||
Matrices with categoricals are a little trickier to access since the C++
|
||||
@@ -801,72 +842,74 @@ parameter.
|
||||
using namespace mlpack;
|
||||
|
||||
typename std::tuple<data::DatasetInfo, arma::mat> TupleType;
|
||||
data::DatasetInfo& di = std::get<0>(IO::GetParam<TupleType>("matrix"));
|
||||
arma::mat& matrix = std::get<1>(IO::GetParam<TupleType>("matrix"));
|
||||
data::DatasetInfo& di = std::get<0>(params.Get<TupleType>("matrix"));
|
||||
arma::mat& matrix = std::get<1>(params.Get<TupleType>("matrix"));
|
||||
@endcode
|
||||
|
||||
These two functions can be used to write an entire program. The third function,
|
||||
@c GetPrintableParam(), can be used to help provide useful output in a program.
|
||||
Typically, this function should be used if you want to provide some kind of
|
||||
error message about a matrix or model parameter, but want to avoid printing the
|
||||
matrix itself. For instance, printing a matrix parameter with
|
||||
@c GetPrintableParam() will print the filename for a command-line binding or the
|
||||
size of a matrix for a Python binding. @c GetPrintableParam() for a model
|
||||
parameter will print the filename for the model for a command-line binding or
|
||||
a simple string representing the type of the model for a Python binding.
|
||||
@c params.GetPrintable(), can be used to help provide useful output in a
|
||||
program. Typically, this function should be used if you want to provide some
|
||||
kind of error message about a matrix or model parameter, but want to avoid
|
||||
printing the matrix itself. For instance, printing a matrix parameter with
|
||||
@c params.GetPrintable() will print the filename for a command-line binding or
|
||||
the size of a matrix for a Python binding. @c params.GetPrintable() for a model
|
||||
parameter will print the filename for the model for a command-line binding or a
|
||||
simple string representing the type of the model for a Python binding.
|
||||
|
||||
Putting all of these ideas together, here is the @c mlpackMain() function that
|
||||
could be created for the "random_numbers" program from earlier sections.
|
||||
Putting all of these ideas together, here is the @c BINDING_FUNCTION() function
|
||||
that could be created for the "random_numbers" program from earlier sections.
|
||||
|
||||
@code
|
||||
// BINDING_NAME should be defined here: ...
|
||||
|
||||
#include <mlpack/core/util/mlpack_main.hpp>
|
||||
|
||||
// BINDING_NAME(), BINDING_SHORT_DESC(), BINDING_LONG_DESC() , BINDING_EXAMPLE(),
|
||||
// BINDING_SEE_ALSO() and PARAM_*() definitions should go here:
|
||||
// ...
|
||||
// BINDING_USER_NAME(), BINDING_SHORT_DESC(), BINDING_LONG_DESC() ,
|
||||
// BINDING_EXAMPLE(), BINDING_SEE_ALSO() and PARAM_*() definitions should go
|
||||
// here: ...
|
||||
|
||||
using namespace mlpack;
|
||||
|
||||
void mlpackMain()
|
||||
void BINDING_FUNCTION(util::Params& params, util::Timers& timers)
|
||||
{
|
||||
// If the user passed an input matrix, tell them that we'll be ignoring it.
|
||||
if (IO::HasParam("input"))
|
||||
if (params.Has("input"))
|
||||
{
|
||||
// Print the filename the user passed, if a command-line binding, or the
|
||||
// size of the matrix passed, if a Python binding.
|
||||
Log::Warn << "The input matrix "
|
||||
<< IO::GetPrintableParam<arma::mat>("input") << " is ignored!"
|
||||
<< params.GetPrintable<arma::mat>("input") << " is ignored!"
|
||||
<< std::endl;
|
||||
}
|
||||
|
||||
// Get the number of samples and also the value we should subtract.
|
||||
const size_t numSamples = (size_t) IO::GetParam<int>("num_samples");
|
||||
const double subtractValue = IO::GetParam<double>("subtract");
|
||||
const size_t numSamples = (size_t) params.Get<int>("num_samples");
|
||||
const double subtractValue = params.Get<double>("subtract");
|
||||
|
||||
// Create the random matrix (1-dimensional).
|
||||
arma::mat output(1, numSamples, arma::fill::randu);
|
||||
output -= subtractValue;
|
||||
|
||||
// Save the output matrix if the user wants.
|
||||
if (IO::HasParam("output"))
|
||||
IO::GetParam<arma::mat>("output") = std::move(output); // Avoid copy.
|
||||
if (params.Has("output"))
|
||||
params.Get<arma::mat>("output") = std::move(output); // Avoid copy.
|
||||
|
||||
// Did the user request a random linear regression model?
|
||||
if (IO::HasParam("output_model"))
|
||||
if (params.Has("output_model"))
|
||||
{
|
||||
LinearRegression lr;
|
||||
lr.Parameters().randu(10); // 10-dimensional (arbitrary).
|
||||
lr.Lambda() = 0.0;
|
||||
lr.Intercept() = false; // No intercept term.
|
||||
|
||||
IO::GetParam<LinearRegression>("output_model") = std::move(lr);
|
||||
params.Get<LinearRegression>("output_model") = std::move(lr);
|
||||
}
|
||||
}
|
||||
@endcode
|
||||
|
||||
@subsection bindings_general_more More documentation on using IO
|
||||
@subsection bindings_general_more More documentation on using @c util::Params
|
||||
|
||||
More documentation for the IO module can either be found on the mlpack::IO
|
||||
More documentation for the IO module can either be found on the util::Params
|
||||
documentation page, or by reading the existing mlpack bindings. These can be
|
||||
found in the @c src/mlpack/methods/ folders, by finding the @c _main.cpp files.
|
||||
For instance, @c src/mlpack/methods/neighbor_search/knn_main.cpp is the
|
||||
@@ -874,14 +917,15 @@ k-nearest-neighbor search program definition.
|
||||
|
||||
@section bindings_structure Structure of IO module and associated macros
|
||||
|
||||
This section describes the internal functionality of the IO module and the
|
||||
associated macros. If you are only interested in writing mlpack programs, this
|
||||
section is probably not worth reading.
|
||||
This section describes the internal functionality of the IO module, which stores
|
||||
all known parameter sets, and the associated macros. If you are only interested
|
||||
in writing mlpack programs, this section is probably not worth reading.
|
||||
|
||||
There are eight main components involved with mlpack bindings:
|
||||
|
||||
- the IO module, a singleton class that stores parameter information
|
||||
- the mlpackMain() function that defines the functionality of the binding
|
||||
- the IO module, a thread-safe singleton class that stores parameter
|
||||
information
|
||||
- the BINDING_FUNCTION() function that defines the functionality of the binding
|
||||
- the BINDING_NAME() macro that defines the binding name
|
||||
- the BINDING_SHORT_DESC() macro that defines the short description
|
||||
- the BINDING_LONG_DESC() macro that defines the long description
|
||||
@@ -889,52 +933,68 @@ There are eight main components involved with mlpack bindings:
|
||||
- (optional) the BINDING_SEE_ALSO() macro that defines "see also" links
|
||||
- the PARAM_*() macros that define parameters for the binding
|
||||
|
||||
The mlpack::IO module is a singleton class that stores, at runtime, the binding
|
||||
name, the documentation, and the parameter information and values. In order to
|
||||
do this, each parameter and the program documentation must make themselves known
|
||||
to the IO singleton. This is accomplished by having the @c BINDING_NAME(),
|
||||
@c BINDING_SHORT_DESC(), @c BINDING_LONG_DESC(), @c BINDING_EXAMPLE(),
|
||||
@c BINDING_SEE_ALSO() and @c PARAM_*() macros declare global variables that,
|
||||
in their constructors, register themselves with the IO singleton.
|
||||
The @c mlpack::IO module is a singleton class that stores, at runtime, the
|
||||
binding name, the documentation, and the parameter information and values for
|
||||
any bindings available in the translation unit. When the binding is called, the
|
||||
@c mlpack::IO class instantiates a @c util::Params and @c util::Timers object,
|
||||
populating them with the correct options for the given binding, then calls
|
||||
@c BINDING_FUNCTION() with those instantiated objects.
|
||||
|
||||
The @c BINDING_NAME() macro declares an object of type mlpack::util::ProgramName.
|
||||
The @c BINDING_SHORT_DESC() macro declares an object of type
|
||||
mlpack::util::ShortDescription.
|
||||
The @c BINDING_LONG_DESC() macro declares an object of type
|
||||
mlpack::util::LongDescription.
|
||||
The @c BINDING_EXAMPLE() macro declares an object of type mlpack::util::Example.
|
||||
The @c BINDING_SEE_ALSO() macro declares an object of type
|
||||
mlpack::util::SeeAlso.
|
||||
The @c ProgramName class constructor calls IO::RegisterProgramName() in order to
|
||||
register the given program name.
|
||||
The @c ShortDescription class constructor calls IO::RegisterShortDescription() in order to
|
||||
register the given short description.
|
||||
The @c LongDescription class constructor calls IO::RegisterLongDescription() in order to
|
||||
register the given long description.
|
||||
The @c Example class constructor calls IO::RegisterExample() in order to
|
||||
register the given example.
|
||||
The @c SeeAlso class constructor calls IO::RegisterSeeAlso() in order to
|
||||
register the given see-also link.
|
||||
In order to do this, each parameter and the program documentation must make
|
||||
themselves known to the IO singleton. This is accomplished by having the @c
|
||||
BINDING_USER_NAME(), @c BINDING_SHORT_DESC(), @c BINDING_LONG_DESC(),
|
||||
@c BINDING_EXAMPLE(), @c BINDING_SEE_ALSO() and @c PARAM_*() macros declare
|
||||
global variables that, in their constructors, register themselves with the IO
|
||||
singleton.
|
||||
|
||||
* The @c BINDING_USER_NAME() macro declares an object of type
|
||||
@c mlpack::util::BindingName.
|
||||
* The @c BINDING_SHORT_DESC() macro declares an object of type
|
||||
@c mlpack::util::ShortDescription.
|
||||
* The @c BINDING_LONG_DESC() macro declares an object of type
|
||||
@c mlpack::util::LongDescription.
|
||||
* The @c BINDING_EXAMPLE() macro declares an object of type
|
||||
@c mlpack::util::Example.
|
||||
* The @c BINDING_SEE_ALSO() macro declares an object of type
|
||||
@c mlpack::util::SeeAlso.
|
||||
* The @c BindingName class constructor calls @c IO::AddBindingName() in order
|
||||
to register the given program name.
|
||||
* The @c ShortDescription class constructor calls @c IO::AddShortDescription()
|
||||
in order to register the given short description.
|
||||
* The @c LongDescription class constructor calls @c IO::AddLongDescription() in
|
||||
order to register the given long description.
|
||||
* The @c Example class constructor calls @c IO::AddExample() in order to
|
||||
register the given example.
|
||||
* The @c SeeAlso class constructor calls @c IO::AddSeeAlso() in order to
|
||||
register the given see-also link.
|
||||
|
||||
All of those macro calls use whatever the value of the @c BINDING_NAME macro is
|
||||
at the time of instantiation. This is why it is important that @c BINDING_NAME
|
||||
is set properly at the time @c mlpack_main.hpp is included and before any
|
||||
options are defined.
|
||||
|
||||
The @c PARAM_*() macros declare an object that will, in its constructor, call
|
||||
IO::Add() to register that parameter with the IO singleton. The specific type
|
||||
of that object will depend on the binding type being used.
|
||||
IO::Add() to register that parameter for the current binding (again specified by
|
||||
the @c BINDING_NAME macro's value) with the IO singleton. The specific type of
|
||||
that object will depend on the binding type being used.
|
||||
|
||||
The IO::Add() function takes an mlpack::util::ParamData object as its input.
|
||||
This @c ParamData object has a number of fields that must be set to properly
|
||||
describe the parameter. Each of the fields is documented and probably
|
||||
self-explanatory, but three fields deserve further explanation:
|
||||
The IO::AddParameter() function takes the name of the binding it is for and an
|
||||
mlpack::util::ParamData object as its input. This @c ParamData object has a
|
||||
number of fields that must be set to properly describe the parameter. Each of
|
||||
the fields is documented and probably self-explanatory, but three fields deserve
|
||||
further explanation:
|
||||
|
||||
- the <tt>std::string tname</tt> member is used to encode the true type of
|
||||
the parameter---which is not known by the IO singleton at runtime. This
|
||||
should be set to <tt>TYPENAME(T)</tt> where @c T is the type of the
|
||||
parameter.
|
||||
|
||||
- the <tt>boost::any value</tt> member is used to hold the actual value of the
|
||||
parameter. Typically this will simply be the parameter held by a
|
||||
@c boost::any object, but for some types it may be more complex. For
|
||||
instance, for a command-line matrix option, the @c value parameter will
|
||||
actually hold a tuple containing both the filename and the matrix itself.
|
||||
- the <tt>ANY value</tt> member (where <tt>ANY</tt> is whatever type was chosen
|
||||
in case <tt>std::any</tt> is not available) is used to hold the actual value
|
||||
of the parameter. Typically this will simply be the parameter held by a
|
||||
@c ANY object, but for some types it may be more complex. For instance, for
|
||||
a command-line matrix option, the @c value parameter will actually hold a
|
||||
tuple containing both the filename and the matrix itself.
|
||||
|
||||
- the <tt>std::string cppType</tt> should be a string containing the type as
|
||||
seen in C++ code. Typically this can be encoded by stringifying a
|
||||
@@ -945,11 +1005,13 @@ arguments into a fully specified @c ParamData object and then call IO::Add()
|
||||
with it.
|
||||
|
||||
With different binding types, different behavior is often required for the
|
||||
@c GetParam<T>(), @c HasParam(), and @c GetPrintableParam<T>() functions. In
|
||||
order to handle this, the IO singleton also holds a function pointer map, so
|
||||
@c params.Get<T>(), @c params.Has(), and @c params.GetPrintable<T>() functions.
|
||||
In order to handle this, the IO singleton also holds a function pointer map, so
|
||||
that a given type of option can call specific functionality for a certain task.
|
||||
This function map is accessible as @c IO::functionMap and is not meant to be
|
||||
used by users, but instead by people writing binding types.
|
||||
Given a @c util::Params object (which can be obtained with
|
||||
@c IO::Parameters("binding_name") ), this function map is accessible as
|
||||
@c params.functionMap, and is not meant to be used by users, but instead by
|
||||
people writing binding types.
|
||||
|
||||
Each function in the map must have signature
|
||||
|
||||
@@ -972,13 +1034,13 @@ and the first map key is the typename (<tt>tname</tt>) of the parameter, and the
|
||||
second map key is the string name of the function. For instance, calling
|
||||
|
||||
@code
|
||||
const util::ParamData& d = IO::Parameters()["param"];
|
||||
IO::GetSingleton().functionMap[d.tname]["GetParam"](d, input, output);
|
||||
const util::ParamData& d = params.Parameters()["param"];
|
||||
params.functionMap[d.tname]["GetParam"](d, input, output);
|
||||
@endcode
|
||||
|
||||
will call the @c GetParam() function for the type of the @c "param" parameter.
|
||||
Examples are probably easiest to understand how this functionality works; see
|
||||
the IO::GetParam<T>() source to see how this might be used.
|
||||
the @c params.Get<T>() source to see how this might be used.
|
||||
|
||||
The IO singleton expects the following functions to be defined in the function
|
||||
map for each type:
|
||||
@@ -996,12 +1058,12 @@ parts of the binding infrastructure for different languages.
|
||||
This section describes the internal functionality of the command-line program
|
||||
binding generator. If you are only interested in writing mlpack programs, this
|
||||
section probably is not worth reading. This section is worth reading only if
|
||||
you want to know the specifics of how the @c mlpackMain() function and macros
|
||||
get turned into a fully working command-line program.
|
||||
you want to know the specifics of how the @c BINDING_FUNCTION() function and
|
||||
macros get turned into a fully working command-line program.
|
||||
|
||||
The code for the command-line bindings is found in @c src/mlpack/bindings/cli.
|
||||
|
||||
@subsection bindings_cli_mlpack_main mlpackMain() definition
|
||||
@subsection bindings_cli_mlpack_main BINDING_FUNCTION() definition
|
||||
|
||||
Any command-line program must be compiled with the @c BINDING_TYPE macro
|
||||
set to the value @c BINDING_TYPE_CLI. This is handled by the CMake macro
|
||||
@@ -1030,31 +1092,40 @@ binding:
|
||||
@code
|
||||
int main(int argc, char** argv)
|
||||
{
|
||||
// Parse the command-line options; put them into IO.
|
||||
mlpack::bindings::cli::ParseCommandLine(argc, argv);
|
||||
// Parse the command-line options; put them into CLI.
|
||||
mlpack::util::Params params =
|
||||
mlpack::bindings::cli::ParseCommandLine(argc, argv);
|
||||
// Create a new timer object for this call.
|
||||
mlpack::util::Timers timers;
|
||||
timers.Enabled() = true;
|
||||
mlpack::Timer::EnableTiming();
|
||||
|
||||
mlpackMain();
|
||||
// A "total_time" timer is run by default for each mlpack program.
|
||||
timers.Start("total_time");
|
||||
BINDING_FUNCTION(params, timers);
|
||||
timers.Stop("total_time");
|
||||
|
||||
// Print output options, print verbose information, save model parameters,
|
||||
// clean up, and so forth.
|
||||
mlpack::bindings::cli::EndProgram();
|
||||
mlpack::bindings::cli::EndProgram(params, timers);
|
||||
}
|
||||
@endcode
|
||||
|
||||
Thus any mlpack command-line binding first processes the command-line arguments
|
||||
with mlpack::bindings::cli::ParseCommandLine(), then runs the binding with
|
||||
@c mlpackMain(), then cleans up with mlpack::bindings::cli::EndProgram().
|
||||
with @c mlpack::bindings::cli::ParseCommandLine(), then runs the binding with
|
||||
@c BINDING_FUNCTION(), then cleans up with
|
||||
@c mlpack::bindings::cli::EndProgram().
|
||||
|
||||
The @c ParseCommandLine() function reads the input parameters and sets the
|
||||
values in IO. For matrix-type and model-type parameters, this reads the
|
||||
filenames from the command-line, but does not load the matrix or model. Instead
|
||||
the matrix or model is loaded the first time it is accessed with
|
||||
@c GetParam<T>().
|
||||
@c params.Get<T>().
|
||||
|
||||
The @c \--help parameter is handled by the mlpack::bindings::cli::PrintHelp()
|
||||
function.
|
||||
|
||||
At the end of program execution, the mlpack::bindings::cli::EndProgram()
|
||||
At the end of program execution, the @c mlpack::bindings::cli::EndProgram()
|
||||
function is called. This writes any output matrix or model parameters to disk,
|
||||
and prints the program parameters and timers if @c \--verbose was given.
|
||||
|
||||
@@ -1065,19 +1136,19 @@ parameters all require special handling, since it is not possible to pass a
|
||||
matrix of any reasonable size or a model on the command line directly.
|
||||
Therefore for a matrix or model parameter, the user specifies the file
|
||||
containing that matrix or model parameter. If the parameter is an input
|
||||
parameter, then the file is loaded when @c GetParam<T>() is called. If the
|
||||
parameter, then the file is loaded when @c params.Get<T>() is called. If the
|
||||
parameter is an output parameter, then the matrix or model is saved to the file
|
||||
when @c EndProgram() is called.
|
||||
|
||||
The actual implementation of this is that the <tt>boost::any value</tt> member
|
||||
The actual implementation of this is that the <tt>ANY value</tt> member
|
||||
of the @c ParamData struct does not hold the model or the matrix, but instead a
|
||||
<tt>std::tuple</tt> containing both the matrix or the model, and the filename
|
||||
associated with that matrix or model.
|
||||
|
||||
This means that functions like @c GetParam<T>() and @c GetPrintableParam<T>()
|
||||
(and all of the other associated functions in the IO function map) must have
|
||||
special handling for matrix or model types. See those implementatipns for more
|
||||
details---the special handling is enforced via SFINAE.
|
||||
This means that functions like @c params.Get<T>() and
|
||||
@c params.GetPrintable<T>() (and all of the other associated functions in the
|
||||
function map) must have special handling for matrix or model types. See those
|
||||
implementations for more details---the special handling is enforced via SFINAE.
|
||||
|
||||
@subsection bindings_cli_parsing Parsing the command line
|
||||
|
||||
@@ -1122,7 +1193,7 @@ numpy matrices as input. Fortunately, numpy Cython bindings already exist,
|
||||
which make it easy to convert from a numpy object to an Armadillo object without
|
||||
copying any data. This code can be found in
|
||||
@c src/mlpack/bindings/python/mlpack/arma_numpy.pyx, and is used by the Python
|
||||
@c GetParam<T>() functionality.
|
||||
@c params.Get<T>() functionality.
|
||||
|
||||
mlpack also supports categorical matrices; in Python, the typical way of
|
||||
representing matrices with categorical features is with Pandas. Therefore,
|
||||
@@ -1147,7 +1218,7 @@ way (this example is taken from the perceptron binding):
|
||||
|
||||
@code
|
||||
cdef extern from "</home/ryan/src/mlpack-rc/src/mlpack/methods/perceptron/perceptron_main.cpp>" nogil:
|
||||
cdef int mlpackMain() nogil except +RuntimeError
|
||||
cdef int mlpack_perceptron(Params, Timers) nogil except +RuntimeError
|
||||
|
||||
cdef cppclass PerceptronModel:
|
||||
PerceptronModel() nogil
|
||||
@@ -1187,7 +1258,7 @@ individually if you like). The file
|
||||
the name of the program and the @c *_main.cpp file to include correctly, then
|
||||
the @c mlpack::bindings::python::PrintPYX() function is called by the program.
|
||||
The @c PrintPYX() function uses the parameters that have been set in the IO
|
||||
singleton by the @c BINDING_NAME(), @c BINDING_SHORT_DESC(),
|
||||
singleton by the @c BINDING_USER_NAME(), @c BINDING_SHORT_DESC(),
|
||||
@c BINDING_LONG_DESC(), @c BINDING_EXAMPLE(), @c BINDING_SEE_ALSO() and
|
||||
@c PARAM_*() macros in order to actually print a fully-working .pyx file that
|
||||
can be compiled. The file has several sections:
|
||||
@@ -1214,9 +1285,9 @@ binding.
|
||||
|
||||
@subsection bindings_python_testing Testing the Python bindings
|
||||
|
||||
We cannot do our tests only from the Boost Unit Test Framework in C++ because we
|
||||
need to see that we are able to load parameters properly from Python and return
|
||||
output correctly.
|
||||
In addition to the C++ tests we have implemented for each binding, we also have
|
||||
tests from Python that ensure that we can successfully transfer parameter values
|
||||
from Python to C++ and return output correctly.
|
||||
|
||||
The tests are in @c src/mlpack/bindings/python/tests/ and test both the actual
|
||||
bindings and also the auxiliary Python code included in
|
||||
|
||||
@@ -54,6 +54,8 @@ set(SOURCES
|
||||
hard_tanh_impl.hpp
|
||||
highway.hpp
|
||||
highway_impl.hpp
|
||||
instance_norm.hpp
|
||||
instance_norm_impl.hpp
|
||||
isrlu.hpp
|
||||
isrlu_impl.hpp
|
||||
join.hpp
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
/**
|
||||
* @file methods/ann/layer/instance_norm.hpp
|
||||
* @author Anjishnu Mukherjee
|
||||
* @author Shah Anwaar Khalid
|
||||
*
|
||||
* Definition of the Instance Normalization layer class.
|
||||
*
|
||||
* 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_ANN_LAYER_INSTANCE_NORM_HPP
|
||||
#define MLPACK_METHODS_ANN_LAYER_INSTANCE_NORM_HPP
|
||||
|
||||
#include <mlpack/prereqs.hpp>
|
||||
#include "layer_types.hpp"
|
||||
|
||||
namespace mlpack {
|
||||
namespace ann /** Artificial Neural Network. */ {
|
||||
|
||||
/**
|
||||
* Declaration of the Instance Normalization layer class. The layer transforms
|
||||
* the input data into zero mean and unit variance and then scales and shifts
|
||||
* the data by parameters, gamma and beta respectively. These parameters are
|
||||
* learnt by the network. The mean and standard-deviation are calculated
|
||||
* per-dimension separately for each object in a mini-batch.
|
||||
*
|
||||
* If deterministic is false (training), the mean and variance are calculated
|
||||
* and the data is normalized. If it is set to true (testing) then
|
||||
* the mean and variance accrued over the training set is used.
|
||||
*
|
||||
* For more information, refer to the following paper,
|
||||
*
|
||||
* @code
|
||||
* @article{Ulyanov17,
|
||||
* author = {Dmitry Ulyanov, Andrea Vedaldi and
|
||||
* Victor Lempitsky},
|
||||
* title = {Instance Normalization:
|
||||
* The Missing Ingredient for Fast Stylization},
|
||||
* year = {2017},
|
||||
* url = {https://arxiv.org/abs/1607.08022}
|
||||
* }
|
||||
* @endcode
|
||||
*
|
||||
* @tparam InputDataType Type of the input data (arma::colvec, arma::mat,
|
||||
* arma::sp_mat or arma::cube).
|
||||
* @tparam OutputDataType Type of the output data (arma::colvec, arma::mat,
|
||||
* arma::sp_mat or arma::cube).
|
||||
*/
|
||||
template <
|
||||
typename InputDataType = arma::mat,
|
||||
typename OutputDataType = arma::mat
|
||||
>
|
||||
class InstanceNorm
|
||||
{
|
||||
public:
|
||||
//! Create the InstanceNorm object.
|
||||
InstanceNorm();
|
||||
|
||||
/**
|
||||
* Create the InstanceNorm layer object with the specified parameters.
|
||||
*
|
||||
* @param size The number of input units / channels.
|
||||
* @param batchSize Size of the minibatch.
|
||||
* @param eps The epsilon added to variance to ensure numerical stability.
|
||||
* @param average Boolean to determine whether cumulative average is used for
|
||||
* updating the parameters or momentum.
|
||||
* @param momentum Parameter used to update the running mean and variance.
|
||||
*/
|
||||
InstanceNorm(const size_t size,
|
||||
const size_t batchSize,
|
||||
const double eps = 1e-5,
|
||||
const bool average = true,
|
||||
const double momentum = 0.1);
|
||||
|
||||
/**
|
||||
* Forward pass of the Instance Normalization layer. Transforms the input data
|
||||
* into zero mean and unit variance, scales the data by a factor gamma and
|
||||
* shifts it by beta.
|
||||
*
|
||||
* @param input Input data for the layer
|
||||
* @param output Resulting output activations.
|
||||
*/
|
||||
template<typename eT>
|
||||
void Forward(const arma::Mat<eT>& input, arma::Mat<eT>& output);
|
||||
|
||||
/**
|
||||
* Backward pass through the layer.
|
||||
*
|
||||
* @param input The input activations
|
||||
* @param gy The backpropagated error.
|
||||
* @param g The calculated gradient.
|
||||
*/
|
||||
template<typename eT>
|
||||
void Backward(const arma::Mat<eT>& input,
|
||||
const arma::Mat<eT>& gy,
|
||||
arma::Mat<eT>& g);
|
||||
|
||||
/**
|
||||
* Calculate the gradient using the output delta and the input activations.
|
||||
*
|
||||
* @param input The input activations
|
||||
* @param error The calculated error
|
||||
* @param gradient The calculated gradient.
|
||||
*/
|
||||
template<typename eT>
|
||||
void Gradient(const arma::Mat<eT>& input,
|
||||
const arma::Mat<eT>& error,
|
||||
arma::Mat<eT>& gradient);
|
||||
|
||||
//! Get the parameters.
|
||||
OutputDataType const& Parameters() const { return batchNorm.Parameters(); }
|
||||
//! Modify the parameters.
|
||||
OutputDataType& Parameters() { return batchNorm.Parameters(); }
|
||||
|
||||
//! Get the output parameter.
|
||||
OutputDataType const& OutputParameter() const
|
||||
{ return batchNorm.OutputParameter(); }
|
||||
|
||||
//! Modify the output parameter.
|
||||
OutputDataType& OutputParameter() { return batchNorm.OutputParameter(); }
|
||||
|
||||
//! Get the delta.
|
||||
OutputDataType const& Delta() const { return batchNorm.Delta(); }
|
||||
//! Modify the delta.
|
||||
OutputDataType& Delta() { return batchNorm.Delta(); }
|
||||
|
||||
//! Get the gradient.
|
||||
OutputDataType const& Gradient() const { return batchNorm.Gradient(); }
|
||||
//! Modify the gradient.
|
||||
OutputDataType& Gradient() { return batchNorm.Gradient(); }
|
||||
|
||||
//! Get the value of deterministic parameter.
|
||||
bool Deterministic() const { return deterministic; }
|
||||
//! Modify the value of deterministic parameter.
|
||||
bool& Deterministic() { return deterministic; }
|
||||
|
||||
//! Get the mean over the training data.
|
||||
OutputDataType const& TrainingMean() const { return runningMean; }
|
||||
//! Modify the mean over the training data.
|
||||
OutputDataType& TrainingMean() { return runningMean; }
|
||||
|
||||
//! Get the variance over the training data.
|
||||
OutputDataType const& TrainingVariance() const { return runningVariance; }
|
||||
//! Modify the variance over the training data.
|
||||
OutputDataType& TrainingVariance() { return runningVariance; }
|
||||
|
||||
//! Get the number of input units / channels.
|
||||
size_t InputSize() const { return size; }
|
||||
//! Modify the input units/ channels.
|
||||
size_t InputSize() {return size; }
|
||||
|
||||
//! Get the epsilon value.
|
||||
double Epsilon() const { return eps; }
|
||||
//! Modify the epsilon value.
|
||||
double Epsilon() { return eps; }
|
||||
|
||||
|
||||
//! Get the momentum value.
|
||||
double Momentum() const { return momentum; }
|
||||
//! Modify the momentum value.
|
||||
double Momentum() { return momentum; }
|
||||
|
||||
//! Get the average parameter.
|
||||
bool Average() const { return average; }
|
||||
//! Modify the average parameter.
|
||||
bool Average() { return average; }
|
||||
|
||||
//! Get the batchSize parameter.
|
||||
bool Batchsize() const { return batchSize; }
|
||||
//! Modify the batchSize parameter.
|
||||
bool Batchsize() { return batchSize; }
|
||||
|
||||
/**
|
||||
* Serialize the layer
|
||||
*/
|
||||
template<typename Archive>
|
||||
void serialize(Archive& ar, const uint32_t /* version */);
|
||||
|
||||
private:
|
||||
//! Locally stored BatchNorm Object.
|
||||
BatchNorm<InputDataType, OutputDataType> batchNorm;
|
||||
|
||||
//! Locally-stored reset parameter used to initialize the layer once.
|
||||
bool reset;
|
||||
|
||||
//! Locally-stored number of input units.
|
||||
size_t size;
|
||||
|
||||
//! Locally-stored epsilon value.
|
||||
double eps;
|
||||
|
||||
//! If true use average else use momentum for computing running mean
|
||||
//! and variance
|
||||
bool average;
|
||||
|
||||
//! Locally-stored value for momentum.
|
||||
double momentum;
|
||||
|
||||
//! Locally stored vale for numFunctions
|
||||
size_t batchSize;
|
||||
|
||||
/**
|
||||
* If true then mean and variance over the training set will be considered
|
||||
* instead of being calculated over the batch.
|
||||
*/
|
||||
bool deterministic;
|
||||
|
||||
//! Locally-stored mean object.
|
||||
OutputDataType runningMean;
|
||||
|
||||
//! Locally-stored variance object.
|
||||
OutputDataType runningVariance;
|
||||
}; // class InstanceNorm
|
||||
|
||||
} // namespace ann
|
||||
} // namespace mlpack
|
||||
|
||||
// Include the implementation.
|
||||
#include "instance_norm_impl.hpp"
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,152 @@
|
||||
/**
|
||||
* @file methods/ann/layer/instance_norm_impl.hpp
|
||||
* @author Anjishnu Mukherjee
|
||||
* @author Shah Anwaar Khalid
|
||||
*
|
||||
* Implementation of the Instance Normalization Layer.
|
||||
*
|
||||
* 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_ANN_LAYER_INSTANCE_NORM_IMPL_HPP
|
||||
#define MLPACK_METHODS_ANN_LAYER_INSTANCE_NORM_IMPL_HPP
|
||||
|
||||
// In case it is not included.
|
||||
#include "instance_norm.hpp"
|
||||
|
||||
namespace mlpack {
|
||||
namespace ann { /** Artificial Neural Network. */
|
||||
|
||||
template<typename InputDataType, typename OutputDataType>
|
||||
InstanceNorm<InputDataType, OutputDataType>::InstanceNorm() :
|
||||
size(0),
|
||||
eps(1e-8),
|
||||
average(true),
|
||||
momentum(0.0),
|
||||
deterministic(false),
|
||||
reset(false)
|
||||
{
|
||||
// Nothing to do here.
|
||||
}
|
||||
|
||||
template <typename InputDataType, typename OutputDataType>
|
||||
InstanceNorm<InputDataType, OutputDataType>::InstanceNorm(
|
||||
const size_t size,
|
||||
const size_t batchSize,
|
||||
const double eps,
|
||||
const bool average,
|
||||
const double momentum) :
|
||||
size(size),
|
||||
batchSize(batchSize),
|
||||
eps(eps),
|
||||
average(average),
|
||||
momentum(momentum),
|
||||
deterministic(false),
|
||||
reset(false)
|
||||
{
|
||||
batchNorm = ann::BatchNorm<> (size * batchSize,
|
||||
eps,
|
||||
average,
|
||||
momentum);
|
||||
runningMean.zeros(size, 1);
|
||||
runningVariance.ones(size, 1);
|
||||
runningVariance = batchNorm.TrainingVariance();
|
||||
}
|
||||
|
||||
template<typename InputDataType, typename OutputDataType>
|
||||
template<typename eT>
|
||||
void InstanceNorm<InputDataType, OutputDataType>::Forward(
|
||||
const arma::Mat<eT>& input,
|
||||
arma::Mat<eT>& output)
|
||||
{
|
||||
// Instance Norm with (N, C, H, W) is same as Batch Norm with (1, N*C, H, W),
|
||||
// where N is the batchSize, C is the number of channels, H and W are the
|
||||
// height and width of each image respectively.
|
||||
if (input.n_cols != batchSize)
|
||||
{
|
||||
Log::Fatal << "Must use the same BatchSize that was used in the constructor."
|
||||
<< std::endl;
|
||||
}
|
||||
|
||||
if (!reset)
|
||||
{
|
||||
batchNorm.Reset();
|
||||
reset = true;
|
||||
}
|
||||
|
||||
const size_t shapeA = input.n_rows;
|
||||
const size_t shapeB = input.n_cols;
|
||||
|
||||
if (deterministic)
|
||||
batchNorm.Deterministic() = true;
|
||||
|
||||
arma::mat inputTemp(const_cast<arma::Mat<eT>&>(input).memptr(),
|
||||
shapeA * shapeB, 1, false, false);
|
||||
batchNorm.Forward(inputTemp, output);
|
||||
output.reshape(shapeA, shapeB);
|
||||
runningMean = batchNorm.TrainingMean();
|
||||
runningMean.reshape(size, shapeB);
|
||||
runningMean = arma::mean(runningMean, 1);
|
||||
runningVariance = batchNorm.TrainingVariance();
|
||||
runningVariance.reshape(size, shapeB);
|
||||
runningVariance = arma::mean(runningVariance, 1);
|
||||
}
|
||||
|
||||
template<typename InputDataType, typename OutputDataType>
|
||||
template<typename eT>
|
||||
void InstanceNorm<InputDataType, OutputDataType>::Backward(
|
||||
const arma::Mat<eT>& input,
|
||||
const arma::Mat<eT>& gy,
|
||||
arma::Mat<eT>& g)
|
||||
{
|
||||
const size_t shapeA = input.n_rows;
|
||||
const size_t shapeB = input.n_cols;
|
||||
|
||||
arma::mat inputTemp(const_cast<arma::Mat<eT>&>(input).memptr(),
|
||||
shapeA * shapeB, 1, false, false);
|
||||
arma::mat gyTemp(const_cast<arma::Mat<eT>&>(gy).memptr(),
|
||||
shapeA * shapeB, 1, false, false);
|
||||
batchNorm.Backward(inputTemp, gyTemp, g);
|
||||
g.reshape(shapeA, shapeB);
|
||||
}
|
||||
|
||||
template<typename InputDataType, typename OutputDataType>
|
||||
template<typename eT>
|
||||
void InstanceNorm<InputDataType, OutputDataType>::Gradient(
|
||||
const arma::Mat<eT>& input,
|
||||
const arma::Mat<eT>& error,
|
||||
arma::Mat<eT>& gradient)
|
||||
{
|
||||
const size_t shapeA = input.n_rows;
|
||||
const size_t shapeB = input.n_cols;
|
||||
|
||||
arma::mat inputTemp(const_cast<arma::Mat<eT>&>(input).memptr(),
|
||||
shapeA * shapeB, 1, false, false);
|
||||
arma::mat errorTemp(const_cast<arma::Mat<eT>&>(error).memptr(),
|
||||
shapeA * shapeB, 1, false, false);
|
||||
batchNorm.Gradient(inputTemp, errorTemp, gradient);
|
||||
}
|
||||
|
||||
template<typename InputDataType, typename OutputDataType>
|
||||
template<typename Archive>
|
||||
void InstanceNorm<InputDataType, OutputDataType>::serialize(
|
||||
Archive& ar, const uint32_t /* version */)
|
||||
{
|
||||
ar(CEREAL_NVP(size));
|
||||
ar(CEREAL_NVP(eps));
|
||||
ar(CEREAL_NVP(average));
|
||||
ar(CEREAL_NVP(momentum));
|
||||
ar(CEREAL_NVP(deterministic));
|
||||
ar(CEREAL_NVP(runningMean));
|
||||
ar(CEREAL_NVP(runningVariance));
|
||||
ar(CEREAL_NVP(reset));
|
||||
ar(CEREAL_NVP(batchNorm));
|
||||
}
|
||||
|
||||
} // namespace ann
|
||||
} // namespace mlpack
|
||||
|
||||
#endif
|
||||
@@ -41,6 +41,7 @@
|
||||
//#include <mlpack/methods/ann/layer/hardshrink.hpp>
|
||||
//#include <mlpack/methods/ann/layer/hard_tanh.hpp>
|
||||
#include <mlpack/methods/ann/layer/highway.hpp>
|
||||
//#include <mlpack/methods/ann/layer/instance_norm.hpp>
|
||||
//#include <mlpack/methods/ann/layer/join.hpp>
|
||||
//#include <mlpack/methods/ann/layer/layer_norm.hpp>
|
||||
#include <mlpack/methods/ann/layer/leaky_relu.hpp>
|
||||
|
||||
@@ -6172,3 +6172,203 @@ TEST_CASE("GradientMultiheadAttentionTest", "[ANNLayerTest]")
|
||||
REQUIRE(CheckGradient(function) <= 3e-06);
|
||||
}
|
||||
*/
|
||||
|
||||
/**
|
||||
* Simple tests for instance normalization layer.
|
||||
*
|
||||
TEST_CASE("InstanceNormLayerTest", "[ANNLayerTest]")
|
||||
{
|
||||
arma::mat input, result, output, delta, deltaExpected;
|
||||
arma::mat runningMean, runningVar;
|
||||
|
||||
// Represents 2 images, each having 3 channels, and shape (3,2).
|
||||
input << 1 << 19 << arma::endr
|
||||
<< 2 << 20 << arma::endr
|
||||
<< 3 << 21 << arma::endr
|
||||
<< 4 << 22 << arma::endr
|
||||
<< 5 << 23 << arma::endr
|
||||
<< 6 << 24 << arma::endr
|
||||
<< 7 << 25 << arma::endr
|
||||
<< 8 << 26 << arma::endr
|
||||
<< 9 << 27 << arma::endr
|
||||
<< 10 << 28 << arma::endr
|
||||
<< 11 << 29 << arma::endr
|
||||
<< 12 << 30 << arma::endr
|
||||
<< 13 << 31 << arma::endr
|
||||
<< 14 << 32 << arma::endr
|
||||
<< 15 << 33 << arma::endr
|
||||
<< 16 << 34 << arma::endr
|
||||
<< 17 << 35 << arma::endr
|
||||
<< 18 << 36 << arma::endr;
|
||||
|
||||
// Output calculated using torch.nn.InstanceNorm2d().
|
||||
result << -1.4638 << -1.4638 << arma::endr
|
||||
<< -0.8783 << -0.8783 << arma::endr
|
||||
<< -0.2928 << -0.2928 << arma::endr
|
||||
<< 0.2928 << 0.2928 << arma::endr
|
||||
<< 0.8783 << 0.8783 << arma::endr
|
||||
<< 1.4638 << 1.4638 << arma::endr
|
||||
<< -1.4638 << -1.4638 << arma::endr
|
||||
<< -0.8783 << -0.8783 << arma::endr
|
||||
<< -0.2928 << -0.2928 << arma::endr
|
||||
<< 0.2928 << 0.2928 << arma::endr
|
||||
<< 0.8783 << 0.8783 << arma::endr
|
||||
<< 1.4638 << 1.4638 << arma::endr
|
||||
<< -1.4638 << -1.4638 << arma::endr
|
||||
<< -0.8783 << -0.8783 << arma::endr
|
||||
<< -0.2928 << -0.2928 << arma::endr
|
||||
<< 0.2928 << 0.2928 << arma::endr
|
||||
<< 0.8783 << 0.8783 << arma::endr
|
||||
<< 1.4638 << 1.4638 << arma::endr;
|
||||
|
||||
// Calculated using torch.nn.InstanceNorm2d().
|
||||
deltaExpected << 1.8367 << 1.8367 << arma::endr
|
||||
<< 0.3967 << 0.3967 << arma::endr
|
||||
<< 0.0147 << 0.0147 << arma::endr
|
||||
<<-0.0147 << -0.0147 << arma::endr
|
||||
<<-0.3967 << -0.3967 << arma::endr
|
||||
<<-1.8367 << -1.8367 << arma::endr
|
||||
<< 1.8367 << 1.8367 << arma::endr
|
||||
<< 0.3967 << 0.3967 << arma::endr
|
||||
<< 0.0147 << 0.0147 << arma::endr
|
||||
<<-0.0147 << -0.0147 << arma::endr
|
||||
<<-0.3967 << -0.3967 << arma::endr
|
||||
<<-1.8367 << -1.8367 << arma::endr
|
||||
<< 1.8367 << 1.8367 << arma::endr
|
||||
<< 0.3967 << 0.3967 << arma::endr
|
||||
<< 0.0147 << 0.0147 << arma::endr
|
||||
<<-0.0147 << -0.0147 << arma::endr
|
||||
<<-0.3967 << -0.3967 << arma::endr
|
||||
<<-1.8367 << -1.8367 << arma::endr;
|
||||
|
||||
// Check Forward and Backward pass in non-deterministic mode.
|
||||
InstanceNorm<> module(3, input.n_cols, 1e-5, false, 0.1);
|
||||
output.zeros(arma::size(input));
|
||||
module.Forward(input, output);
|
||||
CheckMatrices(output, result, 1e-1);
|
||||
|
||||
module.Backward(input, output, delta);
|
||||
CheckMatrices(delta, deltaExpected, 1e-1);
|
||||
|
||||
runningMean = arma::mat(3, 1);
|
||||
runningVar = arma::mat(3, 1);
|
||||
runningMean(0) = 1.2500;
|
||||
runningMean(1) = 1.8500;
|
||||
runningMean(2) = 2.4500;
|
||||
runningVar(0) = 1.2500;
|
||||
runningVar(1) = 1.2500;
|
||||
runningVar(2) = 1.2500;
|
||||
|
||||
CheckMatrices(runningMean, module.TrainingMean(), 1e-1);
|
||||
CheckMatrices(runningVar, module.TrainingVariance(), 1e-1);
|
||||
|
||||
// Check Forward pass in deterministic mode.
|
||||
InstanceNorm<> module1(3, input.n_cols, 1e-5, false, 0.1);
|
||||
module1.Deterministic() = true;
|
||||
output.zeros(arma::size(input));
|
||||
module1.Forward(input, output);
|
||||
|
||||
// Calculated using torch.nn.InstanceNorm2d().
|
||||
result << 1.0000 << 18.9999 << arma::endr
|
||||
<< 2.0000 << 19.9999 << arma::endr
|
||||
<< 3.0000 << 20.9999 << arma::endr
|
||||
<< 4.0000 << 21.9999 << arma::endr
|
||||
<< 5.0000 << 22.9999 << arma::endr
|
||||
<< 6.0000 << 23.9999 << arma::endr
|
||||
<< 7.0000 << 24.9999 << arma::endr
|
||||
<< 8.0000 << 25.9999 << arma::endr
|
||||
<< 9.0000 << 26.9999 << arma::endr
|
||||
<< 10.0000 << 27.9999 << arma::endr
|
||||
<< 10.9999 << 28.9999 << arma::endr
|
||||
<< 11.9999 << 29.9999 << arma::endr
|
||||
<< 12.9999 << 30.9998 << arma::endr
|
||||
<< 13.9999 << 31.9998 << arma::endr
|
||||
<< 14.9999 << 32.9998 << arma::endr
|
||||
<< 15.9999 << 33.9998 << arma::endr
|
||||
<< 16.9999 << 34.9998 << arma::endr
|
||||
<< 17.9999 << 35.9998 << arma::endr;
|
||||
|
||||
CheckMatrices(output, result, 1e-1);
|
||||
}
|
||||
*/
|
||||
|
||||
/**
|
||||
* Test that the functions that can access the parameters of the
|
||||
* Instance Norm layer work.
|
||||
*
|
||||
TEST_CASE("InstanceNormLayerParametersTest", "[ANNLayerTest]")
|
||||
{
|
||||
// Parameter order : size, eps.
|
||||
InstanceNorm<> layer(7, 0, 1e-3);
|
||||
|
||||
// Make sure we can get the parameters successfully.
|
||||
REQUIRE(layer.InputSize() == 7);
|
||||
REQUIRE(layer.Epsilon() == 1e-3);
|
||||
|
||||
arma::mat runningMean(7, 1, arma::fill::randn);
|
||||
arma::mat runningVariance(7, 1, arma::fill::randn);
|
||||
|
||||
layer.TrainingVariance() = runningVariance;
|
||||
layer.TrainingMean() = runningMean;
|
||||
CheckMatrices(layer.TrainingVariance(), runningVariance);
|
||||
CheckMatrices(layer.TrainingMean(), runningMean);
|
||||
}
|
||||
*/
|
||||
|
||||
/**
|
||||
* Instance Norm layer numerical gradient test.
|
||||
*
|
||||
TEST_CASE("GradientInstanceNormLayerTest", "[ANNLayerTest]")
|
||||
{
|
||||
// Add function gradient instantiation.
|
||||
// To make this test robust, check it ten times.
|
||||
bool pass = false;
|
||||
for (size_t trial = 0; trial < 10; trial++)
|
||||
{
|
||||
struct GradientFunction
|
||||
{
|
||||
GradientFunction()
|
||||
{
|
||||
input = arma::randn(16, 1024);
|
||||
arma::mat target;
|
||||
target.ones(1, 1024);
|
||||
|
||||
model = new FFN<NegativeLogLikelihood<>, NguyenWidrowInitialization>();
|
||||
model->Predictors() = input;
|
||||
model->Responses() = target;
|
||||
model->Add<IdentityLayer<> >();
|
||||
model->Add<Convolution<> >(1, 2, 3, 3, 1, 1, 0, 0, 4, 4);
|
||||
model->Add<InstanceNorm<> > (2, 1024);
|
||||
model->Add<Linear<> >(2 * 2 * 2, 2);
|
||||
model->Add<LogSoftMax<> >();
|
||||
}
|
||||
|
||||
~GradientFunction()
|
||||
{
|
||||
delete model;
|
||||
}
|
||||
|
||||
double Gradient(arma::mat& gradient) const
|
||||
{
|
||||
double error = model->Evaluate(model->Parameters(), 0, 1024, false);
|
||||
model->Gradient(model->Parameters(), 0, gradient, 1024);
|
||||
return error;
|
||||
}
|
||||
|
||||
arma::mat& Parameters() { return model->Parameters(); }
|
||||
|
||||
FFN<NegativeLogLikelihood<>, NguyenWidrowInitialization>* model;
|
||||
arma::mat input, target;
|
||||
} function;
|
||||
|
||||
double gradient = CheckGradient(function);
|
||||
if (gradient < 1e-1)
|
||||
{
|
||||
pass = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
REQUIRE(pass);
|
||||
}
|
||||
*/
|
||||
|
||||
Reference in New Issue
Block a user