The Normalization Method for Alleviating Pathological Sharpness in Wide Neural Networks

Ryo Karakida, Shotaro Akaho, Shun-ichi Amari

Introduction

Deep neural networks (DNNs) have performed excellently in various practical applications , but there are still many heuristics and an arbitrariness in their settings and learning algorithms. To proceed further, it would be beneficial to give theoretical elucidation of how and under what conditions deep learning works well in practice.

Normalization methods are widely used to enhance the trainability and generalization ability of DNNs. In particular, batch normalization makes optimization faster with a large learning rate and achieves better generalization in experiments . Recently, some studies have reported that batch normalization changes the shape of the loss function, which leads to better performance . Batch normalization alleviates a sharp change of the loss function and makes the loss landscape smoother , and prevents an explosion of the loss function and its gradient . The flatness of the loss landscape and its geometric characterization have been explored in various topics such as improvement of generalization ability , advantage of skip connections , and robustness against adversarial attacks . Thus, it seems to be an important direction of research to investigate normalization methods from the viewpoint of the geometric characterization. Nevertheless, its theoretical elucidation has been limited to only linear networks and simplified models neglecting the hierarchical structure of DNNs .

One promising approach of analyzing normalization methods is to consider DNNs with random weights and sufficiently wide hidden layers. While theoretical analysis of DNNs often becomes intractable because of hierarchical nonlinear transformations, wide DNNs with random weights can overcome such difficulties and are attracting much attention, especially within the last few years; mean field theory of DNNs , random matrix theory and kernel methods . They have succeeded in predicting hyperparameters with which learning algorithms work well and even used as a kernel function for the Gaussian process. In addition, recent studies on the neural tangent kernel (NTK) have revealed that the Gaussian process with the NTK of random initialization determines even the performance of trained neural networks . Thus, the theory of wide DNNs is becoming a foundation for comprehensive understanding of DNNs. Regarding the geometric characterization, there have been studies on the Fisher information matrix (FIM) of wide DNNs . The FIM widely appears in the context of deep learning because it determines the Riemannian geometry of the parameter space and a local shape of the loss landscape around a certain global minimum. In particular, Karakida et al. have reported that the eigenvalue spectrum of the FIM is strongly distorted in wide DNNs, that is, the largest eigenvalue takes a pathologically large value (Theorem 2.2). This causes pathological sharpness of the landscape and such sharpness seems to be harmful from the perspective of optimization and generalization .

In this study, we focus on the FIM of DNNs and uncover how normalization methods affect it. First, to clarify a condition to alleviate the pathologically large eigenvalues, we identify the eigenspace of the largest eigenvalues (Theorem 3.1). Then, we reveal that batch normalization in the last layer drastically decreases the size of the largest eigenvalues and successfully alleviates the pathological sharpness. This alleviation requires a certain condition on the width and sample size (Theorem 3.3), which is determined by a convergence rate of order parameters. In contrast, we find that batch normalization in the middle layers cannot alleviate pathological sharpness in many settings (Theorem 3.4) and layer normalization cannot either (Theorem 4.1). Thus, we can conclude that batch normalization in the last layer has a vital role in decreasing pathological sharpness. Our experiments suggest that such alleviation of the sharpness is helpful in making gradient descent converge even with a larger learning rate. These results give novel quantitative insight into normalization methods, wide DNNs, and geometric characterization of DNNs and is expected to be helpful in developing a further theory of deep learning.

Preliminaries

We investigate a fully-connected feedforward neural network with random weights and bias parameters. The network consists of one input layer with M0M_{0} units, L−1L-1 hidden layers with MlM_{l} units per layer (l=1,2,...,L−1l=1,2,...,L-1), and one output layer:

The FIM of a DNN is computed by the chain rule in a manner similar to that of the backpropagation algorithm:

where we denote fk=ukLf_{k}=u_{k}^{L} and δk,il:=∂fk/∂uil\delta_{k,i}^{l}:=\partial f_{k}/\partial u_{i}^{l} for k=1,...,Ck=1,...,C. To avoid complicated notation, we omit index kk of the output unit, i.e., δil=δk,il\delta_{i}^{l}=\delta_{k,i}^{l}.

2 Understanding DNNs through order parameters

for l=0,...,L−1l=0,...,L-1. Because input samples generated by Eq. (3) yield q^t0=1\hat{q}^{0}_{t}=1 and q^st0=0\hat{q}_{st}^{0}=0 for all ss and tt, q^stl\hat{q}_{st}^{l} in each layer takes the same value for all s≠ts\neq t, and so does q^tl\hat{q}_{t}^{l} for all tt. The notation Du=duexp⁡(−u2/2)/2πDu=du\exp(-u^{2}/2)/\sqrt{2\pi} means integration over the standard Gaussian density. We use a two-dimensional Gaussian integral given by Iϕ[a,b]:=∫DyDxϕ(ax)ϕ(a(cx+1−c2y))I_{\phi}[a,b]:=\int DyDx\phi(\sqrt{a}x)\phi(\sqrt{a}(cx+\sqrt{1-c^{2}}y)) with c=b/ac=b/a.

The order parameters depend only on σw2\sigma_{w}^{2} and σb2\sigma_{b}^{2}, the types of activation functions, and depth. The recurrence relations require LL iterations of one- and two-dimensional numerical integrals. They are analytically tractable in certain activation functions including the ReLUs .

3 Pathological sharpness of local landscapes

The FIM plays an essential role in the geometry of the parameter space and is a fundamental quantity in both statistics and machine learning. It defines a Riemannian metric of the parameter space, where the infinitesimal difference of statistical models is measured by Kullback-Leibler divergence, as in information geometry . We analyze the eigenvalue statistics of the following FIM of DNNs ,

The FIM is known to determine not only the local distortion of the parameter space but also the loss landscape around a certain global minimum. Suppose the squared loss function E(θ)=(1/2T)∑k=1C∑t=1T(yk(t)−fk(t))2E(\theta)=(1/2T)\sum_{k=1}^{C}\sum_{t=1}^{T}(y_{k}(t)-f_{k}(t))^{2}, where yk(t)y_{k}(t) represents a training label corresponding to the input sample x(t)x(t). The FIM is related to the Hessian of the loss function, H:=∇θ∇θE(θ)H:=\nabla_{\theta}\nabla_{\theta}E(\theta), in the following manner :

The Hessian coincides with the empirical FIM when the parameter converges to the global minimum with zero training error. In that sense, the FIM determines the local shape of the loss landscape around the minimum. This FIM is also known as the Gauss-Newton approximation of the Hessian.

Karakida et al. elucidated hidden relations between the order parameters and basic statistics of the FIM’s eigenvalues. We investigate DNNs satisfying the following condition.

Suppose a DNN with bias terms (σb≠0\sigma_{b}\neq 0) or activation functions satisfying the non-zero Gaussian mean. We refer to this as a non-centered network.

The definition of the non-zero Gaussian mean is ∫Dzϕ(z)≠0\int Dz\phi(z)\neq 0. Non-centered networks include various realistic settings because usual networks include bias terms, and widely used activation functions, such as the sigmoid function and (leaky-) ReLUs, have the non-zero Gaussian mean. Denote the FIM’s eigenvalues as λi\lambda_{i} (i=1,...,Pi=1,...,P) where PP is the number of all parameters. The eigenvalues are non-negative by definition. Their mean is mλ:=∑iPλi/Pm_{\lambda}:=\sum_{i}^{P}\lambda_{i}/P and the maximum is λmax:=max⁡iλi\lambda_{max}:=\max_{i}\lambda_{i}. The following theorem holds:

Suppose a non-centered network and i.i.d. input samples generated by Eq. (3). When MM is sufficiently large, the eigenvalue statistics of FF are asymptotically evaluated as

where α:=∑l=1L−1αlαl−1\alpha:=\sum_{l=1}^{L-1}\alpha_{l}\alpha_{l-1}, and positive constants κ1\kappa_{1} and κ2\kappa_{2} are obtained using order parameters,

The mean is asymptotically close to zero and it implies that most of the eigenvalues are very small. In contrast, λmax\lambda_{max} becomes pathologically large in proportion to the width. We refer to this λmax\lambda_{max} as pathological sharpness since FIM’s eigenvalues determine the local shape of the parameter space and loss landscape. Empirical experiments reported that both of close-to-zero eigenvalues and pathologically large ones appear in trained networks as well .

Pathological sharpness universally appears in various DNNs. Technically speaking, if the network is not non-centered (i.e., a network with no bias terms and zero-Gaussian mean; we call it a centered network), κ2=0\kappa_{2}=0 holds and lower order terms of the eigenvalue statistics become non-negligible , and the pathological sharpness may disappear. For instance, λmax\lambda_{max} is of O(1)O(1) when TT is properly scaled with MM in a centered shallow network . Except for such special centered networks, we cannot avoid pathological sharpness. In practice, it would be better to alleviate the pathologically large λmax\lambda_{max} because it causes the sharp loss landscape. It requires very small learning rates (see Section 3.4) and will lead to worse generalization . In the following section, we reveal that a specific normalization method plays an important role in alleviating pathological sharpness.

Alleviation of pathological sharpness in batch normalization

Before analyzing the effects of normalization methods on the FIM, it will be helpful to characterize the cause of pathological sharpness. We find the following eigenspace of λmax\lambda_{max}’s:

Suppose a non-centered network and i.i.d. input samples generated by Eq. (3). When MM is sufficiently large, the eigenvectors corresponding to λmax\lambda_{max}’s are asymptotically equivalent to

2 Batch normalization in last layer

In this section, we analyze batch normalization in the last layer (LL-th layer):

In the following analysis, we use a widely used assumption for DNNs with random weights:

Supposing this assumption has been a central technique of the mean field theory of DNNs to make the derivation of backward order parameters relatively easy. These studies confirmed that this assumption leads to excellent agreements with experimental results. Moreover, recent studies have succeeded in theoretically justifying that various statistical quantities obtained under this assumption coincide with exact solutions. Thus, Assumption 3.2 is considered to be effective as the first step of the analysis.

First, let us set σk(θ)\sigma_{k}(\theta) as a constant and only consider mean subtraction in the last layer:

Since the σk(θ)\sigma_{k}(\theta) controls the scale of the network output, one may suspect that the contribution of the mean subtraction would only be restrictive for alleviating sharpness. Contrary to this expectation, we find an interesting fact that the mean subtraction is essential to alleviate pathological sharpness:

Suppose a non-centered network with the mean subtraction in the last layer (Eq. (15)) and i.i.d. input samples generated by Eq. (3). In the large MM limit, the mean of the FIM’s eigenvalues is asymptotically evaluated by

The largest eigenvalue is asymptotically evaluated as follows: (i) when T≥2T\geq 2 and T=O(1)T=O(1),

and (ii) when T=O(M)T=O(M) with a constant ρ:=M/T\rho:=M/T, under the gradient independence assumption, we have

for non-negative constants c1c_{1} and c2c_{2}.

The derivation is shown in Supplementary Material B.1. The mean subtraction does not change the order of λmax\lambda_{max} when T=O(1)T=O(1). In contrast, it is interesting that it decreases the order when T=O(M)T=O(M). The decrease in mλm_{\lambda} only appears in the coefficient because κ1>κ1−κ2>0\kappa_{1}>\kappa_{1}-\kappa_{2}>0 hold in the non-centered networks. Thus, we can conclude that the mean subtraction in the last layer plays an essential role in decreasing λmax\lambda_{max} when TT is appropriately scaled to MM.

As shown in Fig.1, we empirically confirmed that λmax\lambda_{max} became of O(1)O(1) in numerical experiments and pathological sharpness disappeared when T=MT=M. Theorem 3.3 is consistent with the numerical experimental results. We numerically computed λmax\lambda_{max} in DNNs with random Gaussian weights, biases, and input samples generated by Eq. (3). We set αl=C=1\alpha_{l}=C=1 and L=3L=3. Variances of parameters were given by (σw2,σb2\sigma_{w}^{2},\sigma_{b}^{2}) = (2,02,0) in the ReLU case, and (3,0.643,0.64) in the tanh case. Each points and error bars show the experimental results over 100100 different ensembles. We show the value of ρα(κ1−κ2)\rho\alpha(\kappa_{1}-\kappa_{2}) as the lower bound of λmax\lambda_{max} (red line). Although this lower bound and the theoretical upper bound of order M\sqrt{M} are relatively loose compared to the experimental results, recall that our purpose is not to obtain the tight bounds but to show the alleviation of λmax\lambda_{max}. The experimental results with the mean subtraction were much lower than those without it as our theory predicts.

We can also add σkL(θ)\sigma_{k}^{L}(\theta) to Theorem 3.3 and obtain the eigenvalue statistics under the normalization (Eq. (13)). When T=O(M)T=O(M), the eigenvalue statistics slightly change to

where Q1:=∑kC1/σk(θ)2Q_{1}:=\sum_{k}^{C}1/\sigma_{k}(\theta)^{2}, Q2:=∑kC1/σk(θ)4Q_{2}:=\sum_{k}^{C}1/\sigma_{k}(\theta)^{4}, c1′c_{1}^{\prime} and c2′c_{2}^{\prime} are non-negative constants. The derivation is shown in Supplementary Material B.2. This clarifies that the variance normalization works only as a constant factor and the mean subtraction is essential to reduce pathological sharpness.

3 Batch normalization in middle layers

To distinguish the effectiveness of normalization in the last layer from those in other layers, we apply batch normalization in all layers except for the last layer :

for all middle layers (l=1,...,L−1)(l=1,...,L-1) while the last layer is kept in an un-normalized manner, i.e., fk(t)=∑jWkjLhjL−1(t)+bkLf_{k}(t)=\sum_{j}W_{kj}^{L}h_{j}^{L-1}(t)+b_{k}^{L}. The variables μil\mu_{i}^{l} and σil\sigma_{i}^{l} depend on weight and bias parameters. For simplicity, we set γil=1\gamma_{i}^{l}=1 and βil=0\beta_{i}^{l}=0. We find a lower bound of λmax\lambda_{max} with order of MM:

Suppose non-negative activation functions and i.i.d. input samples generated by Eq. (3). The largest eigenvalue of the FIM under the normalization (Eq. (20)) is asymptotically lower bounded by

where q^t,BNL−1\hat{q}^{L-1}_{t,BN} and q^st,BNL−1\hat{q}^{L-1}_{st,BN} are positive constants independent of MM.

Because the last layer is unnormalized, we can construct a lower bound composed of the activations in the (L−1)(L-1)-th layer. Note that the set of non-negative activation functions (i.e., ϕ(x)≥0\phi(x)\geq 0) is a subclass of the non-centered networks. It includes sigmoid and ReLU functions which are widely used. The bias term, i.e., σb2\sigma_{b}^{2}, does not affect the theorem because they are canceled out in the mean subtraction of each middle layer. After this batch normalization, λmax\lambda_{max} is still of O(M)O(M) at lowest and the pathological sharpness is unavoidable in that sense. Thus, one can conclude that the normalization in the middle layers cannot alleviate pathological sharpness in many settings.

The constants q^t,BNL−1\hat{q}^{L-1}_{t,BN} and q^st,BNL−1\hat{q}^{L-1}_{st,BN} correspond to feedforward order parameters in batch normalization. The details are shown in Supplementary Material C.1. Although the purpose of our study was to evaluate the order of the eigenvalues, some approaches analytically compute the specific values of the order parameters under certain conditions (see Supplementary Material C.2 for more details). In particular, they are analytically tractable in ReLU networks as follows; q^t,BNL−1=1/2\hat{q}^{L-1}_{t,BN}=1/2 and q^st,BNL−1=12J(−1/(T−1))\hat{q}^{L-1}_{st,BN}=\frac{1}{2}J(-1/(T-1)) where J(x)J(x) is the arccosine kernel .

4 Effect on the gradient descent method

Consider the gradient descent method in a batch regime. Its update rule is given by θt+1←θt−η∇θE(θt)\theta_{t+1}\leftarrow\theta_{t}-\eta\nabla_{\theta}E(\theta_{t}) where η\eta is a constant learning rate. Under some natural assumptions, there exists a necessary condition of the learning rate for the gradient dynamics to converge to a global minimum ;

Because our theory shows that batch normalization in the last layer decreased λmax\lambda_{max}, the appropriate learning rate for convergence becomes larger. To confirm this effect on the learning rate, we did experiments on training with the gradient descent as shown in Fig. 2. we trained DNNs with various widths by using various fixed learning rates, providing i.i.d. Gaussian input samples and labels generated by corresponding teacher networks. It was the same setting as the experiment shown in . Fig. 2 (left) shows the color map of training losses without any normalization method and is just a reproduction of . Losses exploded in the gray area (i.e., were larger than 10310^{3}) and the red line shows the theoretical value of 2/λmax2/\lambda_{max}, which was calculated with the FIM at random initialization. Training above the red line exploded in sufficiently widen DNNs, just as the necessary condition (23) predicts. In contrast, Fig. 2 (right) shows the result of the batch normalization (mean subtraction) in the last layer. We confirmed that it allows larger learning rates for convergence and they are independent of width. We calculated the theoretical line by using the lower bound of λmax\lambda_{max}, i.e., η=2/(ρα(κ1−κ2))\eta=2/(\rho\alpha(\kappa_{1}-\kappa_{2})). Note that Fig. 2 shows the results on the single trial of training with fixed initialization. It caused the stripe pattern of color map depending on the random seed of each width, especially in the case of normalized networks. As shown in Fig. S.2 of Supplementary Material D, accumulation of multiple trials achieves lower losses regardless of the width. Thus, the batch normalization is helpful to set larger learning rates, which could be expected to speed-up the training of neural networks .

Pathological sharpness in layer normalization

It is an interesting question to investigate the effect of other normalization methods on pathological sharpness. Let us consider layer normalization :

for all layers (l=1,...,Ll=1,...,L). The network output is normalized as fk(t)=uˉkL(t)f_{k}(t)=\bar{u}^{L}_{k}(t). While batch normalization (20) normalizes the pre-activation of each unit across batch samples, layer normalization (24) normalizes that of each sample across the units in the same layer. Although layer normalization is the method typically used in recurrent neural networks, we show its effectiveness in feedforward networks to contrast the effect of batch normalization on the FIM. For simplicity, we set γil=1\gamma_{i}^{l}=1 and βil=0\beta_{i}^{l}=0. Then, we find

Suppose a non-centered network, i.i.d. input samples generated by Eq. (3), and the gradient independence assumption. When MM is sufficiently large and C>2C>2, the eigenvalue statistics of the FIM under the normalization (Eq. (24)) are asymptotically evaluated as

where κ1′\kappa_{1}^{\prime}, κ2′\kappa_{2}^{\prime} and ηi\eta_{i} (i=1,2,3)(i=1,2,3) are constants independent of MM.

Related work

Normalization and geometric characterization. Batch normalization is believed to perform well because it suppresses the internal covariate shift . Recent extensive studies, however, have reported alternative explanations on how batch normalization works . Santurkar et al. empirically found that batch normalization decreases a sharp change of the loss function and makes the loss landscape smoother. Bjorck et al. reported that batch normalization works to prevent an explosion of the loss and gradients. While some theoretical studies analyzed FIMs in un-normalized DNNs , analysis in normalized DNNs has been limited. Santurkar et al. analyzed gradients and Hessian under batch normalization in a single layer and theoretically evaluated their worst case bounds, but its inequality was too general to quantify the decrease of sharpness. In particular, it misses the special effect of the last layer, as we found in this study. The original paper of layer normalization analyzed the FIM in generalized linear models (GLMs) and argued that the normalization could decrease curvature of the parameter space. While a GLM corresponds to the single layer model, shallow and deep networks have hidden layers. As the hidden layers become wide, pathological sharpness appears and layer normalization suffers from it.

Gradient descent method. There are other related works in addition to those mentioned in Section 3.4. Bjorck et al. speculated that larger learning rates realized by batch normalization may help stochastic gradient descent avoid sharp minima and it leads to better generalization. Wei et al. estimated λmax\lambda_{max} and η\eta under a special type of batch-wise normalization. Because their normalization method approximates a chain rule of backpropagation by neglecting the contribution of mean subtraction, it suffers from pathological sharpness and requires smaller learning rates.

Neural tangent kernel. The FIM and NTK satisfy a kind of duality, and share the same non-zero eigenvalues. Our proofs on the eigenvalue statistics use NTK with standard parameterization, i.e., F∗F^{*} in Supplementary Material A.1. The NTK at random initialization is known to determine the gradient dynamics of a sufficiently wide DNN in function space. The sufficiently wide network can achieve a zero training error and it means that there is always a global minimum sufficiently close to random initialization. In the parameter space, Lee et al. proved that NTK dynamics is sufficiently approximated by the gradient descent of a linearized model expanded around random initialization θ0\theta_{0}: f(x;θt)=f(x;θ0)+∇θf(x;θ0)⊤ωtf(x;\theta_{t})=f(x;\theta_{0})+\nabla_{\theta}f(x;\theta_{0})^{\top}\omega_{t}, where ωt:=θt−θ0\omega_{t}:=\theta_{t}-\theta_{0} and tt means the step of the gradient descent. Naively speaking, this suggests that the optimization of the wide DNN approximately becomes convex and the loss landscape is dominated by a quadratic form with the FIM, i.e., ωt⊤Fωt\omega_{t}^{\top}F\omega_{t}.

Discussion

There remain a number of directions for extending our theoretical framework. Recent studies on wide DNNs have revealed that the NTK of random initialization dominates the training dynamics and even the performance of trained networks . Since the NTK is defined as a right-to-left reversed Gram matrix of the FIM under a special parameterization, the convergence speed of the training dynamics is essentially governed by the eigenvalues of the FIM at the random initialization. Analyzing these dynamics under normalization remains to be uncovered. For further analysis, random matrix theory will also be helpful in obtaining the whole eigenvalue spectrum or deriving tighter bounds of the largest eigenvalues. Although random matrix theory has been limited to a single layer or shallow networks , it will be an important direction to extend it to deeper and normalized networks.

There may be potential properties of normalization methods that are not detected in our framework. Kohler et al. analyzed the decoupling of the weight vector to its direction and length as in batch normalization and weight normalization. They revealed that such decoupling could contribute to accelerating the optimization. Bjorck et al. discussed that deep linear networks without bias terms suffer from the explosion of the feature vectors and speculated that batch normalization is helpful in reducing this explosion. This implies that batch normalization may be helpful to improve optimization performance even in a centered network. Yang et al. developed an excellent mean-field framework for batch normalization through all layers and found that the gradient explosion is induced by batch normalization in networks with extreme depth. Even if batch normalization alleviates pathological sharpness regarding the width, the coefficients of order evaluation can become very large when the network is extremely deep. It may cause another type of sharpness. It is also interesting to explore SGD training under normalization and quantify how the alleviation of sharpness affects appropriate sizes of learning rate and mini-batch, which have been mainly investigated in SGD training without normalization . Further studies on such phenomena in wide DNNs would be helpful for further understanding and development of normalization methods.

This work was partially supported by a Grant-in-Aid for Young Scientists (19K20366) from the Japan Society for the Promotion of Science (JSPS).

References

Supplementary Materials

We prepare the following two lemmas to prove the theorems in the main text.

An FIM is a P×PP\times P matrix, where PP is the dimension of all parameters. Define a P×CTP\times CT matrix RR by

Its columns are the gradients on each input, i.e., ∇θfk(t)\nabla_{\theta}f_{k}(t) (t=1,...,T)(t=1,...,T). One can represent an empirical FIM by

Let us refer to the following CT×CTCT\times CT matrix as a reversed FIM:

which is the right-to-left reversed Gram matrix of FF. This F∗F^{*} is essentially the same as the NTK . The FF and F∗F^{*} have the same non-zero eigenvalues by definition. Karakida et al. introduced F∗F^{*} to derive the eigenvalue statistics in Theorem 2.2. Technically speaking, they derived the eigenvalue statistics under the gradient independence assumption (Assumption 3.2). However, Yang recently succeeded in proving that this assumption is unnecessary. Therefore, Theorem 2.2 is free from this assumption.

To evaluate the effects of batch normalization, we need to take a more careful look into F∗F^{*} than done in previous studies. As shown in Supplementary Material B, the FIM under batch normalization in the last layer requires information on how fast backward order parameters asymptotically converge in the large MM limit. Let us introduce the following variables depending on MM:

Between the reversed FIM and convergence rate qq, we found that the following lemma holds. This lemma is a minor extension of Supplementary Material A in into the case without the gradient independence assumption.

Suppose a non-centered network and i.i.d. input samples generated by Eq. (3). When MM is sufficiently large, the F∗F^{*} can be partitioned into C2C^{2} block matrices whose (k,k′)(k,k^{\prime})-th block is a T×TT\times T matrix defined by

where q∗=min⁡{q,1/2}q^{*}=\min\{q,1/2\}, k,k′=1,...,Ck,k^{\prime}=1,...,C and δkk′\delta_{kk^{\prime}} is the Kronecker delta. The matrix KK has entries given by

Proof. We have the parameter set θ={Wijl,bil}\theta=\{W^{l}_{ij},b^{l}_{i}\} but the number of bias parameters (of O(M)O(M)) is much less than that of weight parameters (of O(M2)O(M^{2})). Therefore, the contribution of the FIM corresponding the bias terms are negligibly small in the large MM limit , and what we should analyze is weight parts of the FIM, that is, ∇Wijlfk\nabla_{W_{ij}^{l}}f_{k}. The (k,k′)(k,k^{\prime})-th block of F∗F^{*} has the (s,t)(s,t)-th entry as

for s,t=1,...,Ts,t=1,...,T. In the large MM limit, we can apply the central limit theorem to the feedforward propagation because the pre-activation uilu_{i}^{l} is a weighted sum of independent random weights : q^M,stl=q^stl+O(1/M)\hat{q}^{l}_{M,st}=\hat{q}^{l}_{st}+O(1/\sqrt{M}). This convergence rate of 1/21/2 is also known in non-asymptotic evaluations . We then have

The current work essentially differs from in the point that the evaluation of F∗F^{*} includes the convergence rate. The previous work investigated DNNs without any normalization method and such cases allow us to focus on the first term of the right-had side of Eq. (S.6). This is because the second term becomes asymptotically negligible in the large MM limit. In contrast, batch normalization in the last layer makes the first term comparable to the second term and requires careful evaluation of the second term. Thus, eigenvalues statistics become dependent on the convergence rate.

The previous work showed that the matrix KK in the first term of (S.6) determines the eigenvalue statistics such as mλm_{\lambda} and λmax\lambda_{max} in the large MM limit. The assumption of i.i.d. input samples makes the structure of matrix KK easy to analyze, i.e., all the diagonal terms take the same κ1\kappa_{1} and all the non-diagonal terms take κ2\kappa_{2}. Using this matrix KK, we can also derive the eigenvectors of F∗F^{*} corresponding to λmax\lambda_{max}:

It should be remarked that the above results require κ2>0\kappa_{2}>0. Technically speaking, the second term of Eq. (S.6) is negligible because κ1\kappa_{1} is positive by definition and κ2\kappa_{2} is also positive in a non-centered network. If one considers a centered network, however, the initialization of recurrence relations, i.e., q^st0=0\hat{q}^{0}_{st}=0, recursively yields

To prove Theorem 3.1, we use the eigenvector νk\nu_{k} obtained in Lemma A.2. The eigenspace of FF corresponding to λmax\lambda_{max} is constructed from νk\nu_{k}. Let us denote an eigenvector of FF as vv satisfying Fv=λmaxvFv=\lambda_{max}v. By multiplying R⊤R^{\top} by both sides, we have

for k=1,...,Ck=1,...,C. The first term of the right-hand side of Eq. (S.16) is of O(M1/2)O(M^{1/2}) in non-centered networks and asymptotically larger than the second term. Thus, we obtain Theorem 3.1.

B Batch normalization in last layer

The FIM under the mean subtraction (Eq. (15)) is expressed by

where Rˉ\bar{R} is a CT×PCT\times P matrix whose kk-th column is given by a vector ∇θμi/T\nabla_{\theta}\mu_{i}/\sqrt{T} ((i−1)T+1≤k≤iT(i-1)T+1\leq k\leq iT, i=1,2,...,Ci=1,2,...,C). Note that the hyperparameter βk\beta_{k} disappears since βk\beta_{k} is independent of θ\theta. Here, we define the projector

which satisfies G2=GG^{2}=G. Using this projector, we have RG=R−RˉRG=R-\bar{R} and

where ICI_{C} is a C×CC\times C identity matrix and ⊗\otimes is the Kronecker product. We introduce a reversed Gram matrix of the FIM under the mean subtraction:

Let us partition FL,mBN∗F_{L,mBN}^{*} into C2C^{2} block matrices and denote its (k,k′)(k,k^{\prime})-th block as a T×TT\times T matrix FL,mBN∗(k,k′)F^{*}_{L,mBN}(k,k^{\prime}). Substituting the F∗F^{*} (S.6) into the above, we obtain these blocks as

We assume T≥2T\geq 2 since T=1T=1 is trivial.

The mean of eigenvalues is asymptotically obtained by

First, we obtain a lower bound of λmax\lambda_{max}. In general, we have

where ν\nu is a CTCT-dimensional vector whose ((i−1)T+1)((i-1)T+1)-th entries are 1/2C1/\sqrt{2C}, ((i−1)T+2)((i-1)T+2)-th entries are −1/2C-1/\sqrt{2C}, and the others are (i=1,...,Ci=1,...,C). We then have

Next, we obtain an upper bound of λmax\lambda_{max}. In general, the maximum eigenvalue is denoted as the spectral norm ∣∣⋅∣∣2||\cdot||_{2}, i.e., λmax=∣∣FL,mBN∗∣∣2\lambda_{max}=||F^{*}_{L,mBN}||_{2}. Using the triangle inequality, we have

where ∣∣⋅∣∣F||\cdot||_{F} is the Frobenious norm. These lead to

Finally, sandwiching λmax\lambda_{max} by bounds (S.27) and (S.31), we asymptotically obtain

Note that κ1>κ2\kappa_{1}>\kappa_{2} holds in our settings. We can easily observe q^tl>q^stl\hat{q}^{l}_{t}>\hat{q}^{l}_{st} from the Cauchy–Schwarz inequality and it leads to κ1>κ2\kappa_{1}>\kappa_{2} (strictly speaking, when ϕ(x)\phi(x) is a constant function, its equality holds and we have q^tl=q^stl\hat{q}^{l}_{t}=\hat{q}^{l}_{st} and κ1=κ2\kappa_{1}=\kappa_{2}. However, we do not suppose the constant function as an ”activation” function and then κ1>κ2\kappa_{1}>\kappa_{2} holds).

This case requires a careful consideration of the O(M1−q∗)O(M^{1-q^{*}}) term in the reversed FIM (S.21). This is because the non-diagonal term of KL,mBNK_{L,mBN} asymptotically decreases to zero in the large MM limit and the O(M1−q∗)O(M^{1-q^{*}}) term becomes non-negligible. We found the following theorem without using the gradient independence assumption;

Suppose a non-centered network with the mean subtraction in the last layer (Eq. (15)) and i.i.d. input samples generated by Eq. (3). When T=O(M)T=O(M) with a constant ρ:=M/T\rho:=M/T, the largest eigenvalue in the large MM limit is asymptotically evaluated as

for q∗=min⁡{q,1/2}q^{*}=\min\{q,1/2\}, non-negative constants c1c_{1} and c2c_{2}.

Proof. To evaluate the largest eigenvalue, we use the second moment of the eigenvalues, i.e., sλ:=∑iPλi2/Ps_{\lambda}:=\sum_{i}^{P}\lambda_{i}^{2}/P. Because FL,mBN∗F^{*}_{L,mBN} is positive semi-definite, we have ∑iλi2=∑st((FL,mBN∗)st)2\sum_{i}\lambda_{i}^{2}=\sum_{st}((F^{*}_{L,mBN})_{st})^{2} and obtain

where we denote the second term of Eq. (S.21) as F0=O(M1−q∗)F_{0}=O(M^{1-q^{*}}). We then have

When T=O(Mp)T=O(M^{p}) (p≥0p\geq 0), the first term is of O(1/Mp)O(1/M^{p}), the second and third terms are of O(1/Mp+q∗)O(1/M^{p+q^{*}}), and the fourth term is of O(1/M2q∗)O(1/M^{2q^{*}}). Therefore, the second and third terms are negligible compared to the first term for all pp and q∗q^{*}. The fifth term is non-negative by definition. Although we can make the bounds of λmax\lambda_{max} for all pp, we focus on p=1p=1 for simplicity. In the large MM limit, we have asymptotically

The constant c0c_{0} comes from the fourth term of Eq. (S.35) and is non-negative.

The lower bound of λmax\lambda_{max} is derived from λmax≥∑iλi2/∑iλi=sλ/mλ\lambda_{max}\geq\sum_{i}\lambda_{i}^{2}/\sum_{i}\lambda_{i}=s_{\lambda}/m_{\lambda}, that is,

The upper bound comes from λmax≤∑iλi2=Psλ\lambda_{max}\leq\sqrt{\sum_{i}\lambda_{i}^{2}}=\sqrt{Ps_{\lambda}} and we have

The non-negative constants c1c_{1} and c2c_{2} come from c0c_{0}. ∎

Thus, we find that λmax\lambda_{max} is of order M1−2q∗M^{1-2q^{*}} at least and of order M1−q∗M^{1-q^{*}} at most. Since we have 0<q∗≤1/20<q^{*}\leq 1/2 by definition, the order of λmax\lambda_{max} is always lower than order of MM. Therefore, we can conclude that the mean subtraction alleviates the pathological sharpness for any qq.

From the theoretical perspective, we found that the gradient independence assumption achieves q=q∗=1/2q=q^{*}=1/2 and leads to a constant lower bound independent of MM.

The gradient independent assumption yields q=q∗=1/2q=q^{*}=1/2.

The bounds for λmax\lambda_{max} in Theorem 3.3 are immediately obtained from Theorem B.1 and Lemma B.2.

B.2 Mean subtraction and variance normalization

Define uˉk(t)=:ukL(t)−μk(θ)\bar{u}_{k}(t)=:u_{k}^{L}(t)-\mu_{k}(\theta). The derivatives of output units are given by

QQ is a CT×CTCT\times CT matrix whose (k,k′)(k,k^{\prime})-th block is given by a T×TT\times T matrix,

where ITI_{T} is a T×TT\times T identity matrix, σk2\sigma_{k}^{2} means σk(θ)2\sigma_{k}(\theta)^{2}, and Q(k,k)Q(k,k) is a projector to the vector uˉk\bar{u}_{k}. FL,BNF_{L,BN} and the following matrix have the same non-zero eigenvalues,

FL,BN∗F^{*}_{L,BN} is a CT×CTCT\times CT matrix and partitioned into C2C^{2} block matrices. Using Eq. (S.21), we obtain the (k,k′)(k,k^{\prime})-th block as

where the independence assumption yields q∗=1/2q^{*}=1/2. The first term is easy to evaluate,

by using the fact of ∑tuˉkL(t)=0\sum_{t}\bar{u}_{k}^{L}(t)=0. Suppose the case of ρ=M/T=const.\rho=M/T=const. Regarding the diagonal entries of Q(k,k′)KL,mBNQ(k,k^{\prime})K_{L,mBN}, the contribution of 1σk2T(1−1T)ukuk⊤\frac{1}{\sigma^{2}_{k}T}(1-\frac{1}{T})u_{k}u_{k}^{\top} is negligible to that of ITI_{T} in the large TT limit. Thus, we asymptotically obtain

The bounds of the largest eigenvalue are straightforwardly obtained from the second moment as in the deviation of Theorem B.1. Since the second moment sλ=∑iλi2/Ps_{\lambda}=\sum_{i}\lambda_{i}^{2}/P is given by a trace of the squared matrix in general, we have

The lower bound is given by λmax≥sλ/mλ\lambda_{max}\geq s_{\lambda}/m_{\lambda} and the upper bound by λmax≤Psλ\lambda_{max}\leq\sqrt{Ps_{\lambda}}.

C Batch normalization in middle layers

Batch normalization makes the chain of backward signals more complicated as follows. Suppose the tt-th input sample is given. Then, the activation in each layer depends not only on the tt-th sample but also on the whole of all samples. This is because batch normalization includes μl\mu^{l} and σl\sigma^{l}, which depend on the whole of all samples in the batch. Therefore, we should compute derivatives as

Recently, Yang et al. investigated a gradient explosion of the above chain rule in extremely deep networks although it requires a complicated formulation of mean field equations and is analytically intractable in general cases. In the following, we demonstrate an approach to batch normalization in the middle layers by avoiding the complicated analysis of the chain rule.

The derivative with respect to the LL-th layer is independent of the complicated chain of batch normalization because we do not normalize the last layer and have

where we used δk,iL(t;a)=δkiδta\delta^{L}_{k,i}(t;a)=\delta_{ki}\delta_{ta}. The lower bound of λmax\lambda_{max} is derived as follows:

The matrix KLK_{L} is defined by (KL)st:=q^t,BNL−1  (s=t),  q^st,BNL−1  (s≠t)(K_{L})_{st}:=\hat{q}^{L-1}_{t,BN}\ \ (s=t),\ \ \hat{q}^{L-1}_{st,BN}\ \ (s\neq t) where we denote feedroward order parameters for batch normalization as

The evaluation of the order parameters are shown in the following subsection. When the activation function is non-negative, the order paramters are positive. In particular, they are analytically tractable in ReLU networks.

Order parameters for batch normalization in the middle layers (S.66) require a careful integral over a TT-dimensional Gaussian distribution . This is because the pre-activation uˉil\bar{u}_{i}^{l} depends on all of uil(t)u_{i}^{l}(t) (t=1,...,T)(t=1,...,T) which share the same weight WijlW_{ij}^{l}. Therefore, we generally need the integration of ϕ(uˉil(t))\phi(\bar{u}_{i}^{l}(t)) over the TT-dimensional Gaussian distribution, that is,

where ul=(ul(1),ul(2),...,ul(T)){u}^{l}=({u}^{l}(1),{u}^{l}(2),...,{u}^{l}(T)) is a TT dimensional vector and ul∼N(0,σw2Σl−1){u}^{l}\sim\mathcal{N}(0,\sigma_{w}^{2}\Sigma_{l-1}). The T×TT\times T covariance matrix is defined by (Σl−1)st=q^st,BNl−1(\Sigma_{l-1})_{st}=\hat{q}^{l-1}_{st,BN} (s≠ts\neq t), q^t,BNl−1\hat{q}^{l-1}_{t,BN} (s=ts=t). These order parameters are positive when the activation function is non-negative (strictly speaking, non-negative and ϕ(x)>0\phi(x)>0 for certain xx).

Although the above integral is analytically intractable in many activation functions, Yang et al. gave profound insight into the integral. For instance, Corollary F.10 in revealed that the ReLU activation is more tractable, and we have

where J(x):=(1−x2+(π−arccos⁡(x))x)/πJ(x):=(\sqrt{1-x^{2}}+(\pi-\arccos(x))x)/\pi is known as the arccosine kernel. Wei et al. proposed a mean field approximation on the computation of order parameters for batch normalization, which is consistent with the above order parameters in the large TT limit. The previous study also proposed some methods to evaluate the order parameters in more general activation functions.

D Additional experiment on gradient descent training

E Layer normalization

We show that the order parameters under layer normalization are quite similar to those without the normalization. This is because the random weights and biases make the contribution of layer normalization relatively easy. In the large MM limit, we asymptotically have

for l=1,...,L−1l=1,...,L-1. Let us denote feedforward order parameters as

The same calculation as in the feedforward propagation without normalization leads to

The backward order parameters are also very similar to those without layer normalization. Let us consider the chain rule which appears in a FIM:

Ommiting index kk in δk,il(t)\delta_{k,i}^{l}(t) to avoid complicated notation, we have

where we define Pkil(t):=∂uˉkl∂uil(t)P_{ki}^{l}(t):=\frac{\partial\bar{u}^{l}_{k}}{\partial{u}^{l}_{i}}(t), which is an essential effect of layer normalization on the chain, and it becomes

where we substituted σl(t)=σl(s)=σw2q^tl−1+σb2\sigma^{l}(t)=\sigma^{l}(s)=\sqrt{\sigma_{w}^{2}\hat{q}^{l-1}_{t}+\sigma_{b}^{2}} and defined

The first term of Γk,kl(s,t)\Gamma_{k,k}^{l}(s,t) is dominant in the large MM limit because other terms are of order 1/M1/M. Then, we have

After applying the central limit theorem to ∑kϕk′l(s)ϕk′l(t)Ml\sum_{k}\frac{\phi^{\prime l}_{k}(s)\phi^{\prime l}_{k}(t)}{M_{l}}, we have

E.2 FIM

Denote the mean subtraction in the last layer as uˉk(t)=:ukL(t)−μL(t)\bar{u}_{k}(t)=:u_{k}^{L}(t)-\mu^{L}(t). The derivatives in the last layer are given by

where σ(t)2:=∑kuˉk(t)2/C\sigma(t)^{2}:=\sum_{k}\bar{u}_{k}(t)^{2}/C. Then, the FIM is given by

We can represent FL,BNF_{L,BN} in a matrix representation. Define a P×CTP\times CT matrix RR by

Its columns are the gradients on each input sample, i.e., ∇θukL(t)\nabla_{\theta}{u}_{k}^{L}(t) (t=1,...,T)(t=1,...,T). We then have

where Rˉ\bar{R} is defined as a CT×PCT\times P matrix whose ((k−1)T+t)((k-1)T+t)-th column is given by a vector ∇θμL(t)\nabla_{\theta}\mu^{L}(t) (t=1,...,Tt=1,...,T, k=1,...,Ck=1,...,C). We also defined a CT×CTCT\times CT matrix QQ whose (k,k′)(k,k^{\prime})-th block matrix is given by the following T×TT\times T matrix:

for k,k′=1,...,Ck,k^{\prime}=1,...,C. This Q(k,k′)Q(k,k^{\prime}) is a diagonal matrix. Compared to the matrix QQ in batch normalization (Eq. (S.47)), QQ in layer normalization is not block-diagonal. This is because layer normalization in the last layer yields interaction between different output units.

We introduce the following matrix which has the same non-zero eigenvalues as FLNF_{LN}:

This FmLN∗F^{*}_{mLN} corresponds to the mean subtraction in layer normalization. Its entries are given by

after doing the same calculation as Eq. (S.12) and using the order parameters obtained in Section E.1. We have q∗=1/2q^{*}=1/2 due to the gradient independence assumption. The reversed FIM becomes

where η1:=1T∑t1σ(t)2\eta_{1}:=\frac{1}{T}\sum_{t}\frac{1}{\sigma(t)^{2}}.

The largest eigenvalue is evaluated using the second moment of the eigenvalues. Since the second moment sλ=∑iλi2/Ps_{\lambda}=\sum_{i}\lambda_{i}^{2}/P is given by a trace of the squared matrix in general, we have

where we define g(t,t′):=1C∑auˉa(t)uˉa(t′)σ(t)σ(t′).g(t,t^{\prime}):=\frac{1}{C}\frac{\sum_{a}\bar{u}_{a}(t)\bar{u}_{a}(t^{\prime})}{\sigma(t)\sigma(t^{\prime})}. Substituting Eq. (S.112) into Eq. (S.109), we obtain

where η2:=∑tT1σ(t)4\eta_{2}:=\frac{\sum_{t}}{T}\frac{1}{\sigma(t)^{4}} and η3:=1T2∑t,t′g(t,t′)σ(t)σ(t′).\eta_{3}:=\frac{1}{T^{2}}\sum_{t,t^{\prime}}\frac{g(t,t^{\prime})}{\sigma(t)\sigma(t^{\prime})}. The lower bound is given by λmax≥sλ/mλ\lambda_{max}\geq s_{\lambda}/m_{\lambda} and the upper bound by λmax≤Psλ\lambda_{max}\leq\sqrt{Ps_{\lambda}}.

Remark on C=2C=2: Because we have a special symmetry, i.e., uˉ1(t)=−uˉ2(t)=(u1L(t)−u2L(t))/2\bar{u}_{1}(t)=-\bar{u}_{2}(t)=(u_{1}^{L}(t)-u_{2}^{L}(t))/2 in C=2C=2, the gradient (Eq. (S.88)) becomes zero. This is caused by the mean subtraction and variance normalization in the last layer. This makes the FIM a zero matrix. The case of C>3C>3 is non-trivial and the FIM becomes non-zero, as we revealed. Similarly, the gradient (Eq. (S.43)) in batch normalization becomes zero when T=2T=2 due to the same symmetry . Such an exceptional case of batch normalization is not our interest because we focus on the sufficiently large TT in Eq. (19).