Norm matters: efficient and accurate normalization schemes in deep networks
Elad Hoffer, Ron Banner, Itay Golan, Daniel Soudry
Introduction
Deep neural networks are known to benefit from normalization between consecutive layers. This was made noticeable with the introduction of Batch-Normalization (BN) , which normalizes the output of each layer to have zero mean and unit variance for each channel across the training batch. This idea was later developed to act across channels instead of the batch dimension in Layer-normalization and improved in certain tasks with methods such as Batch-Renormalization , Instance-normalization and Group-Normalization . In addition, normalization methods are also applied to the layer parameters instead of their outputs. Methods such as Weight-Normalization , and Normalization-Propagation targeted the layer weights by normalizing their per-channel norm to have a fixed value. Instead of explicit normalization, effort was also made to enable self-normalization by adapting activation function so that intermediate activations will converge towards zero-mean and unit variance .
Batch-normalization, despite its merits, suffers from several issues, as pointed out by previous work . These issues are not yet solved in current normalization methods.
Batch normalization typically improves generalization performance and is therefore considered a regularization mechanism. Other regularization mechanisms are typically used in conjunction. For example, weight decay, also known as regularization, is a common method which adds a penalty proportional to the weights’ norm. Weight decay was proven to improve generalization in various problems , but, so far, not for non-linear deep neural networks. There, performed an extensive set of experiments on regularization and concluded that explicit regularization, such as weight decay, may improve generalization performance, but is neither necessary nor, by itself, sufficient for reducing generalization error. Therefore, it is not clear how weight decay interacts with BN, or if weight decay is even really necessary given that batch norm already constrains the output norms ).
A key assumption in BN is the independence between samples appearing in each batch. While this assumption seems to hold for most convolutional networks used to classify images in conventional datasets, it falls short when employed in domains with strong correlations between samples, such as time-series prediction, reinforcement learning, and generative modeling. For example, BN requires modifications to work in recurrent networks , for which alternatives such as weight-normalization and layer-normalization were explicitly devised, without reaching the success and wide adoption of BN. Another example is Generative adversarial networks, which are also noted to suffer from the common form of BN. GAN training with BN proved unstable in some cases, decreasing the quality of the trained model . Instead, it was replaced with virtual-BN , weight-norm and spectral normalization . Also, BN may be harmful even in plain classification tasks, when using unbalanced classes, or correlated instances. In addition, while BN is defined for the training phase of the models, it requires a running estimate for the evaluation phase – causing a noticeable difference between the two . This shortcoming was addressed later by batch-renormalization , yet still requiring the original BN at the early steps of training.
From the computational perspective, BN is significant in modern neural networks, as it requires several floating point operations across the activation of the entire batch for every layer in the network. Previous analysis by Gitman & Ginsburg measured BN to constitute up to of the computation time needed for the entire model. It is also not easily parallelized, as it is usually memory-bound on currently employed hardware. In addition, the operation requires saving the pre-normalized activations for back-propagation in the general case , thus using roughly twice the memory as a non-BN network in the training phase. Other methods, such as Weight-Normalization have a much smaller computational cost but typically achieve significantly lower accuracy when used in large-scale tasks such as ImageNet .
As the use of deep learning continues to evolve, the interest in low-precision training and inference increases . Optimized hardware was designed to leverage benefits of low-precision arithmetic and memory operations, with the promise of better, more efficient implementations . Although most mathematical operations employed in neural-networks are known to be robust to low-precision and quantized values, the current normalization methods are notably not suited for these cases. As far as we know, this has remained an unanswered issue, with no suggested alternatives. Specifically, all normalization methods, including BN, use an normalization (variance computation) to control the activation scale for each layer. The operation requires a sum of power-of-two floating point variables, a square-root function, and a reciprocal operation. All of these require both high-precision to avoid zero variance, and a large range to avoid overflow when adding large numbers. This makes BN an operation that is not easily adapted to low-precision implementations. Using norm spaces other than can alleviate these problems, as we shall see later.
2 Contributions
In this paper we make the following contributions, to address the issues explained in the previous section:
We find the mechanism through which weight decay before BN affects learning dynamics: we demonstrate that by adjusting the learning rate or normalization method we can exactly mimic the effect of weight decay on the learning dynamics. We suggest this happens since certain normalization methods, such as a BN, disentangle the effect of weight vector norm on the following activation layers.
We show that we can replace the standard BN with certain and based variations of BN, which do not harm accuracy (on CIFAR and ImageNet) and even somewhat improve training speed. Importantly, we demonstrate that such norms can work well with low precision (16bit), while does not. Notably, for these normalization schemes to work well, precise scale adjustment is required, which can be approximated analytically.
We show that by bounding the norm in a weight-normalization scheme, we can significantly improve its performance in convnets (on ImageNet), and improve baseline performance in LSTMs (on WMT14 de-en). This method can alleviate several task-specific limitations of BN, and reduce its computational and memory costs (e.g., allowing to work with significantly larger batch sizes). Importantly, for the method to work well, we need to carefully choose the scale of the weights using the scale of the initialization.
Together, these findings emphasize that the learning dynamics in neural networks are very sensitive to the norms of the weights. Therefore, it is an important goal for future research to search for precise and theoretically justifiable methods to adjust the scale for these norms.
Consequences of the scale invariance of Batch-Normalization
When BN is applied after a linear layer, it is well known that the output is invariant to the channel weight vector norm. Specifically, denoting a channel weight vector with and , channel input as and for batch-norm, we have
This invariance to the weight vector norm means that a BN applied after a layer renders its norm irrelevant to the inputs of consecutive layers. The same can be easily shown for the per-channel weights of a convolutional layer. The gradient in such case is scaled by :
When a layer is rescaling invariant, the key feature of the weight vector is its direction.
During training, the weights are typically incremented through some variant of stochastic gradient descent, according to the gradient of the loss at mini-batch , with learning rate
Claim. During training, the weight direction is updated according to
Proof. Denote . Note that, from eqs. 2 and 3 we have
Therefore, the step size of the weight direction is approximately proportional to
in the case of linear layer followed by BN, and for small learning rate . Note that a similar conclusion was reached by van Laarhoven , who implicitly assumed , though this is only approximately true. Here we show this conclusion is still true without such an assumption. This analysis continues to hold for non-linear functions that do not affect scale, such as the commonly used ReLU function. In addition, although stated for the case of vanilla SGD, similar argument can be made for adaptive methods such as Adagrad or Adam .
Connection between weight-decay, learning rate and normalization
We claim that when using batch-norm (BN), weight decay (WD) improves optimization only by fixing the norm to a small range of values, leading to a more stable step size for the weight direction (“effective step size”). Fixing the norm allows better control over the effective step size through the learning rate . Without WD, the norm grows unbounded , resulting in a decreased effective step size, although the learning rate hyper-parameter remains unchanged.
We show empirically that the accuracy gained by using WD can be achieved without it, only by adjusting the learning rate. Given statistics on norms of each channel from a training with WD and BN, similar results can be achieved without WD by mimicking the effective step size using the following correction on the learning rate:
where is the weights’ vector of a single channel, and is the weights’ vector of the corresponding channel in a training with WD. This correction requires access to the norms of a training with WD, hence it is not a practical method to replace WD but just a tool to demonstrate our claim on the connection between weights’ norm, WD and step size.
We conducted multiple experiments on CIFAR-10 to show this connection. Figure 1 reports the test accuracy during the training of all experiments. We were able to show that WD results can be mimicked with step size adjustments using the correction formula from Eq. 5. In another experiment, we replaced the learning rate scheduling with norm scheduling. To do so, after every gradient descent step we normalized the norm of each convolution layer channel to be the same as the norm of the corresponding channel in training with WD and keep the learning rate constant. When learning rate is multiplied by 0.1 in the WD training, we instead multiply the norm by , leading to an effective step size of . As expected, when applying the correction on step-size or replacing learning rate scheduling with norm scheduling, the accuracy is similar to the training with WD throughout the learning process, suggesting that WD affects the training process only indirectly, by modulating the learning rate. Implementation details appear in supplementary material.
We suggested above that the main function of BN is to neutralize the effect of the preceding layer’s weights. If this hypothesis is true, then other operations might be able to replace BN, as long as they remain similarly scale invariant (as in eq. (1)) — and if we keep the same scale as BN. Following this reasoning, we next aim to replace the use of norm with scale-invariant alternatives which are more appealing computationally and for low-precision implementations.
Batch normalization aims at regularizing the input so that sum of deviations from the mean would be standardized according to the Euclidean norm metric. For a layer with dimensional input , batch norm normalizes each dimension
In this section, we suggest alternative metrics for BN. We focus on the and due to their appealing speed and memory computations. In our simulations, we were able to train models faster and with fewer GPUs using the above normalizations. Strikingly, by proper adjustments of these normalizations, we were able to train various complicated models without hurting the classification performance. We begin with the -norm metric.
For a layer with dimensional input , batch normalization normalize each dimension
where is the expectation over , is the batch size and is a normalization term.
Unlike traditional batch normalization that computes the average squared deviation from the mean (variance), batch normalization computes only the average absolute deviation from the mean. This has two major advantages. First, batch normalization eliminates the computational efforts required for the square and square root operations. Second, as the square of an -bit number is generally of bits, the absence of these square computations makes it much more suitable for low-precision training that has been recognized to drastically reduce memory size and power consumption on dedicated deep learning hardware .
As can be seen in equation 7, the batch normalization quantifies the variability with the normalized average absolute deviation . To calculate an appropriate value for the constant , we assume the input follows Gaussian distribution . This is a common approximation (e.g., Soudry et al. ), based on the fact that the neural input is a sum of many inputs, so we expect it to be approximately Gaussian from the central limit theorem. In this case, follows the distribution . Therefore, for each example it holds that follows a half-normal distribution with expectation . Accordingly, the expected variability measure is related to the traditional standard deviation measure normally used with batch normalization as follows:
Figure 3 presents the validation accuracy of ResNet-18 and ResNet-50 on ImageNet using and batch norms. While the use of batch norm is more efficient in terms of resource usage, power, and speed, they both share the same classification accuracy. We additionally verified layer-normalization to work on Transformer architecture . Using an layer-norm we achieved a final perplexity of 5.2 vs. 5.1 for original layer-norm using the base model on the WMT14 dataset.
We note the importance of to the performance of normalization method. For example, using helps the network to reach 20% validation error more than twice faster than an equivalent configuration without this normalization term. With the network converges at the same rate and to the same accuracy as batch norm. It is somewhat surprising that this constant can have such an impact on performance, considering the fact that it is so close to one (). A demonstration of this effect can be found in the supplementary material (Figure 1).
We also note that the use of norm improved both running time and memory consumption for models we tested. These benefits can be attributed to the fact that absolute-value operation is computationally more efficient compared to the costly square and sqrt operations. Additionally, the derivative of is the operation . Therefore, in order to compute the gradients, we only need to cache the sign of the values (not the actual values), allowing for substantial memory savings.
Another alternative measure for variability that avoids the discussed limitations of the traditional batch norm is the maximum absolute deviation. For a layer with dimensional input , batch normalization normalize each dimension
where is the expectation over , is batch size and is computed similarly to (derivation appears in appendix).
While normalizing according to the maximum absolute deviation offers a major performance advantage, we found it somewhat less robust to noise compared to and normalization.
3 Batch norm at half precision
Due to numerical issues, prior attempts to train neural networks at low precision had to leave batch norm operations at full precision (float 32) as described by Micikevicius et al. , Das et al. , thus enabling only mixed precision training. This effectively means that low precision hardware still needs to support full precision data types. The sensitivity of BN to low precision operations can be attributed to both the numerical operations of square and square-root used, as well as the possible overflow of the sum of many large positive values. To overcome this overflow, we may further require a wide accumulator with full precision.
We provide evidence that by using arithmetic, batch normalization can also be quantized to half precision with no apparent effect on validation accuracy, as can be seen in figure 3. Using the standard BN in low-precision leads to overflow and significant quantization noise that quickly deteriorate the whole training process, while BN allows training with no visible loss of accuracy.
Improving weight normalization
Trying to address several of the limitations of BN, Salimans & Kingma suggested weight normalization as its replacement. As weight-norm requires an normalization over the output channels of the weight matrix, it alleviates both computational and task-specific shortcomings of BN, ensuring no dependency on the current batch of sample activations within a layer.
While this alternative works well for small-scale problems, as demonstrated in the original work, it was noted by Gitman & Ginsburg to fall short in large-scale usage. For example, in the ImageNet classification task, weight-norm exhibited unstable convergence and significantly lower performance (67% accuracy on ResNet50 vs. 75% for original).
An additional modification of weight-norm called "normalization propagation" adds additional multiplicative and additive corrections to address the change of activation distribution introduced by the ReLU non-linearity used between layers in the network. These modifications are not trivially applied to architectures with complex structure elements such as residual connections .
So far, we’ve demonstrated that the key to the performance of normalization techniques lies in their property to neutralize the effect of weight’s norm. Next, we will use this reasoning to overcome the shortcoming of weight-norm.
2 Norm bounded weight-normalization
We return to the original parametrization suggested for weight norm, for a given initialized weight matrix with output channels:
where is a parameterized weight for the th output channel, composed from an normalized vector and scalar .
Weight-norm successfully normalized each output channel’s weights to reside on the sphere. However, it allowed the weights scale to change freely through the scalar . Following reasoning presented earlier in this work, we wish to make the weight’s norm completely disjoint from its values. We can achieve this by keeping the norm fixed as follows:
where is a fixed scalar for each layer that is determined by its size (number of input and output channels). A simple choice for is by the initial norm of the weights, e.g , thus employing the various successful heuristics used to initialize modern networks . We also note that when using non-linearity with no scale sensitivity (e.g ReLU), these constants can be instead incorporated into only the final classifier’s weights and biases throughout the network.
Previous works demonstrated that weight-normalized networks converge faster when augmented with mean only batch normalization. We follow this regime, although noting that similar final accuracy can be achieved without mean normalization but at the cost of slower convergence, or with the use of zero-mean preserving activation functions .
After this modification, we now find that weight-norm can be improved substantially, solving the stability issues for large-scale task observed by Gitman & Ginsburg and achieving comparable accuracy (although still behind BN). Results on Imagenet using Resnet50 are described in Figure 4, using the original settings and training regime . We believe the still apparent margin between the two methods can be further decreased using hyper-parameter tuning, such as a modified learning rate schedule.
It is also interesting to observe BWN’s effect in recurrent networks, where BN is not easily applicable . We compare weight-norm vs. the common implementation (with layer-norm) of an attention-based LSTM model on the WMT14 en-de translation task . The model consists of 2 LSTM cells for both encoder and decoder, with an attention mechanism. We also compared BWN on the Transformer architecture to replace layer-norm, again achieving comparable final performance (26.5 vs. 27.3 BLEU score on the original base model). Both sequence-to-sequence models were tested using beam-search decoding with a beam size of and length penalty of . Additional results for BWN can be found in the supplementary material (Figure 2 and Table 1).
As we did for BN, we can consider weight-normalization over norms other than such that
where computing the constant over desired (vector) norm will ensure proper scaling that was required in the BN case. We find that similarly to BN, the norm can serve as an alternative to original weight-norm, where using cause a noticeable degradation when using its proper form (top-1 absolute maximum).
Discussion
In this work, we analyzed common normalization techniques used in deep learning models, with BN as their prime representative. We considered a novel perspective on the role of these methods, as tools to decouple the weights’ norm from training objective. This perspective allowed us to re-evaluate the necessity of regularization methods such as weight decay, and to suggest new methods for normalization, targeting the computational, numerical and task-specific deficiencies of current techniques.
Specifically, we showed that the use of and -based normalization schemes could provide similar results to the standard BN while allowing low-precision computation. Such methods can be easily implemented and deployed to serve in current and future network architectures, low-precision devices. A similar normalization scheme to ours was recently introduced by Wu et al. , appearing in parallel to us (within a week). In contrast to Wu et al. , we found that the normalization constant is crucial for achieving the same performance as (see Figure 1 in supplementary). We additionally demonstrated the benefits of normalization: it allowed us to perform BN in half-precision floating-point, which was noted to fail in previous works and required full and mixed precision hardware.
Moreover, we suggested a bounded weight normalization method, which achieves improved results on large-scale tasks (ImageNet) and is nearly comparable with BN. Such a weight normalization scheme improves computational costs and can enable improved learning in tasks that were not suited for previous methods such as reinforcement-learning and temporal modeling.
We further suggest that insights gained from our findings can have an additional impact on the way neural networks are devised and trained. As previous works demonstrated, a strong connection appears between the batch-size used and the optimal learning rate regime and between the weight-decay factor and learning-rate . We deepen this connection and suggest that all of these factors, including the effective norm (or temperature), are mutually affecting one another. It is plausible, given our results, that some (or all) of these hyper-parameters can be fixed given another, which can potentially ease the design and training of modern models.
Acknowledgments
This research was supported by the Israel Science Foundation (grant No. 31/1031), and by the Taub foundation. A Titan Xp used for this research was donated by the NVIDIA Corporation.
References
Appendix
For all experiments, we used weight decay on the last layer with . The network architecture was VGG11 with batch-norm after every convolution layer. Learning rate started from 0.1 and divided by 10 every 20 epochs (except for the norm scheduling experiment). The same random seed was used.
Appendix B Importance of normalization constants
Figure 5 shows the scale adjustment is essential even for relatively "easy" data sets with small images such as CIFAR-10, and the use of smaller/bigger adjustments degrade classification accuracy.
To derive we assume again the input to the normalization layer follows a Gaussian distribution . Then, the maximum absolute deviation is bounded on expectation as follows :
Therefore, by multiplying the three sides of inequality with the normalization term , the batch norm in equation 9 approximates an expectation the original standard deviation measure as follows:
where , and .
Appendix C Bounded-weight-norm experiments
Figure 6 depicts the impact of bounded-weight norm for the training of recurrent network on WMT14 de-en task. Additional results are summarized in Table 1.