Fixup Initialization: Residual Learning Without Normalization
Hongyi Zhang, Yann N. Dauphin, Tengyu Ma
Introduction
Artificial intelligence applications have witnessed major advances in recent years. At the core of this revolution is the development of novel neural network models and their training techniques. For example, since the landmark work of He et al. (2016), most of the state-of-the-art image recognition systems are built upon a deep stack of network blocks consisting of convolutional layers and additive skip connections, with some normalization mechanism (e.g., batch normalization (Ioffe & Szegedy, 2015)) to facilitate training and generalization. Besides image classification, various normalization techniques (Ulyanov et al., 2016; Ba et al., 2016; Salimans & Kingma, 2016; Wu & He, 2018) have been found essential to achieving good performance on other tasks, such as machine translation (Vaswani et al., 2017) and generative modeling (Zhu et al., 2017). They are widely believed to have multiple benefits for training very deep neural networks, including stabilizing learning, enabling higher learning rate, accelerating convergence, and improving generalization.
Despite the enormous empirical success of training deep networks with normalization, and recent progress on understanding the working of batch normalization (Santurkar et al., 2018), there is currently no general consensus on why these normalization techniques help training residual neural networks. Intrigued by this topic, in this work we study
without normalization, can a deep residual network be trained reliably? (And if so,)
without normalization, can a deep residual network be trained with the same learning rate, converge at the same speed, and generalize equally well (or even better)?
Perhaps surprisingly, we find the answers to both questions are Yes. In particular, we show:
Why normalization helps training. We derive a lower bound for the gradient norm of a residual network at initialization, which explains why with standard initializations, normalization techniques are essential for training deep residual networks at maximal learning rate. (Section 2)
Training without normalization. We propose Fixup, a method that rescales the standard initialization of residual branches by adjusting for the network architecture. Fixup enables training very deep residual networks stably at maximal learning rate without normalization. (Section 3)
Image classification. We apply Fixup to replace batch normalization on image classification benchmarks CIFAR-10 (with Wide-ResNet) and ImageNet (with ResNet), and find Fixup with proper regularization matches the well-tuned baseline trained with normalization. (Section 4.2)
Machine translation. We apply Fixup to replace layer normalization on machine translation benchmarks IWSLT and WMT using the Transformer model, and find it outperforms the baseline and achieves new state-of-the-art results on the same architecture. (Section 4.3)
In the remaining of this paper, we first analyze the exploding gradient problem of residual networks at initialization in Section 2. To solve this problem, we develop Fixup in Section 3. In Section 4 we quantify the properties of Fixup and compare it against state-of-the-art normalization methods on real world benchmarks. A comparison with related work is presented in Section 5.
Problem: ResNet with Standard Initializations Lead to Exploding Gradients
Standard initialization methods (Glorot & Bengio, 2010; He et al., 2015; Xiao et al., 2018) attempt to set the initial parameters of the network such that the activations neither vanish nor explode. Unfortunately, it has been observed that without normalization techniques such as BatchNorm they do not account properly for the effect of residual connections and this causes exploding gradients. Balduzzi et al. (2017) characterizes this problem for ReLU networks, and we will generalize this to residual networks with positively homogenous activation functions. A plain (i.e. without normalization layers) ResNet with residual blocks and input computes the activations as
As we will show, at initialization, the gradient norm of certain activations and weight tensors is lower bounded by the cross-entropy loss up to some constant. Intuitively, this implies that blowup in the logits will cause gradient explosion. Our result applies to convolutional and linear weights in a neural network with ReLU nonlinearity (e.g., feed-forward network, CNN), possibly with skip connections (e.g., ResNet, DenseNet), but without any normalization.
Our analysis utilizes properties of positively homogeneous functions, which we now introduce.
Let be the set of parameters of and . We call a positively homogeneous set (of first degree) (p.h. set) if for any , , where denotes .
Intuitively, a p.h. set is a set of parameters in function such that for any fixed input and fixed parameters , is a p.h. function.
Examples of p.h. functions are ubiquitous in neural networks, including various kinds of linear operations without bias (fully-connected (FC) and convolution layers, pooling, addition, concatenation and dropout etc.) as well as ReLU nonlinearity. Moreover, we have the following claim:
A function that is the composition of p.h. functions is itself p.h.
is a sequential composition of network blocks , i.e. , each of which is composed of p.h. functions.
Weight elements in the FC layer are i.i.d. sampled from a zero-mean symmetric distribution.
These assumptions hold at initialization if we remove all the normalization layers in a residual network with ReLU nonlinearity, assuming all the biases are initialized at .
Our results are summarized in the following two theorems, whose proofs are listed in the appendix:
Denote the input to the -th block by . With Assumption 1, we have
where is the softmax probabilities and denotes the Shannon entropy.
Since is upper bounded by and is small in the lower blocks, blowup in the loss will cause large gradient norm with respect to the lower block input. Our second theorem proves a lower bound on the gradient norm of a p.h. set in a network.
Furthermore, with Assumptions 1 and 2, we have
It remains to identify such p.h. sets in a neural network. In Figure 2 we provide three examples of p.h. sets in a ResNet without normalization. Theorem 2 suggests that these layers would suffer from the exploding gradient problem, if the logits blow up at initialization, which unfortunately would occur in a ResNet without normalization if initialized in a traditional way. This motivates us to introduce a new initialization in the next section.
Fixup: Update a Residual Network Θ(η)Θ𝜂\Theta(\eta) per SGD Step
Our analysis in the previous section points out the failure mode of standard initializations for training deep residual network: the gradient norm of certain layers is in expectation lower bounded by a quantity that increases indefinitely with the network depth. However, escaping this failure mode does not necessarily lead us to successful training — after all, it is the whole network as a function that we care about, rather than a layer or a network block. In this section, we propose a top-down design of a new initialization that ensures proper update scale to the network function, by simply rescaling a standard initialization. To start, we denote the learning rate by and set our goal:
We define the Shortcut as the shortest path from input to output in a residual network. The Shortcut is typically a shallow network with a few trainable layers.For example, in the ResNet architecture (e.g., ResNet-50, ResNet-101 or ResNet-152) for ImageNet classification, the Shortcut is always a 6-layer network with five convolution layers and one fully-connected layer, irrespective of the total depth of the whole network. We assume the Shortcut is initialized using a standard method, and focus on the initialization of the residual branches.
To start, we first make an important observation that the SGD update to each residual branch changes the network output in highly correlated directions. This implies that if a residual network has residual branches, then an SGD step to each residual branch should change the network output by on average to achieve an overall update. We defer the formal statement and its proof until Section B.1.
Through deriving the constraints for to make updates, we will also discover how to rescale the weight layers of a standard initialization as desired. In particular, we show the SGD update to is if and only if the initialization satisfies the following constraint:
We defer the derivation until Section B.2.
Equation 5 suggests new methods to initialize a residual branch through rescaling the standard initialization of i-th layer in a residual branch by its corresponding scalar . For example, we could set . Alternatively, we could start the residual branch as a zero function by setting and . In the second option, the residual branch does not need to “unlearn” its potentially bad random initial state, which can be beneficial for learning. Therefore, we use the latter option in our experiments, unless otherwise specified.
With proper rescaling of the weights in all the residual branches, a residual network is supposed to be updated by per SGD step — our goal is achieved. However, in order to match the training performance of a corresponding network with normalization, there are two more things to consider: biases and multipliers.
Using biases in the linear and convolution layers is a common practice. In normalization methods, bias and scale parameters are typically used to restore the representation power after normalization.For example, in batch normalization gamma and beta parameters are used to affine-transform the normalized activations per each channel. Intuitively, because the preferred input/output mean of a weight layer may be different from the preferred output/input mean of an activation layer, it also helps to insert bias terms in a residual network without normalization. Empirically, we find that inserting just one scalar bias before each weight layer and nonlinear activation layer significantly improves the training performance.
Multipliers scale the output of a residual branch, similar to the scale parameters in batch normalization. They have an interesting effect on the learning dynamics of weight layers in the same branch. Specifically, as the stochastic gradient of a layer is typically almost orthogonal to its weight, learning rate decay tends to cause the weight norm equilibrium to shrink when combined with L2 weight decay (van Laarhoven, 2017). In a branch with multipliers, this in turn causes the growth of the multipliers, increasing the effective learning rate of other layers. In particular, we observe that inserting just one scalar multiplier per residual branch mimics the weight norm dynamics of a network with normalization, and spares us the search of a new learning rate schedule.
Put together, we propose the following method to train residual networks without normalization:
Fixup initialization (or: How to train a deep residual network without normalization) 1. Initialize the classification layer and the last layer of each residual branch to . 2. Initialize every other layer using a standard method (e.g., He et al. (2015)), and scale only the weight layers inside residual branches by . 3. Add a scalar multiplier (initialized at 1) in every branch and a scalar bias (initialized at 0) before each convolution, linear, and element-wise activation layer. It is important to note that Rule 2 of Fixup is the essential part as predicted by Equation 5. Indeed, we observe that using Rule 2 alone is sufficient and necessary for training extremely deep residual networks. On the other hand, Rule 1 and Rule 3 make further improvements for training so as to match the performance of a residual network with normalization layers, as we explain in the above text.It is worth noting that the design of Fixup is a simplification of the common practice, in that we only introduce parameters beyond convolution and linear weights (since we remove bias terms from convolution and linear layers), whereas the common practice includes (Ioffe & Szegedy, 2015; Salimans & Kingma, 2016) or (Ba et al., 2016) additional parameters, where is the number of layers, is the max number of channels per layer and are the spatial dimension of the largest feature maps. We find ablation experiments confirm our claims (see Section C.1).
Our initialization and network design is consistent with recent theoretical work Hardt & Ma (2016); Li et al. (2018), which, in much more simplified settings such as linearized residual nets and quadratic neural nets, propose that small initialization tend to stabilize optimization and help generalizaiton. However, our approach suggests that more delicate control of the scale of the initialization is beneficial.For example, learning rate smaller than our choice would also stabilize the training, but lead to lower convergence rate.
Experiments
Figure 3 shows the test accuracy at the first epoch as depth increases. Observe that Fixup matches the performance of BatchNorm at the first epoch, even with 10,000 layers. LSUV and -scaling are not able to train with the same learning rate as BatchNorm past 100 layers.
2 Image classification
In this section, we evaluate the ability of Fixup to replace batch normalization in image classification applications. On the CIFAR-10 dataset, we first test on ResNet-110 (He et al., 2016) with default hyper-parameters; results are shown in Table 1. Fixup obtains 7% relative improvement in test error compared with standard initialization; however, we note a substantial difference in the difficulty of training. While network with Fixup is trained with the same learning rate and converge as fast as network with batch normalization, we fail to train a Xavier initialized ResNet-110 with 0.1x maximal learning rate.Personal communication with the authors of (Shang et al., 2017) confirms our observation, and reveals that the Xavier initialized network need more epochs to converge. The test error gap in Table 1 is likely due to the regularization effect of BatchNorm rather than difficulty in optimization; when we train Fixup networks with better regularization, the test error gap disappears and we obtain state-of-the-art results on CIFAR-10 and SVHN without normalization layers (see Section C.2).
On the ImageNet dataset, we benchmark Fixup with the ResNet-50 and ResNet-101 architectures (He et al., 2016), trained for 100 epochs and 200 epochs respectively. Similar to our finding on the CIFAR-10 dataset, we observe that (1) training with Fixup is fast and stable with the default hyperparameters, (2) Fixup alone significantly improves the test error of standard initialization, and (3) there is a large test error gap between Fixup and BatchNorm. Further inspection reveals that Fixup initialized models obtain significantly lower training error compared with BatchNorm models (see Section C.3), i.e., Fixup suffers from overfitting. We therefore apply stronger regularization to the Fixup models using Mixup (Zhang et al., 2017). We find it is beneficial to reduce the learning rate of the scalar multiplier and bias by x when additional large regularization is used. Best Mixup coefficients are found through cross-validation: they are , and for BatchNorm, GroupNorm (Wu & He, 2018) and Fixup respectively. We present the results in Table 2, noting that with better regularization, the performance of Fixup is on par with GroupNorm.
3 Machine translation
To demonstrate the generality of Fixup, we also apply it to replace layer normalization (Ba et al., 2016) in Transformer (Vaswani et al., 2017), a state-of-the-art neural network for machine translation. Specifically, we use the fairseq library (Gehring et al., 2017) and follow the Fixup template in Section 3 to modify the baseline model. We evaluate on two standard machine translation datasets, IWSLT German-English (de-en) and WMT English-German (en-de) following the setup of Ott et al. (2018). For the IWSLT de-en dataset, we cross-validate the dropout probability from and find to be optimal for both Fixup and the LayerNorm baseline. For the WMT’16 en-de dataset, we use dropout probability . All models are trained for k updates.
It was reported (Chen et al., 2018) that “Layer normalization is most critical to stabilize the training process… removing layer normalization results in unstable training runs”. However we find training with Fixup to be very stable and as fast as the baseline model. Results are shown in Table 3. Surprisingly, we find the models do not suffer from overfitting when LayerNorm is replaced by Fixup, thanks to the strong regularization effect of dropout. Instead, Fixup matches or supersedes the state-of-the-art results using Transformer model on both datasets.
Related Work
Normalization methods have enabled training very deep residual networks, and are currently an essential building block of the most successful deep learning architectures. All normalization methods for training neural networks explicitly normalize (i.e. standardize) some component (activations or weights) through dividing activations or weights by some real number computed from its statistics and/or subtracting some real number activation statistics (typically the mean) from the activations.For reference, we include a brief history of normalization methods in Appendix D. In contrast, Fixup does not compute statistics (mean, variance or norm) at initialization or during any phase of training, hence is not a normalization method.
Training very deep neural networks is an important theoretical problem. Early works study the propagation of variance in the forward and backward pass for different activation functions (Glorot & Bengio, 2010; He et al., 2015).
Recently, the study of dynamical isometry (Saxe et al., 2013) provides a more detailed characterization of the forward and backward signal propogation at initialization (Pennington et al., 2017; Hanin, 2018), enabling training 10,000-layer CNNs from scratch (Xiao et al., 2018). For residual networks, activation scale (Hanin & Rolnick, 2018), gradient variance (Balduzzi et al., 2017) and dynamical isometry property (Yang & Schoenholz, 2017) have been studied. Our analysis in Section 2 leads to the similar conclusion as previous work that the standard initialization for residual networks is problematic. However, our use of positive homogeneity for lower bounding the gradient norm of a neural network is novel, and applies to a broad class of neural network architectures (e.g., ResNet, DenseNet) and initialization methods (e.g., Xavier, LSUV) with simple assumptions and proof.
Hardt & Ma (2016) analyze the optimization landscape (loss surface) of linearized residual nets in the neighborhood around the zero initialization where all the critical points are proved to be global minima. Yang & Schoenholz (2017) study the effect of the initialization of residual nets to the test performance and pointed out Xavier or He initialization scheme is not optimal. In this paper, we give a concrete recipe for the initialization scheme with which we can train deep residual networks without batch normalization successfully.
Despite its popularity in practice, batch normalization has not been well understood. Ioffe & Szegedy (2015) attributed its success to “reducing internal covariate shift”, whereas Santurkar et al. (2018) argued that its effect may be “smoothing loss surface”. Our analysis in Section 2 corroborates the latter idea of Santurkar et al. (2018) by showing that standard initialization leads to very steep loss surface at initialization. Moreover, we empirically showed in Section 3 that steep loss surface may be alleviated for residual networks by using smaller initialization than the standard ones such as Xavier or He’s initialization in residual branches. van Laarhoven (2017); Hoffer et al. (2018) studied the effect of (batch) normalization and weight decay on the effective learning rate. Their results inspire us to include a multiplier in each residual branch.
Gehring et al. (2017); Balduzzi et al. (2017) proposed to address the initialization problem of residual nets by using the recurrence . Mishkin & Matas (2015) proposed a data-dependent initialization to mimic the effect of batch normalization in the first forward pass. While both methods limit the scale of activation and gradient, they would fail to train stably at the maximal learning rate for very deep residual networks, since they fail to consider the accumulation of highly correlated updates contributed by different residual branches to the network function (Section B.1). Srivastava et al. (2015); Hardt & Ma (2016); Goyal et al. (2017); Kingma & Dhariwal (2018) found that initializing the residual branches at (or close to) zero helped optimization. Our results support their observation in general, but Equation 5 suggests additional subtleties when choosing a good initialization scheme.
Conclusion
In this work, we study how to train a deep residual network reliably without normalization. Our theory in Section 2 suggests that the exploding gradient problem at initialization in a positively homogeneous network such as ResNet is directly linked to the blowup of logits. In Section 3 we develop Fixup initialization to ensure the whole network as well as each residual branch gets updates of proper scale, based on a top-down analysis. Extensive experiments on real world datasets demonstrate that Fixup matches normalization techniques in training deep residual networks, and achieves state-of-the-art test performance with proper regularization.
Our work opens up new possibilities for both theory and applications. Can we analyze the training dynamics of Fixup, which may potentially be simpler than analyzing models with batch normalization is? Could we apply or extend the initialization scheme to other applications of deep learning? It would also be very interesting to understand the regularization benefits of various normalization methods, and to develop better regularizers to further improve the test performance of Fixup.
The authors would like to thank Yuxin Wu, Kaiming He, Aleksander Madry and the anonymous reviewers for their helpful feedback.
References
Appendix A Proofs for Section 2
We use to denote the composition , so that for all . Note that is p.h. with respect to the input of each network block, i.e. for . This allows us to compute the gradient of the cross-entropy loss with respect to the scaling factor at as
A.2 Gradient norm lower bound for positively homogeneous sets
The proof idea is similar. Recall that if is a p.h. set, then is a p.h. function. We therefore have
hence we again invoke the directional derivative argument to show
In order to estimate the scale of this lower bound, recall the FC layer weights are i.i.d. sampled from a symmetric, mean-zero distribution, therefore has a symmetric probability density function with mean . We hence have
Appendix B Proofs for Section 3
A common theme in previous analysis of residual networks is the scale of activation and gradient (Balduzzi et al., 2017; Yang & Schoenholz, 2017; Hanin & Rolnick, 2018). However, it is more important to consider the scale of actual change to the network function made by a (stochastic) gradient descent step. If the updates to different layers cancel out each other, the network would be stable as a whole despite drastic changes in different layers; if, on the other hand, the updates to different layers align with each other, the whole network may incur a drastic change in one step, even if each layer only changes a tiny amount. We now provide analysis showing that the latter scenario more accurately describes what happens in reality at initialization.
For our result in this section, we make the following assumptions:
is a sequential composition of network blocks , i.e. , consisting of fully-connected weight layers, ReLU activation functions and residual branches.
is a fully-connected layer with weights i.i.d. sampled from a zero-mean distribution.
For , let be the input to and be a branch in with layers. Without loss of generality, we study the following specific form of network architecture:
For the last block we denote and .
We have the following result on the gradient update to :
The first insight to prove our result is to note that conditioning on a specific input , we can replace each ReLU activation layer by a diagonal matrix and does not change the forward and backward pass. (In fact, this is valid even after we apply a gradient descent update, as long as the learning rate is sufficiently small so that all positive preactivation remains positive. This observation will be essential for our later analysis.) We thus have the gradient w.r.t. the -th weight layer in the -th block is
where denotes the Kronecker product. The second insight is to note that with our assumptions, a network block and its gradient w.r.t. its input have the following relation:
B.2 What scalar branch has Θ(η/L)Θ𝜂𝐿\Theta(\eta/L) updates?
For this section, we focus on the proper initialization of a scalar branch . We have the following result:
We start by calculating the gradient of each parameter:
and a first-order approximation of :
where we conveniently abuse some notations by defining
Denote as and as , we have
and therefore by rearranging Equation 17 and letting we get
i.e. . Hence the “only if” part is proved. For the “if” part, we apply Equation 19 to Equation 17 and observe that by Equation 15
The result of this theorem provides useful guidance on how to rescale the standard initialization to achieve the desired update scale for the network function.
Appendix C Additional experiments
In this section we present the training curves of different architecture designs and initialization schemes. Specifically, we compare the training accuracy of batch normalization, Fixup, as well as a few ablated options: (1) removing the bias parameters in the network; (2) use x the suggested initialization scale and no bias parameters; (3) use x the suggested initialization scale and no bias parameters; and (4) remove all the residual branches. The results are shown in Figure 4. We see that initializing the residual branch layers at a smaller scale (or all zero) slows down learning, whereas training fails when initializing them at a larger scale; we also see the clear benefit of adding bias parameters in the network.
C.2 CIFAR and SVHN with better regularization
We perform additional experiments to validate our hypothesis that the gap in test error between Fixup and batch normalization is primarily due to overfitting. To combat overfitting, we use Mixup (Zhang et al., 2017) and Cutout (DeVries & Taylor, 2017) with default hyperparameters as additional regularization. On the CIFAR-10 dataset, we perform experiments with WideResNet-40-10 and on SVHN we use WideResNet-16-12 (Zagoruyko & Komodakis, 2016), all with the default hyperparameters. We observe in Table 4 that models trained with Fixup and strong regularization are competitive with state-of-the-art methods on CIFAR-10 and SVHN, as well as our baseline with batch normalization.
C.3 Training and test curves on ImageNet
Figure 5 shows that without additional regularization Fixup fits the training set very well, but overfits significantly. We see in Figure 6 that Fixup is competitive with networks trained with normalization when the Mixup regularizer is used.
Appendix D Additional references: A brief history of normalization methods
The first use of normalization in neural networks appears in the modeling of biological visual system and dates back at least to Heeger (1992) in neuroscience and to Pinto et al. (2008); Lyu & Simoncelli (2008) in computer vision, where each neuron output is divided by the sum (or norm) of all of the outputs, a module called divisive normalization. Recent popular normalization methods, such as local response normalization (Krizhevsky et al., 2012), batch normalization (Ioffe & Szegedy, 2015) and layer normalization (Ba et al., 2016) mostly follow this tradition of dividing the neuron activations by their certain summary statistics, often also with the activation mean subtracted. An exception is weight normalization (Salimans & Kingma, 2016), which instead divides the weight parameters by their statistics, specifically the weight norm; weight normalization also adopts the idea of activation normalization for weight initialization. The recently proposed actnorm (Kingma & Dhariwal, 2018) removes the normalization of weight parameters, but still use activation normalization to initialize the affine transformation layers.