diff --git a/src/mlpack/methods/ann/layer/CMakeLists.txt b/src/mlpack/methods/ann/layer/CMakeLists.txt index 58065d9189..c96c48799c 100644 --- a/src/mlpack/methods/ann/layer/CMakeLists.txt +++ b/src/mlpack/methods/ann/layer/CMakeLists.txt @@ -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 diff --git a/src/mlpack/methods/ann/layer/instance_norm.hpp b/src/mlpack/methods/ann/layer/instance_norm.hpp new file mode 100644 index 0000000000..1a4dce2926 --- /dev/null +++ b/src/mlpack/methods/ann/layer/instance_norm.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 +#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 + void Forward(const arma::Mat& input, arma::Mat& output); + + /** + * Backward pass through the layer. + * + * @param input The input activations + * @param gy The backpropagated error. + * @param g The calculated gradient. + */ + template + void Backward(const arma::Mat& input, + const arma::Mat& gy, + arma::Mat& 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 + void Gradient(const arma::Mat& input, + const arma::Mat& error, + arma::Mat& 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 + void serialize(Archive& ar, const uint32_t /* version */); + + private: + //! Locally stored BatchNorm Object. + BatchNorm 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 diff --git a/src/mlpack/methods/ann/layer/instance_norm_impl.hpp b/src/mlpack/methods/ann/layer/instance_norm_impl.hpp new file mode 100644 index 0000000000..352726239d --- /dev/null +++ b/src/mlpack/methods/ann/layer/instance_norm_impl.hpp @@ -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 +InstanceNorm::InstanceNorm() : + size(0), + eps(1e-8), + average(true), + momentum(0.0), + deterministic(false), + reset(false) +{ + // Nothing to do here. +} + +template +InstanceNorm::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 +template +void InstanceNorm::Forward( + const arma::Mat& input, + arma::Mat& 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&>(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 +template +void InstanceNorm::Backward( + const arma::Mat& input, + const arma::Mat& gy, + arma::Mat& g) +{ + const size_t shapeA = input.n_rows; + const size_t shapeB = input.n_cols; + + arma::mat inputTemp(const_cast&>(input).memptr(), + shapeA * shapeB, 1, false, false); + arma::mat gyTemp(const_cast&>(gy).memptr(), + shapeA * shapeB, 1, false, false); + batchNorm.Backward(inputTemp, gyTemp, g); + g.reshape(shapeA, shapeB); +} + +template +template +void InstanceNorm::Gradient( + const arma::Mat& input, + const arma::Mat& error, + arma::Mat& gradient) +{ + const size_t shapeA = input.n_rows; + const size_t shapeB = input.n_cols; + + arma::mat inputTemp(const_cast&>(input).memptr(), + shapeA * shapeB, 1, false, false); + arma::mat errorTemp(const_cast&>(error).memptr(), + shapeA * shapeB, 1, false, false); + batchNorm.Gradient(inputTemp, errorTemp, gradient); +} + +template +template +void InstanceNorm::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 diff --git a/src/mlpack/methods/ann/layer/layer.hpp b/src/mlpack/methods/ann/layer/layer.hpp index a50766cc73..8aca0ad270 100644 --- a/src/mlpack/methods/ann/layer/layer.hpp +++ b/src/mlpack/methods/ann/layer/layer.hpp @@ -40,6 +40,7 @@ #include "hard_tanh.hpp" #include "hardshrink.hpp" #include "highway.hpp" +#include "instance_norm.hpp" #include "join.hpp" #include "layer_norm.hpp" #include "layer_types.hpp" diff --git a/src/mlpack/methods/ann/layer/layer_types.hpp b/src/mlpack/methods/ann/layer/layer_types.hpp index ec7fd20e84..a1970bd6e1 100644 --- a/src/mlpack/methods/ann/layer/layer_types.hpp +++ b/src/mlpack/methods/ann/layer/layer_types.hpp @@ -27,6 +27,7 @@ #include #include #include +#include #include #include #include @@ -251,7 +252,8 @@ using MoreTypes = boost::variant< ISRLU*, BicubicInterpolation*, NearestInterpolation*, - GroupNorm* + GroupNorm*, + InstanceNorm* >; template diff --git a/src/mlpack/tests/ann_layer_test.cpp b/src/mlpack/tests/ann_layer_test.cpp index bc67ecf686..b9d5452c41 100644 --- a/src/mlpack/tests/ann_layer_test.cpp +++ b/src/mlpack/tests/ann_layer_test.cpp @@ -5408,3 +5408,200 @@ 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, NguyenWidrowInitialization>(); + model->Predictors() = input; + model->Responses() = target; + model->Add >(); + model->Add >(1, 2, 3, 3, 1, 1, 0, 0, 4, 4); + model->Add > (2, 1024); + model->Add >(2 * 2 * 2, 2); + model->Add >(); + } + + ~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, NguyenWidrowInitialization>* model; + arma::mat input, target; + } function; + + double gradient = CheckGradient(function); + if (gradient < 1e-1) + { + pass = true; + break; + } + } + + REQUIRE(pass); +}