diff --git a/src/mlpack/methods/ann/layer/layer_norm.hpp b/src/mlpack/methods/ann/layer/layer_norm.hpp index 41afac3be9..de182b0d0b 100644 --- a/src/mlpack/methods/ann/layer/layer_norm.hpp +++ b/src/mlpack/methods/ann/layer/layer_norm.hpp @@ -184,6 +184,9 @@ class LayerNorm //! Locally-stored normalized input. OutputDataType normalized; + + //! Locally-stored input with 0 mean. + OutputDataType inputMean; }; // class LayerNorm } // namespace ann diff --git a/src/mlpack/methods/ann/layer/layer_norm_impl.hpp b/src/mlpack/methods/ann/layer/layer_norm_impl.hpp index f0d94e1cf6..bf12916880 100644 --- a/src/mlpack/methods/ann/layer/layer_norm_impl.hpp +++ b/src/mlpack/methods/ann/layer/layer_norm_impl.hpp @@ -63,7 +63,7 @@ void LayerNorm::Forward( // Normalize the input. output = input.each_row() - mean; - + inputMean = output; output.each_row() /= arma::sqrt(variance + eps); // Reused in the backward and gradient step. @@ -79,7 +79,6 @@ template void LayerNorm::Backward( const arma::Mat&& input, arma::Mat&& gy, arma::Mat&& g) { - const arma::mat inputMean = input.each_row() - mean; const arma::mat stdInv = 1.0 / arma::sqrt(variance + eps); // dl / dxhat