Disentangling Trainability and Generalization in Deep Neural Networks

Lechao Xiao, Jeffrey Pennington, Samuel S. Schoenholz

Introduction

Machine learning models based on deep neural networks have attained state-of-the-art performance across a dizzying array of tasks including vision (Cubuk et al., 2019), speech recognition (Park et al., 2019), machine translation (Bahdanau et al., 2014), chemical property prediction (Gilmer et al., 2017), diagnosing medical conditions (Raghu et al., 2019), and playing games (Silver et al., 2018). Historically, the rampant success of deep learning models has lacked a sturdy theoretical foundation: architectures, hyperparameters, and learning algorithms are often selected by brute force search (Bergstra & Bengio, 2012) and heuristics (Glorot & Bengio, 2010). Recently, significant theoretical progress has been made on several fronts that have shown promise in making neural network design more systematic. In particular, in the infinite width (or channel) limit, the distribution of functions induced by neural networks with random weights and biases has been precisely characterized before, during, and after training.

The study of infinite networks dates back to seminal work by Neal (1994) who showed that the distribution of functions given by single hidden-layer networks with random weights and biases in the infinite-width limit are Gaussian Processes (GPs). Recently, there has been renewed interest in studying random, infinite, networks starting with concurrent work on “conjugate kernels” (Daniely et al., 2016; Daniely, 2017) and “mean-field theory” (Poole et al., 2016; Schoenholz et al., 2017). Among numerous contributions, the pair of papers by Daniely et al. argued that the empirical covariance matrix of pre-activations becomes deterministic in the infinite-width limit and called this the conjugate kernel of the network. Meanwhile, from a mean-field perspective, the latter two papers studied the properties of these limiting kernels. In particular, the spectrum of the conjugate kernel of wide, fully-connected, networks approaches a well-defined and data-independent limit when the depth exceeds a certain scale, ξ\xi. Networks with tanh⁡\tanh-nonlinearities (among other bounded activations) exhibit a phase transition between two limiting spectral distributions of the conjugate kernel as a function of their hyperparameters with ξ\xi diverging at the transition. It was additionally hypothesized that networks were un-trainable when the conjugate kernel was sufficiently close to its limit.

Since then this analysis has been extended to include a wide range for architectures such as convolutions (Xiao et al., 2018), recurrent networks (Chen et al., 2018; Gilboa et al., 2019), networks with residual connections (Yang & Schoenholz, 2017), networks with quantized activations (Blumenfeld et al., 2019), the spectrum of the fisher (Karakida et al., 2018), a range of activation functions (Hayou et al., 2018), and batch normalization (Yang et al., 2019). In each case, it was observed that the spectra of the kernels correlated strongly with whether or not the architectures were trainable. While these papers studied the properties of the conjugate kernels, especially the spectrum in the large-depth limit, a branch of concurrent work took a Bayesian perspective: that many networks converge to Gaussian Processes as their width becomes large (Lee et al., 2018; Matthews et al., 2018; Novak et al., 2019b; Garriga-Alonso et al., 2018; Yang, 2019). In this case, the Conjugate Kernel was referred to as the Neural Network Gaussian Process (NNGP) kernel, which is used to train neural networks in a fully Bayesian fashion. As such, the NNGP kernel characterizes performance of the corresponding NNGP.

Together this work offered a significant advance to our understanding of wide neural networks; however, this theoretical progress was limited to networks at initialization or after Bayesian posterior estimation and provided no link to gradient descent. Moreover, there was some preliminary evidence that suggested the situation might be more nuanced than the qualitative link between the NNGP spectrum and trainability might suggest. For example, Philipp et al. (2017) showed that deep tanh⁡\tanh FCNs could be trained after the kernel reached its large-depth, data-independent, limit but that these networks did not generalize to unseen data.

Recently, significant theoretical clarity has been reached regarding the relationship between the GP prior and the distribution following gradient descent. In particular, Jacot et al. (2018) along with followup work (Lee et al., 2019; Chizat et al., 2019) showed that the distribution of functions induced by gradient descent for infinite-width networks is a Gaussian Process with a particular compositional kernel known as the Neural Tangent Kernel (NTK). In addition to characterizing the distribution over functions following gradient descent in the wide network limit, the learning dynamics can be solved analytically throughout optimization.

In this paper, we leverage these developments and revisit the relationship between architecture, hyperparameters, trainability, and generalization in the large-depth limit for a variety of neural networks. In particular, we make the following contributions:

Trainability. We compute the large-depth asymptotics of several quantities related to trainability, including the largest/smallest eigenvalue of the NTK, λmax/min\lambda_{\text{max}/\text{min}}, and the condition number κ=λmax/λmin\kappa=\lambda_{\text{max}}/\lambda_{\text{min}}; see Table 1.

Generalization. We characterize the mean predictor P(Θ)P(\Theta), which is intimately related to the prediction of wide neural networks on the test set following gradient descent training. As such, the mean predictor is intimately related to the model’s ability to generalize. In particular, we argue that networks fail to generalize if the mean predictor becomes data-independent.

We show that the ordered and chaotic phases identified in Poole et al. (2016) lead to markedly different limiting spectra of the NTK. In the ordered phase the trainability of neural networks degrades at large depths, but their ability to generalize persists. By contrast, in the chaotic phase we show that trainability improves with depth, but generalization degrades and neural networks behave like hash functions.

A corollary of these differences in the spectra is that, as a function of depth, the optimal learning rates ought to decay exponentially in the chaotic phase, linearly on the order-to-chase trainsition line, and remain roughly a constant in the ordered phase.

We examine the differences in the above quantities for fully-connected networks (FCNs) and convolutional networks (CNNs) with and without pooling and precisely characterize the effect of pooling on the interplay between trainability, generalization, and depth.

In each case, we provide empirical evidence to support our theoretical conclusions. Together these results provide a complete, analytically tractable, and dataset-independent theory for learning in very deep and wide networks. Philosophically, we find that trainability and generalization are distinct notions that are, at least in this case, at odds with one another. Indeed, good conditioning of the NTK (which is a necessary condition for training) seems necessarily to lead to poor generalization performance. It will be interesting to see whether these results carry over in shallower and narrower networks. The tractable nature of the wide and deep regime leads us to conclude that these models will be an interesting testbed to investigate various theories of generalization in deep learning.

Related Work

Recent work Jacot et al. (2018); Du et al. (2018b); Allen-Zhu et al. (2018); Du et al. (2018a); Zou et al. (2018) and many others proved global convergence of over-parameterized deep networks by showing that the NTK essentailly remains a constant over the course of training. However, in a different scaling limit the NTK changes over the course of training and global convergence is much more difficult to obtain and is known for neural networks with one hidden layer Mei et al. (2018); Chizat & Bach (2018); Sirignano & Spiliopoulos (2018); Rotskoff & Vanden-Eijnden (2018). Therefore, understanding the training and generalization properties in this scaling limit remains a very challenging open question.

Another two excellent recent works (Hayou et al., 2019; Jacot et al., 2019) also study the dynamics of Θ(l)(x,x′)\Theta^{(l)}(x,x^{\prime}) for FCNs (and deconvolutions in (Jacot et al., 2019)) as a function of depth and variances of the weights and biases. (Hayou et al., 2019) investigates role of activation functions (smooth v.s. non-smooth) and skip-connection. (Jacot et al., 2019) demonstrate that batch normalization helps remove the “ordered phase” (as in (Yang et al., 2019)) and a layer-dependent learning rate allows every layer in a network to contribute to learning.

Background

Equation 2 describes a dynamical system on positive semi-definite matrices K\mathcal{K}. It was shown in Poole et al. (2016) that fixed points, K∗(x,x′)\mathcal{K}^{*}(x,x^{\prime}), of these dynamics exist such that lim⁡l→∞K(l)(x,x′)=K∗(x,x′)\lim_{l\to\infty}\mathcal{K}^{(l)}(x,x^{\prime})=\mathcal{K}^{*}(x,x^{\prime}) with K∗(x,x′)=q∗[δx,x′+c∗(1−δx,x′)]\mathcal{K}^{*}(x,x^{\prime})=q^{*}[\delta_{x,x^{\prime}}+c^{*}(1-\delta_{x,x^{\prime}})] independent of the inputs xx and x′x^{\prime}. The values of q∗q^{*} and c∗c^{*} are determined by the hyperparameters, σw\sigma_{w} and σb\sigma_{b}. However Equation 2 admits multiple fixed points (e.g. c∗=0,1c^{*}=0,1) and the stability of these fixed points plays a significant role in determining the properties of the network. Generically, there are large regions of the (σw,σb)(\sigma_{w},\sigma_{b}) plane in which the fixed-point structure is constant punctuated by curves, called phase transitions, where the structure changes; see Fig 5 for tanh⁡\tanh-networks.

The rate at which K(x,x′)\mathcal{K}(x,x^{\prime}) approaches or departs K∗(x,x′)\mathcal{K}^{*}(x,x^{\prime}) can be determined by expanding Equation 2 about its fixed point, δK(x,x′)=K(x,x′)−K∗(x,x′)\delta\mathcal{K}(x,x^{\prime})=\mathcal{K}(x,x^{\prime})-\mathcal{K}^{*}(x,x^{\prime}) to find

for train and test points respectively; see Section 2 in Lee et al. (2019). Here Θtest, train\Theta_{\text{test, train}} denotes the NTK between the test inputs XtestX_{\text{test}} and training inputs XtrainX_{\text{train}} and Θtrain, train\Theta_{\text{train, train}} is defined similarly. Since Θ^\hat{\Theta} converges to Θ\Theta as the network’s width approaches infinity, the gradient flow dynamics of real network also converge to the dynamics described by Equation 5 and Equation 6 (Jacot et al., 2018; Lee et al., 2019; Chizat et al., 2019; Yang, 2019; Arora et al., 2019; Huang & Yau, 2019). As the training time, tt, tends to infinity we note that these equations reduce to μ(Xtrain)=Ytrain\mu(X_{\text{train}})=Y_{\text{train}} and μ(Xtest)=Θtest, trainΘtrain, train−1Ytrain\mu(X_{\text{test}})=\Theta_{\text{test, train}}\Theta_{\text{train, train}}^{-1}Y_{\text{train}}. Consequently we call

the “mean predictor”. We can also compute the mean predictor of the NNGP kernel, P(K)P(\mathcal{K}), which analogously can be used to find the mean of the posterior after Bayesian inference. We will discuss the connection between the mean predictor and generalization in the next section.

In addition to showing that the NTK describes networks during gradient descent, Jacot et al. (2018) showed that the NTK could be computed in closed form in terms of T\mathcal{T}, T˙\dot{\mathcal{T}}, and the NNGP as,

where Θ(l)\Theta^{(l)} is the NTK for the pre-activations at layer-ll.

Metrics for Trainability and Generalization at Large Depth

We begin by discussing the interplay between the conditioning of Θtrain, train\Theta_{\text{train, train}} and the trainability of wide networks. We can write Equation 5 in terms of the spectrum of Θtrain, train\Theta_{\text{train, train}}. To do this we write the eigendecomposition of Θtrain, train\Theta_{\text{train, train}} as Θtrain, train=UTDU\Theta_{\text{train, train}}=U^{T}DU with DD a diagonal matrix of eigenvalues and UU a unitary matrix. In this case Equation 5 can be written as,

We will see that at large depths, the spectrum of Θtrain, train\Theta_{\text{train, train}} typically features a single large eigenvalue, λmax\lambda_{\text{max}}, and then a gap that is large compared with the rest of the spectrum. We therefore will often refer to a typical eigenvalue in the bulk as λbulk\lambda_{\rm bulk} and approximate the condition number as κ=λmax/λbulk\kappa=\lambda_{\text{max}}/\lambda_{\rm bulk}.

We now turn our attention to generalization. At large depths, we will see that Θtest, train(l)\Theta^{(l)}_{\text{test, train}} and Θtrain, train(l)\Theta^{(l)}_{\text{train, train}} converge their fixed points independent of the data distribution. Consequently it is often the case that P(Θ∗)P(\Theta^{*}) will be data-independent and the network will fail to generalize. In this case, by symmetry, it is necessarily true that P(Θ∗)P(\Theta^{*}) will be a constant matrix. Contracting this matrix with a vector of labels YtrainY_{\text{train}} that have been standardized to have zero mean it will follow that P(Θ∗)Ytrain=0P(\Theta^{*})Y_{\text{train}}=0 and the network will output zero in expectation on all test points. Clearly, in this setting the network will not be able to generalize. At large, but finite, depths the generalization performance of the network can be quantified by considering the rate at which P(Θ(l))YtrainP(\Theta^{(l)})Y_{\text{train}} decays to zero. There are cases, however, where despite the data-independence of Θ∗\Theta^{*}, lim⁡l→∞P(Θ(l))Ytrain\lim_{l\to\infty}P(\Theta^{(l)})Y_{\text{train}} remains nonzero and the network can continue to generalize even in the asymptotic limit. In either case, we will show that precisely characterizing P(Θ(l))YtrainP(\Theta^{(l)})Y_{\text{train}} allows us to understand exactly where networks can, and cannot, generalize.

Our goal is therefore to characterize the evolution of the two metrics κ(l)\kappa^{(l)} and P(Θ(l))P(\Theta^{(l)}) in ll. We follow the methodology outlined in Schoenholz et al. (2017); Xiao et al. (2018) to explore the spectrum of the NTK as a function of depth. We will use this to make precise predictions relating trainability and generalization to the hyperparameters (σw,σb,l)(\sigma_{w},\sigma_{b},l). Our main results are summarized in Table 1 which describes the evolution of λmax(l)\lambda_{\rm max}^{(l)} (the largest eigenvalue of Θ(l)\Theta^{(l)}), λbulk(l)\lambda_{\rm bulk}^{(l)} (the remaining eigenvalues), κ(l)\kappa^{(l)}, and P(Θ(l))P(\Theta^{(l)}) as a function of depth for three different network configurations (the ordered phase, the chaotic phase, and the phase transition). We study the dependence on: the size of the training set, mm; the choices of architecture including fully-connected networks (FCN), convolutional networks with flattening (CNN-F), and convolutions with pooling (CNN-P); and the size, dd, of the window in the pooling layer (which we always take to be the penultimate layer).

Before discussing the methodology it is useful to first give a qualitative overview of the phenomenology. We find identical phenomenology between FCNs and CNN-F architectures. In the ordered phase, Θ(l)→p∗11T\Theta^{(l)}\to p^{*}\bm{1}\bm{1}^{T}, λmax(l)→mp∗\lambda_{\rm max}^{(l)}\to mp^{*} and λbulk(l)=O(lχ1l)\lambda_{\rm bulk}^{(l)}=\mathcal{O}(l\chi_{1}^{l}). At large depths since χ1<1\chi_{1}<1 it follows that κ(l)≳mp∗/(lχ1l)\kappa^{(l)}\gtrsim mp^{*}/(l\chi_{1}^{l}) and so the condition number diverges exponentially quickly. Thus, in the ordered phase we expect networks not to be trainable (or, specifically, the time they take to learn will grow exponentially in their depth). Here P(Θ(l))P(\Theta^{(l)}) converges to a data dependent constant independent of depth; thus, in the ordered phase networks fail to train but can generalize indefinitely.

By contrast, in the chaotic phase we see that there is no gap between λmax(l)\lambda_{\rm max}^{(l)} and λbulk(l)\lambda_{\rm bulk}^{(l)} and networks become perfectly conditioned and are trainable everywhere. However, in this regime we see that the mean predictor scales as l(χc∗/χ1)ll(\chi_{c^{*}}/\chi_{1})^{l}. Since in the chaotic phase χc∗<1\chi_{c^{*}}<1 and χ1>1\chi_{1}>1 it follows that P(Θ(l))→0P(\Theta^{(l)})\to 0 over a depth ξ∗=−1/log⁡(χc∗/χ1)\xi_{*}=-1/\log(\chi_{c^{*}}/\chi_{1}). Thus, in the chaotic phase, networks fail to generalize at a finite depth but remain trainable indefinitely. Finally, introducing pooling modestly augments the depth over which networks can generalize in the chaotic phase but reduces the depth in the ordered phase. We will explore all of these predictions in detail in section 7.

A Toy Example: RBF Kernel

To provide more intuition about our analysis, we present a toy example using RBF kernels which already shares some core observations for deep neural networks. Consider a Gaussian process along with the RBF kernel given by,

where x,x′∈Xtrainx,x^{\prime}\in X_{\text{train}} along with a bandwidth h>0h>0. Note that Kh(x,x)=1K_{h}(x,x)=1 for all hh and xx. Considering the following two cases.

If the bandwidth is given by h=2lh=2^{l} and l→∞l\to\infty, then Kh(x,x′)≈1−2−l∥x−x′∥22K_{h}(x,x^{\prime})\approx 1-2^{-l}\|x-x^{\prime}\|^{2}_{2} which converges to 11 exponentially fast. Thus, the largest eigenvalue of KhK_{h} is λmax≈∣Xtrain∣\lambda_{\text{max}}\approx|X_{\text{train}}| and the bulk is of order λbulk≈2−l\lambda_{\text{bulk}}\approx 2^{-l}. Thus the condition number κ≳2l\kappa\gtrsim 2^{l} which diverges with ll. We will see in the Ordered Phase Θ(l)\Theta^{(l)} behaves qualitatively similar to this setting.

On the other hand, if the bandwidth is given by h=1/lh=1/l and l→∞l\to\infty then the off-diagonals Kh(x,x′)=exp⁡(−l∥x−x′∥22)→0K_{h}(x,x^{\prime})=\exp(-l\|x-x^{\prime}\|_{2}^{2})\to 0. For large ll, KhK_{h} is very close to the identity matrix and the condition number of it is almost 1. In the Chaotic Phase, Θ(l)\Theta^{(l)} is qualitatively similar to KhK_{h}.

Large-Depth Asymptotics of the NNGP and NTK

We now give a brief derivation of the results in Table 1. Details can be found in Sec.B, D in the appendix. To simplify notation we will discuss fully-connected networks and then extend the results to CNNs with pooling (CNN-P) and without pooling (CNN-F).

As in Sec. 3, we will be concerned with the fixed points of Θ\Theta as well as the linearization of Equation 8 about its fixed point. Recall that the fixed point structure is invariant within a phase so it suffices to consider the ordered phase, the chaotic phase, and the critical line separately. In cases where a stable fixed point exists, we will describe how Θ\Theta converges to the fixed point. We will see that in the chaotic phase and on the critical line, Θ\Theta has no stable fixed point and in that case we will describe its divergence. As above, in each case the fixed points of Θ\Theta have a simple structure with Θ∗=p∗((1−c^∗)Id+c^∗11T)\Theta^{*}=p^{*}((1-\hat{c}^{*})\textbf{Id}+\hat{c}^{*}\bm{1}\bm{1}^{T}).

To simplify the forthcoming analysis, without a loss of generality, we assume the inputs are normalized to have variance q∗q^{*} It has been observed in previous works (Poole et al., 2016; Schoenholz et al., 2017) that the diagonals converge much faster than the off-diagonals for tanh⁡\tanh- or erf- networks.. As such, we can treat T\mathcal{T} and T˙\dot{\mathcal{T}}, restricted on {K(l)}l\{\mathcal{K}^{(l)}\}_{l}, as a point-wise functions. To see this note that with this normalization K(l)(x,x)=q∗\mathcal{K}^{(l)}(x,x)=q^{*} for all ll and xx. It follows that both T(K(l+1))(x,x′)\mathcal{T}(\mathcal{K}^{(l+1)})(x,x^{\prime}) and T˙(K(l+1))(x,x′)\dot{\mathcal{T}}(\mathcal{K}^{(l+1)})(x,x^{\prime}) depend only on K(l)(x,x′)\mathcal{K}^{(l)}(x,x^{\prime}).

Since all of the off-diagonal elements approach the same fixed point at the same rate, we use qab(l)≡K(l)(x,x′)q^{(l)}_{ab}\equiv\mathcal{K}^{(l)}(x,x^{\prime}) and pab(l)≡Θ(l)(x,x′)p^{(l)}_{ab}\equiv\Theta^{(l)}(x,x^{\prime}) to denote any off diagonal entry of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} respectively. We will similarly use qab∗q^{*}_{ab} and pab∗p_{ab}^{*} to denote the limits, lim⁡l→∞qab(l)=qab∗=c∗q∗\lim_{l\to\infty}q^{(l)}_{ab}=q^{*}_{ab}=c^{*}q^{*} and lim⁡l→∞pab(l)=pab∗=c^∗p∗\lim_{l\to\infty}p^{(l)}_{ab}=p_{ab}^{*}=\hat{c}^{*}p^{*}. Finally, although the diagonal entries of K(l)\mathcal{K}^{(l)} are all q∗q^{*}, the diagonal entries of Θ(l)\Theta^{(l)} can vary and we denote them p(l)p^{(l)}.

In what follows, we split the discussion into three sections according to the values of χ1≡σω2T˙(q∗)\chi_{1}\equiv\sigma_{\omega}^{2}\dot{\mathcal{T}}(q^{*}) recalling that in Poole et al. (2016); Schoenholz et al. (2017) it was shown that χ1\chi_{1} controls the fixed point structure. In each section, we analyze the evolution of (1) the entries of Θ(l)\Theta^{(l)}, i.e., p(l)p^{(l)}, pab(l)p^{(l)}_{ab}, (2) the spectrum λmax(l)\lambda_{\rm max}^{(l)} and λbulk(l)\lambda_{\rm bulk}^{(l)}, (3) the trainability and generalization metrics κ(l)\kappa^{(l)} and P(Θ(l))P(\Theta^{(l)}), and finally (4) discuss the impact on finite width networks.

The chaotic phase is so-named because it has a stable fixed-point c∗<1c^{*}<1; as such similar inputs become increasingly uncorrelated as they pass through the network. Our first result is to show that (see Sec. B.1),

Note that χc∗\chi_{c^{*}} controls the convergence of the qab(l)q^{(l)}_{ab} and is always less than 1 in the chaotic phase (Poole et al., 2016; Schoenholz et al., 2017; Xiao et al., 2018). Since χ1>1\chi_{1}>1, p(l)p^{(l)} diverges with rate χ1l\chi_{1}^{l} while pab(l)p^{(l)}_{ab} remains finite. It follows that (p(l))−1Θ(l)→Id(p^{(l)})^{-1}\Theta^{(l)}\to\textbf{Id} as l→∞l\to\infty. Thus, in the chaotic phase, the spectrum of the NTK for very deep networks approaches the diverging constant multiplying the identity. This implies

Figure 1(a) plots the evolution of κ(l)\kappa^{(l)} in this phase, confirming κ(l)→1\kappa^{(l)}\to 1 for all three different architectures (FCN, CNN-F and CNN-P).

We now describe the asymptotic behavior of the mean predictor. Since Θtest, trainl\Theta_{\text{test, train}}^{l} has no diagonal elements, it follows that it remains finite at large depths and so P(Θ∗)Ytrain=0P(\Theta^{*})Y_{\text{train}}=0. It follows that in the chaotic phase, the predictions of asymptotically deep neural networks on unseen test points will converge to zero exponentially quickly (see Sec. D.1),

Neglecting the relatively slowly varying polynomial term, this implies that we expect chaotic networks to fail to generalize when their depth is much larger than a scale set by ξ∗=−1/log⁡(χc∗/χ1)\xi_{*}=-1/\log(\chi_{c^{*}}/\chi_{1}). We confirm this scaling in Fig 1(e).

We confirm these predictions for finite-width neural network training using SGD as well as gradient-flow on infinite networks in the experimental results; see Fig 2.

The ordered phase is defined by the stability of the c∗=1c^{*}=1 fixed point. Here disparate inputs will end up converging to the same output at the end of the network. We show in Sec. B.2 that elements of the NNGP kernel and NTK have asymptotic dynamics given by,

where p∗=q∗/(1−χ1)p^{*}=q^{*}/(1-\chi_{1}). Here all of the entries of Θ(l)\Theta^{(l)} converge to the same value, p∗p^{*}, and the limiting kernel has the form Θ∗=p∗1n1mT\Theta^{*}=p^{*}\bf{1}_{n}\bf{1}_{m}^{T} where 1m\bf{1_{m}} is the all-ones vector of dimension mm (typically mm will correspond to the number of datapoints in the training set). The NNGP kernel has the same structure with p∗↔q∗p^{*}\leftrightarrow q^{*}. Consequently both the NNGP kernel and the NTK are highly singular and feature a single non-zero eigenvalue, λmax=mp∗\lambda_{\text{max}}=mp^{*}, with eigenvector 1m\bf{1_{m}}.

For large-but-finite depths, Θ(l)\Theta^{(l)} has (approximately) two eigenspaces: the first eigenspace corresponds to finite-depth corrections to λmax\lambda_{\text{max}},

The second eigenspace comes from lifting the degenerate zero-modes has dimension (m−1)(m-1) with eigenvalues that scale like λbulk(l)=O(p(l)−pab(l))=O(lχ1l).\lambda_{\rm bulk}^{(l)}=\mathcal{O}(p^{(l)}-p^{(l)}_{ab})=\mathcal{O}(l\chi_{1}^{l}). It follows that κ(l)≳(lχ1l)−1\kappa^{(l)}\gtrsim(l\chi_{1}^{l})^{-1} and so the conditioning number explodes exponentially quickly. We confirm the presence of the 1/l1/l correction term in κ(l)\kappa^{(l)} by plotting χ1lκ(l)\chi_{1}^{l}\kappa^{(l)} against ll in Figure 1(b). Neglecting this correction, we expect networks in the ordered phase to become untrainable when their depth exceeds a scale given by ξ1=−1/log⁡χ1\xi_{1}=-1/\log\chi_{1}.

We now turn our discussion to the mean predictor. Equation 14 shows that we can write the finite-depth corrections to the NTK as Θ(l)=p∗11T+A(l)lχ1l\Theta^{(l)}=p^{*}{\bf 11}^{T}+\bm{A}^{(l)}l\chi_{1}^{l}. Here A(l){\bm{A}}^{(l)} is the data-dependent piece that lifts the zero eigenvalues. In the appendix, A(l)\bm{A}^{(l)} converges to A\bm{A} as l→∞l\to\infty; see Lemma 2. In Sec. D.3 we show that despite the singular nature of Θ∗\Theta^{*}, the mean has a well-defined limit as,

where A^\hat{\bm{A}} is some correction term. Thus, the mean predictor remains well-behaved and data dependent even in the infinite-depth limit. Thus, we suspect that networks in the ordered phase should be able to generalize whenever they can be trained. We confirm the asymptotic data-dependence of the mean predictor in Fig 1(f).

On the critical line the c∗=1c^{*}=1 fixed point is marginally stable and dynamics become powerlaw. Here, both the diagonal and the off-diagonal elements of Θ(l)\Theta^{(l)} diverge linearly in the depth with 1lΘ(l)→q∗3(11T+2Id)\frac{1}{l}\Theta^{(l)}\to\frac{q^{*}}{3}(\bm{1}\bm{1}^{T}+2\textbf{Id}). The condition number κ(l)\kappa^{(l)} converges to a finite value and the network is always trainable. However, the mean predictor decreases linearly with depth. In particular we show in Sec. B.3,

For large ll it follows that Θ(l)\Theta^{(l)} essentially has two eigenspaces: one has dimension one and the other has dimension (m−1)(m-1) with

It follows that the condition number κ(l)=m+22+mO(l−1)→m+22\kappa^{(l)}=\frac{m+2}{2}+m\mathcal{O}(l^{-1})\to\frac{m+2}{2} as l→∞l\to\infty. Unlike in the chaotic and ordered phases, here κ(l)\kappa^{(l)} converges with rate O(l−1)\mathcal{O}(l^{-1}). Figure 1(c) confirms the κ(l)→m+22\kappa^{(l)}\to\frac{m+2}{2} for both FCN and CNN-F (the global average pooling in CNN introduces a correction term that we will discuss below). A similar calculation gives P(Θ(l))=O(l−1)P(\Theta^{(l)})=\mathcal{O}(l^{-1}) on the critical line.

In summary, κ(l)\kappa^{(l)} converges to a finite number and the network ought to be trainable for arbitrary depth but the mean predictor P(Θ(l))P(\Theta^{(l)}) decays as a powerlaw. Decay as l−1l^{-1} is much slower than exponential and is slow on the scale of neural networks. This explains why critically initialized networks with thousands of layers could still generalize (Xiao et al., 2018).

4 The Effect of Convolutions

The above theory can be extended to CNNs. We will provide an informal description here, with details in Sec. F. For an input-images of size (m,k,k,3)(m,k,k,3) the NTK and NNGP kernels will have shape (m,k,k,m,k,k)(m,k,k,m,k,k) and will contain information about the covariance between each pair of pixels in each image. For convenience we will let d=k2d=k^{2}. In the large depth setting deviations of both kernels from their fixed point decomposes via Fourier transform in the spatial dimensions as,

where qq denotes the Fourier mode with q=0q=0 being the zero-frequency (uniform) mode and ρq\rho_{q} are eigenvalues of certain convolution operator. Here δΘ(l)(q)\delta\Theta^{(l)}(q) are deviations from the fixed-point for the qqth mode with δΘ(l)(q)∝δΘFCN(l)\delta\Theta^{(l)}(q)\propto\delta\Theta^{(l)}_{\text{FCN}} the fully-connected deviation described above. We show that ρq=0=1\rho_{q=0}=1 and ∣ρq≠0∣<1|\rho_{q\neq 0}|<1 which implies that asymptotically the nonuniform modes become subleading as ρql→0\rho_{q}^{l}\to 0. Thus, at large depths different pixels evolve identically as FCNs.

In Sec. F.2 we discuss the differences that arise when one combines a CNN with a flattening layer compared with an average pooling layer at the readout. In the case of flattening, the pixel-pixel correlations are discarded and ΘCNN−F(l)≈ΘFCN(l)\Theta^{(l)}_{\rm CNN-F}\approx\Theta^{(l)}_{\rm FCN}. The plots in the first row of Figure 1 confirm that the κ(l)\kappa^{(l)} of ΘCNN−F(l)\Theta^{(l)}_{\rm CNN-F} and of ΘFCN(l)\Theta^{(l)}_{\rm FCN} evolve almost identically in all phases. Note that this clarifies an empirical observation in Xiao et al. (2018) (Figure 3 of Xiao et al. (2018)) that test performance of critically initialized CNNs degrades towards that of FCNs as depth increases. This is because (i) in the large width limit, the prediction of neural networks is characterized by the NTK and (ii) the NTKs of the two models are almost identical for large depth. However, when CNNs are combined with global average pooling a correction to the spectrum of the NTK (NNGP) emerges oweing to pixel-pixel correlations; this alters the dynamics of κ(l)\kappa^{(l)} and P(Θ(l))P(\Theta^{(l)}). In particular, we find that global average pooling increases κ(l)\kappa^{(l)} by a factor of dd in the ordered phase and on the critical line; see Table 1 for the exact correction as well as Figures 1(d) for experimental evidence of this correction.

5 Dropout, Relu and Skip-connection

Adding a dropout to the penultimate layer has a similar effect to adding a diagonal regularization term to the NTK, which significantly improves the conditioning of the NTK in the ordered phase. In particular, adding a single dropout layer can cause κ(l)\kappa^{(l)} to converge to a finite κ∗\kappa^{*} rather than diverges exponentially; see Figure 4 and Sec. E.

For critically initialized Relu networks (aka, He’s initialization (He et al., 2015)), the entries of the NTK also diverges linearly and κ(l)→m+33\kappa^{(l)}\to\frac{m+3}{3} and P(Θ(l))=O(1/l)P(\Theta^{(l)})=\mathcal{O}(1/l); see Table 2 and Figure 3. In addition, adding skip-connections makes all entries of the NTK to diverge exponentially, resulting exploding of gradients. However, we find that skip connections do not alter the dynamics of κ(l)\kappa^{(l)}. Finally, layer normalization could help address the issue of exploding of gradients; see Sec. C.

Experiments

Evolution of κ(l)\kappa^{(l)} (Figure 1). We randomly sample inputs with shape (m,k,k,3)(m,k,k,3) where m∈{12,20}m\in\{12,20\} and k=6k=6. We compute the exact NTK with activation function Erf using the Neural Tangents library (Novak et al., 2019a). We see excellent agreement between the theoretical calculation of κ(l)\kappa^{(l)} in Sec. 6 (summarized in Table 1) and the experimental results Figure 1.

Maximum Learning Rates (Figure 2 (c)). In practice, given a set of hyper-parameters of a network, knowing the range of feasible learning rates is extremely valuable. As discussed above, in the infinite width setting, Equation 5 implies the maximal convergent learning rate is given by ηtheory≡2/λmax(l)\eta_{\rm theory}\equiv 2/{\lambda^{(l)}_{\rm max}}. From our theoretical results above, varying the hyperparameters of our network allows us to vary λmax(l)\lambda^{(l)}_{\rm max} over a wide range and test this hypothesis. This is shown for depth 10 networks varying σw2\sigma_{w}^{2} with η=ρηtheory\eta=\rho\eta_{\rm theory}. We see that networks become untrainable when ρ\rho exceeds 2 as predicted.

Trainability vs Generalization (Figure 2 (a,b)). We conduct an experiment training finite-width CNN-F networks with 1k training samples from CIFAR-10 with 20×2020\times 20 different (σω2,l)(\sigma_{\omega}^{2},l) configurations. We train each network using SGD with batch size b=256b=256 and learning rate η=0.1ηtheory\eta=0.1\eta_{\rm theory}. We see in Figure 2 (a) that deep in the chaotic phase we see that all configurations reach perfect training accuracy, but the network completely fails to generalize in the sense test accuracy is around 10%10\%. As expected, in the ordered phase we see that although the training accuracy degrades generalization improves. As expected we see that the depth-scales ξ1\xi_{1} and ξ∗\xi_{*} control trainability in the ordered phase and generalization in the chaotic phase respectively. We also conduct extra experiments for FCN with more training points (16k); see Figure 6.

CNN-P v.s. CNN-F: spatial correction (Figure 2 (d-f)). We compute the test accuracy using the analytic equations for gradient flow, Equation 6, which corresponds to the test accuracy of ensemble of gradient descent trained neural networks taking the width to infinity. As above, we use 1k1k training points and consider a 20×2020\times 20 grid of configurations for (σω2,l)(\sigma_{\omega}^{2},l). We plot the test performance of CNN-P and CNN-F and the performance difference in Fig 2 (d-f). As expected, we see that the performance of both CNN-P and CNN-F are captured by ξ1=−1/log⁡(χ1)\xi_{1}=-1/\log(\chi_{1}) in the ordered phase and by ξ∗=−1/(log⁡ξc−log⁡ξ1)\xi_{*}=-1/(\log\xi_{c}-\log\xi_{1}) in the chaotic phase. We see that the test performance difference between CNN-P and CNN-F exhibits a region in the ordered phase (a blue strip) where CNN-F outperforms CNN-P by a large margin. This performance difference is due to the correction term dd as predicted by the P(Θ(l))P(\Theta^{(l)})-row of Table 1. We also conduct extra experiments densely varying σb2\sigma_{b}^{2}; see Sec. G.4. Together these results provide an extremely stringent test of our theory.

Conclusion and Future Work

In this work, we identify several quantities (λmax\lambda_{\rm max}, λbulk\lambda_{\rm bulk}, κ\kappa, and P(Θ(l))P(\Theta^{(l)})) related to the spectrum of the NTK that control trainability and generalization of deep networks. We offer a precise characterization of these quantities and provide substantial experimental evidence supporting their role in predicting the training and generalization performance of deep neural networks. Future work might extend our framework to other architectures (for example, residual networks with batch-norm or attention architectures). Understanding the role of the nonuniform Fourier modes in the NTK in determining the test performance of CNNs is also an important research direction.

In practice, the correspondence between the NTK and neural networks is often broken due to, e.g., insufficient width, using a large learning rate, or changing the parameterization. Our theory does not directly apply to this setting. As such, developing an understanding of training and generalization away from the NTK regime remains an important research direction.

Acknowledgements

We thank Jascha Sohl-dickstein, Greg Yang, Ben Adlam, Jaehoon Lee, Roman Novak and Yasaman Bahri for useful discussions and feedback. We also thank anonymous reviewers for feedback that helped improve the manuscript.

References

Appendix A Related Work

Recent work Jacot et al. (2018); Du et al. (2018b); Allen-Zhu et al. (2018); Du et al. (2018a); Zou et al. (2018) proved global convergence of over-parameterized deep networks by showing that the NTK essentailly remains a constant over the course of training. However, in a different scaling limit the NTK changes over the course of training and global convergence is much more difficult to obtain and is known for neural networks with one hidden layer Mei et al. (2018); Chizat & Bach (2018); Sirignano & Spiliopoulos (2018); Rotskoff & Vanden-Eijnden (2018). Therefore, understanding the training and generalization properties in this scaling limit remains a very challenging open question.

Two excellent concurrent works (Hayou et al., 2019; Jacot et al., 2019) also study the dynamics of Θ(l)(x,x′)\Theta^{(l)}(x,x^{\prime}) for FCNs (and deconvolutions in (Jacot et al., 2019)) as a function of depth and variances of the weights and biases. (Hayou et al., 2019) investigates role of activation functions (smooth v.s. non-smooth) and skip-connection. (Jacot et al., 2019) demonstrate that batch normalization helps remove the “ordered phase” (as in (Yang et al., 2019)) and a layer-dependent learning rate allows every layer in a network to contribute to learning. As opposed to these contributions, here we focus our effort on understanding trainability and generalization in this context. We also provide a theory for a wider range of architectures than these other efforts.

Appendix B Signal propagation of NNGP and NTK

In this section, we assume that the activation function ϕ\phi has a continuous third derivative. Recall that the recursive formulas for NNGP K(l)\mathcal{K}^{(l)} and the NTK Θ(l)\Theta^{(l)} are given by

Note that we have normalized each input to have variance q∗q^{*} and the diagonals of K(l)\mathcal{K}^{(l)} are equal to q∗q^{*} for all \l\l. The off-diagonal terms of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} are denoted by qab(l)q^{(l)}_{ab} and pab(l)p^{(l)}_{ab}, resp. and the diagonal terms are q(l)q^{(l)} and p(l)p^{(l)}, resp. The above equations can be simplified to

In what follows, we compute the evolution of qab(l)q^{(l)}_{ab}, pab(l)p^{(l)}_{ab}, p(l)p^{(l)} and the spectrum and condition numbers of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)}. We will use λmax(Θ(l))/λmax(K(l))\lambda_{\rm max}(\Theta^{(l)})/\lambda_{\rm max}(\mathcal{K}^{(l)}), λbulk(Θ(l))/λbulk(K(l))\lambda_{\rm bulk}(\Theta^{(l)})/\lambda_{\rm bulk}(\mathcal{K}^{(l)}) and κ(Θ(l))/κ(K(l))\kappa(\Theta^{(l)})/\kappa(\mathcal{K}^{(l)}) to denote the maximum eigenvalues, the bulk eigenvalues and the condition number of Θ(l)/K(l)\Theta^{(l)}/\mathcal{K}^{(l)}, resp.

The diagonal terms are relatively simple to compute. Equation 24 gives

In the chaotic phase, χ1>1\chi_{1}>1 and p(l)≈χ1l−1q∗p^{(l)}\approx\chi_{1}^{l-1}q^{*}, i.e. diverges exponentially quickly.

Now we compute the off-diagonal terms. Since χc∗=σω2T˙(qab∗)<1\chi_{c^{*}}=\sigma_{\omega}^{2}\dot{\mathcal{T}}(q^{*}_{ab})<1 in the chaotic, pab∗p_{ab}^{*} exists and is finite. Indeed, letting l→∞l\to\infty in equation 23, we have

To compute the finite depth correction, let

Applying Taylor’s expansion to the first equation of 23 gives

Thus qab(l)q^{(l)}_{ab} converges to qab∗q^{*}_{ab} exponentially quickly with

Similarly, applying Taylor’s expansion to the second equation of 23 gives

where χc∗,2=σω2T¨(qab∗)\chi_{c^{*},2}=\sigma_{\omega}^{2}\ddot{\mathcal{T}}(q^{*}_{ab}). This implies

Note that δab(l)\delta_{ab}^{(l)} contains a polynomial correction term and decays like lχc∗ll\chi_{c^{*}}^{l}.

There exist a finite number ζab\zeta_{ab} such that

We want to emphasize that the limits are data-dependent, which was verified in Fig. 1(e) and 1(f) empirically.

Let ζab(l)=χc∗−lϵab(l)\zeta^{(l)}_{ab}=\chi_{c^{*}}^{-l}\epsilon_{ab}^{(l)}. We will show ζab(l)\zeta^{(l)}_{ab} is a Cauchy sequence. For any k>lk>l

Thus ζab≡lim⁡l→∞χc∗−lϵab(l)\zeta_{ab}\equiv\lim_{l\to\infty}\chi_{c^{*}}^{-l}\epsilon_{ab}^{(l)} exists and

Let ηab(l)=χc∗−lδab(l)−l(1+χc∗,2χc∗pab∗)ζab\eta^{(l)}_{ab}=\chi_{c^{*}}^{-l}\delta_{ab}^{(l)}-l(1+\frac{\chi_{c^{*},2}}{\chi_{c^{*}}}p_{ab}^{*})\zeta_{ab}. Coupled the above equation with Equation 41, we have

B.1.2 The Spectrum of the NNGP and NTK

We consider the spectrum of K\mathcal{K} and Θ\Theta in this phase. For K(l)\mathcal{K}^{(l)}, we have qab∗=c∗q∗q^{*}_{ab}=c^{*}q^{*} (with c∗<1c^{*}<1), q(l)=q∗q^{(l)}=q^{*} and qab(l)=qab∗+O(χc∗l)q^{(l)}_{ab}=q^{*}_{ab}+\mathcal{O}(\chi_{c^{*}}^{l}). Thus

The NNGP K∗\mathcal{K}^{*} has two different eigenvalues: q∗(1+(m−1)c∗)q^{*}(1+(m-1)c^{*}) of order 1 and q∗(1−c∗)q^{*}(1-c^{*}) of order (m−1)(m-1), where mm is the size of the dataset. For large ll, since the spectral norm of El\mathcal{E}^{l} is O(χc∗l)\mathcal{O}(\chi_{c^{*}}^{l}), the spectrum and condition number of K(l)\mathcal{K}^{(l)} are

For Θ(l)\Theta^{(l)}, we have pab(l)=pab∗+O(lχc∗l)→pab∗<∞p^{(l)}_{ab}=p_{ab}^{*}+\mathcal{O}(l\chi_{c^{*}}^{l})\to p_{ab}^{*}<\infty and p(l)=1−χ1l1−χ1q∗→∞p^{(l)}=\frac{1-\chi_{1}^{l}}{1-\chi_{1}}q^{*}\to\infty, i.e.

Thus Θ(l)\Theta^{(l)} is essentially a diverging constant multiplying the identity and

B.2 Ordered Phase

In the ordered phase, qab(l)→q∗q^{(l)}_{ab}\to q^{*}, q(l)=q∗q^{(l)}=q^{*}, p(l)→p∗p^{(l)}\to p^{*} and pab(l)→p∗p^{(l)}_{ab}\to p^{*}. Indeed, letting \l→∞\l\to\infty in the equations 24 and 23,

The correction of the diagonal terms are p(l)=p∗−χ1l−χ11−χ1q∗p^{(l)}=p^{*}-\frac{\chi_{1}^{l}-\chi_{1}}{1-\chi_{1}}q^{*}. Same calculation as in the chaotic phase implies

where χ1,2=σω2T¨(q∗)\chi_{1,2}=\sigma_{\omega}^{2}\ddot{\mathcal{T}}(q^{*}). Note that δab(l)\delta_{ab}^{(l)} contains also a polynomial correction term and decays like lχ1ll\chi_{1}^{l}.

Similar, in the ordered phase we have the following.

Since the proof is almost identical to Lemma 1, we omit the details.

B.2.2 The Spectrum of the NNGP and NTK

For K(l)\mathcal{K}^{(l)}, we have qab∗=q∗q^{*}_{ab}=q^{*}, qab(l)=q∗+O(χ1l)q^{(l)}_{ab}=q^{*}+\mathcal{O}(\chi_{1}^{l}) and q(l)=q∗q^{(l)}=q^{*}. Thus

For Θ(l)\Theta^{(l)}, pab(l)=p∗+O(lχ1l)p^{(l)}_{ab}=p^{*}+\mathcal{O}(l\chi_{1}^{l}) and p(l)=p∗−χ1l−χ11−χ1q∗=p∗+O(χ1l)p^{(l)}=p^{*}-\frac{\chi_{1}^{l}-\chi_{1}}{1-\chi_{1}}q^{*}=p^{*}+\mathcal{O}(\chi_{1}^{l}). Thus

B.3 The critical line.

We have χ1=1\chi_{1}=1 on the critical line. Equation 24 implies p(l)=lq∗p^{(l)}=lq^{*}, i.e. the diagonal terms diverge linearly. To capture the linear divergence of pab(l)p^{(l)}_{ab}, define

We need to expand the first equation of 23 to the second order

Here we assume T\mathcal{T} has a continuous third derivative (which is sufficient to assume the activation ϕ\phi to have a continuous third derivative.) The above equation implies

Plugging Equation 73 into the above equation gives

B.3.2 The Spectrum of NNGP and NTK

For K(l)\mathcal{K}^{(l)}, qab(l)=q∗+O(l−1)q^{(l)}_{ab}=q^{*}+\mathcal{O}(l^{-1}) and q(l)=q∗q^{(l)}=q^{*}. Thus

For Θ(l)\Theta^{(l)}, pab(l)=13q∗l+O(1)p^{(l)}_{ab}=\frac{1}{3}q^{*}l+\mathcal{O}(1) and p(l)=lq∗p^{(l)}=lq^{*}. Thus

Appendix C NNGP and NTK of Relu networks.

We only consider the critical initialization (i.e. He’s initialization (He et al., 2015)) σω2=2\sigma_{\omega}^{2}=2 and σb2=0\sigma_{b}^{2}=0, which preserves the norm of an input from layer to layer. We also normalize the inputs to have unit variance, i.e. q∗=q(l)=q(0)=1q^{*}=q^{(l)}=q^{(0)}=1. Recall that

which gives p(l)=lp^{(l)}=l. Using the equations in Appendix C of (Lee et al., 2019) gives

and taking the derivative w.r.t. ϵ\epsilon

This is enough to conclude (similar to the above calculation)

Recall that the diagonals of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} are q(l)=1q^{(l)}=1 and p(l)=lp^{(l)}=l, resp. Therefore the spectrum and the condition numbers of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} for large ll are

C.2 Residual Relu

We consider the following “continuum” residual network

Using the fact that q(0)=1q^{(0)}=1 (i.e. the inputs have unit variance), we can compute the diagonal terms q(t)=etq^{(t)}=e^{t} and p(t)=tetp^{(t)}=te^{t}. Letting qab(t)=etcab(t)q_{ab}^{(t)}=e^{t}c_{ab}^{(t)} and applying the above fractional Taylor expansion to T\mathcal{T} and T˙\dot{\mathcal{T}}, we have

Ignoring the higher order term and set y(t)=(1−cab(t))y(t)=(1-c_{ab}^{(t)}), we have

Solving this gives y(t)=9π22t−2y(t)=\frac{9\pi^{2}}{2}t^{-2} (note that y(∞)=0y(\infty)=0), which implies

Applying this estimate to Equation 97 gives

Thus the limiting condition number of the NTK is m/3+1m/3+1. This is the same as the above non-residual Relu case although the entries of K(t)\mathcal{K}^{(t)} and Θ(t)\Theta^{(t)} blow up exponentially with tt.

C.3 Residual Relu + Layer Norm

As we saw above, all the entries of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} of a residual Relu network blow up exponentially, so do its gradients. In what follows, we show that normalization could help to avoid this issue. We consider the following “continuum” residual network with “layer norm”

Using the fact that q(0)=1q^{(0)}=1 (i.e. the inputs have unit variance) and the mapping 2T2\mathcal{T} is norm preserving, we see that q(t)=1q^{(t)}=1 because

This implies p(t)=tp^{(t)}=t (note that p˙(t)=q(t)=1\dot{p}^{(t)}=q^{(t)}=1 and we assume the initial value p(0)=0p^{(0)}=0.) The off-diagonal terms can be computed similarly and

Thus the condition number of the NTK is m/3+1m/3+1. This is the same as the non-residual Relu case discussed above.

To keep the notation simple, we denote Xd=XtrainX_{d}=X_{\text{train}}, Yd=YtrainY_{d}=Y_{\text{train}}, Θtd=Θtest, train\Theta_{td}=\Theta_{\text{test, train}}, Θdd=Θtrain, train\Theta_{dd}=\Theta_{\text{train, train}}. Recall that

We split our calculation into three parts.

In this case the diagonal p(l)p^{(l)} diverges exponentially and the off-diagonals pab(l)p^{(l)}_{ab} converges to a bounded constant pab∗p_{ab}^{*}. We further assume the input labels are centered in the sense YdY_{d} contains the same number of positive (+1) and negative (-1) labelsWhen the number of classes is greater than two, we require YdY_{d} to have mean zero along the batch dimension for each class.. We expand Θ(l)\Theta^{(l)} about its “fixed point”

In the last equation, we have used the fact 11TYd=0\bm{1}\bm{1}^{T}Y_{d}={\bf 0} and Θtd∗Yd=0\Theta_{td}^{*}Y_{d}={\bf 0} since YdY_{d} is balanced. Therefore

Without centering the labels YdY_{d} and normalizing each input in XdX_{d} to have the same variance, we will get a χ1l\chi_{1}^{l} decay for P(Θ(l))YdP(\Theta^{(l)})Y_{d} instead of l(χc∗/χ1)ll(\chi_{c^{*}}/\chi_{1})^{l}.

D.2 Critical line

Note that in this phase, both the diagonals and the off-diagonals diverge linearly. In this case

Here we use 1d\bm{1}_{d} to denote the all ‘1’ (column) vector with length equal to the number of training points in XdX_{d} and 1t\bm{1}_{t} is defined similarly. Note that the constant matrix BB is invertible. By Equation 77

The term 1t1dTB−1\bm{1}_{t}\bm{1}_{d}^{T}B^{-1} is independent of the inputs and 1t1dTB−1Yd=0\bm{1}_{t}\bm{1}_{d}^{T}B^{-1}Y_{d}=0 when YdY_{d} is centered. Thus

D.3 Ordered Phase

In the ordered phase, we have that Θdd(l)=p∗1d1dT+lχ1lAdd(l)\Theta^{(l)}_{dd}=p^{*}\bm{1}_{d}\bm{1}^{T}_{d}+l\chi_{1}^{l}\bm{A}^{(l)}_{dd} where Add(l)\bm{A}^{(l)}_{dd}, a symmetric matrix, represents the data-dependent piece of Θdd(l)\Theta^{(l)}_{dd}. By Lemma 2, Add(l)→Add\bm{A}^{(l)}_{dd}\to\bm{A}_{dd} as l→∞l\to\infty. To simply the notation, in the calculation below we will replace Add(l)\bm{A}^{(l)}_{dd} by Add\bm{A}_{dd}. We also assume Add\bm{A}_{dd} is invertible. To compute the mean predictor, P(Θ(l))P(\Theta^{(l)}), asymptotically we begin by computing (Θdd(l))−1(\Theta^{(l)}_{dd})^{-1} via the Woodbury identity,

and ai=1m∑jAij−1\bm{a}_{i}=\frac{1}{m}\sum_{j}\bm{A}_{ij}^{-1}. Noting that Θtd(l)=p∗1t1dT+lχ1lAtd\Theta^{(l)}_{td}=p^{*}\bm{1}_{t}\bm{1}_{d}^{T}+l\chi_{1}^{l}\bm{A}_{td} we can compute the mean predictor,

Note that there is no divergence in P(Θ(l))P(\Theta^{(l)}) as l→∞l\to\infty and the limit is well-defined. The term p^1taT\hat{p}\bm{1}_{t}\bm{a}^{T} is independent from the input data.

We therefore see that even in the infinite-depth limit the mean predictor retains its data-dependence and we expect these networks to be able generalize indefinitely.

Appendix E Dropout

In this section, we investigate the effect of adding a dropout layer to the penultimate layer. Let 0<ρ≤10<\rho\leq 1 and γj(L)(x)\gamma^{(L)}_{j}(x) be iid random variables

where Wij(l)W^{(l)}_{ij} and bi(l)b^{(l)}_{i} are iid Gaussians N(0,1)\mathcal{N}(0,1). Since no dropout is applied in the first LL layers, the NNGP kernel K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} can be computed using Equation 20 and Equation 8. Let Kρ(L+1)\mathcal{K}_{\rho}^{(L+1)} and Θρ(L+1)\Theta_{\rho}^{(L+1)} denote the NNGP and NTK of the (L+1)(L+1)-th layer. Note that when ρ=1\rho=1, K1(L+1)=K(L+1)\mathcal{K}_{1}^{(L+1)}=\mathcal{K}^{(L+1)} and Θ1(L+1)=Θ(L+1)\Theta_{1}^{(L+1)}=\Theta^{(L+1)} . We will compute the correction induced by ρ<1\rho<1. The fact

implies that the NNGP kernel Kρ(L+1)\mathcal{K}^{(L+1)}_{\rho} (Schoenholz et al., 2017) is

Now we compute the NTK Θρ(L+1)\Theta^{(L+1)}_{\rho}, which is a sum of two terms

Here θ(L+1)\theta^{(L+1)} denote the parameters in the (L+1)(L+1) layer, namely, Wij(L+1)W_{ij}^{(L+1)} and bi(L+1)b^{(L+1)}_{i} and θ(≤L)\theta^{(\leq L)} the remaining parameters. Note that the first term in Equation 138 is equal to Kρ(L+1)(x,x′)\mathcal{K}^{(L+1)}_{\rho}(x,x^{\prime}). Using the chain rule, the second term is equal to

In sum, we see that dropout only modifies the diagonal terms

In Fig 4, we plot the evolution of κρ(L)\kappa^{(L)}_{\rho} for ρ=0.8,0.95,0.99\rho=0.8,0.95,0.99 and 11, confirming Equation 145.

Appendix F Convolutions

In this section, we compute the evolution of Θ(l)\Theta^{(l)} for CNNs.

General setup. For simplicity of presentation we consider 1D convolutional networks with circular padding as in Xiao et al. (2018). We will see that this reduces to the fully-connected case introduced above if the image size is set to one and as such we will see that many of the same concepts and equations carry over schematically from the fully-connected case. The theory of two-or higher-dimensional convolutions proceeds identically but with more indices.

Random weights and biases. The parameters of the network are the convolutional filters and biases, ωij,β(l)\omega^{(l)}_{ij,\beta} and μi(l)\mu^{(l)}_{i}, respectively, with outgoing (incoming) channel index ii (jj) and filter relative spatial location β∈[±k]≡{−k,…,0,…,k}\beta\in[\pm k]\equiv\{-k,\dots,0,\dots,k\}.We will use Roman letters to index channels and Greek letters for spatial location. We use letters i,j,i′,j′i,j,i^{\prime},j^{\prime}, etc to denote channel indices, α,α′\alpha,\alpha^{\prime}, etc to denote spatial indices and β,β′\beta,\beta^{\prime}, etc for filter indices. As above, we will assume a Gaussian prior on both the filter weights and biases,

As above, σω2\sigma^{2}_{\omega} and σb2\sigma^{2}_{b} are hyperparameters that control the variance of the weights and biases respectively. N(l)N^{(l)} is the number of channels (filters) in layer ll, 2k+12k+1 is the filter size.

Here K(l)≡[Kα,α′(l)(x,x′)]α,α′∈[d],x,x′∈X\mathcal{K}^{(l)}\equiv[\mathcal{K}^{(l)}_{\alpha,\alpha^{\prime}}(x,x^{\prime})]_{\alpha,\alpha^{\prime}\in[d],x,x^{\prime}\in\mathcal{X}}, T\mathcal{T} is a non-linear transformation related to its fully-connected counterpart, and A\mathcal{A} a convolution acting on Xd×Xd\mathcal{X}d\times\mathcal{X}d PSD matrices

To understand how the neural tangent kernel evolves with depth, we define the NTK of the ll-th hidden layer to be Θ^(l){\hat{\Theta}}^{(l)}

where θ≤l\theta^{\leq l} denotes all of the parameters in layers at-or-below the ll’th layer. It does not matter which channel index ii is used because as the number of channels approach infinity, this kernel will also converge in distribution to a deterministic kernel Θ(l+1)\Theta^{(l+1)} (Yang, 2019), which can also be computed recursively in a similar manner to the NTK for fully-connected networks as (Yang, 2019; Arora et al., 2019),

where T˙\dot{\mathcal{T}} is given by Equation 151 with ϕ\phi replaced by its derivative ϕ′\phi^{\prime}. We will also normalize the variance of the inputs to q∗q^{*} and hence treat T\mathcal{T} and T˙\dot{\mathcal{T}} as pointwise functions. We will only present the treatment in the chaotic phase to showcase how to deal with the operator A\mathcal{A}. The treatment of other phases are similar. Note that the diagonal entries of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} are exactly the same as the fully-connected setting, which are q∗q^{*} and p(l)=lq∗p^{(l)}=lq^{*}, respectively. We only need to consider the off-diagonal terms. Letting l→∞l\to\infty in Equation 154 we see that all the off-diagonal terms also converge pab∗p_{ab}^{*}. Note that A\mathcal{A} does not mix terms from different diagonals and it suffices to handle each off-diagonal separately. Let ϵab(l)\epsilon_{ab}^{(l)} and δab(l)\delta_{ab}^{(l)} denote the correction of the jj-th diagonal of K(l)\mathcal{K}^{(l)} and Θ(l)\Theta^{(l)} to the fixed points. Linearizing Equation 150 and Equation 154 gives

Next let {ρα}α\{\rho_{\alpha}\}_{\alpha} be the eigenvalues of A\mathcal{A} and ϵab,α(l)\epsilon_{ab,\alpha}^{(l)} and δab,α(l)\delta_{ab,\alpha}^{(l)} be the projection of ϵab(l)\epsilon_{ab}^{(l)} and δab(l)\delta_{ab}^{(l)} onto the α\alpha-th eigenvector of A\mathcal{A}, respectively. Then for each α\alpha,

Therefore, the correction Θ(l)−Θ∗\Theta^{(l)}-\Theta^{*} propagates independently through different Fourier modes. In each mode, up to a scaling factor ραl\rho_{\alpha}^{l}, the correction is the same as the correction of FCN. Since the subdominant modes (with ∣ρα∣<1|\rho_{\alpha}|<1) decay exponentially faster than the dominant mode (with ρα=1\rho_{\alpha}=1), for large depth, the NTK of CNN is essentially the same as that of FCN.

F.2 The effect of pooling and flattening of CNNs

where Θflatten(l)\Theta^{(l)}_{\rm flatten} (Θpool(l)\Theta^{(l)}_{\rm pool}) denotes the NTK right after flattening (pooling) the last convolution. We will also use Θfc(l)\Theta^{(l)}_{\rm fc} to denote the NTK of FC. Kflatten(l)\mathcal{K}^{(l)}_{\rm flatten}, Kpool(l)\mathcal{K}^{(l)}_{\rm pool} and Kfc(l)\mathcal{K}^{(l)}_{\text{f}c} are defined similarly. As discussed above, in the large depth setting, all the diagonals Θα,α(l)(x,x)=p(l)\Theta^{(l)}_{\alpha,\alpha}(x,x)=p^{(l)} (since the inputs are normalized to have variance q∗q^{*} for each pixel) and similar to Θfc(l)\Theta^{(l)}_{\rm fc}, all the off-diagonals Θα′,α(l)(x,x′)\Theta^{(l)}_{\alpha^{\prime},\alpha}(x,x^{\prime}) are almost equal (in the sense they have the same order of correction to pab∗p_{ab}^{*} if exists.) Without loss of generality, we assume all off-diagonals are the same and equal to pab(l)p^{(l)}_{ab} (the leading correction of qab(l)q^{(l)}_{ab} for CNN and FCN are of the same order.) Applying flattening and pooling, the NTKs become

respectively. As we can see, Θflatten(l)\Theta^{(l)}_{\rm flatten} is essentially the same as its FCN counterpart Θfc(l)\Theta^{(l)}_{\rm fc} up to sub-dominant Fourier modes which decay exponentially faster than the dominant Fourier modes. Therefore the spectrum properties of Θflatten(l)\Theta^{(l)}_{\rm flatten} and Θfc(l)\Theta^{(l)}_{\rm fc} are essentially the same for large ll; see Figure 1 (a - c).

However, pooling alters the NTK/NNGP spectrum in an interesting way. Noticeably, the contribution from p(l)p^{(l)} is discounted by a factor of dd. On the critical line, asymptotically, the on- and off-diagonal terms are

Here we use blue color to indicate the changes of such quantities against their Θflatten(l)\Theta^{(l)}_{\rm flatten} counterpart. Alternatively, one can consider Θflatten(l)\Theta^{(l)}_{\rm flatten} as a special version (with {\color[rgb]{0,0,1}{d}}=1) of Θpool(l)\Theta^{(l)}_{\rm pool}. Thus pooling decreases λbulk(l)\lambda_{\rm bulk}^{(l)} roughly by a factor of {\color[rgb]{0,0,1}{d}} and increases the condition number by a factor of {\color[rgb]{0,0,1}{d}} comparing to flattening. In the chaotic phase, pooling does not change the off-diagonals qab(l)=O(1)q^{(l)}_{ab}=\mathcal{O}(1) but does slow down the growth of the diagonals by a factor of dd, i.e. p^{(l)}=\mathcal{O}(\chi_{1}^{l}\color[rgb]{0,0,1}{/d}). This improves P(Θ(l))P(\Theta^{(l)}) by a factor of {\color[rgb]{0,0,1}{d}}. This suggests, in the chaotic phase, there exists a transient regime of depths, where CNN-F hardly perform while CNN-P performs well. In the ordered phase, the pooling does not affect λmax(l)\lambda_{\rm max}^{(l)} much but does decrease λbulk(l)\lambda_{\rm bulk}^{(l)} by a factor of {\color[rgb]{0,0,1}{d}} and the condition number κ(l)\kappa^{(l)} grows approximately like {\color[rgb]{0,0,1}{d}}l\chi_{1}^{-l}, {\color[rgb]{0,0,1}{d}} times bigger than its flattening and fully-connected network counterparts. This suggests the existence of a transient regime of depths, in which CNN-F outperforms CNN-P. This might be surprising because it is commonly believed CNN-P usually outperforms CNN-F. These statements are supported empirically in Figure 2.

Appendix G Figure Zoo

We plot the phase diagrams for the Erf function and the tanh⁡\tanh function (adopted from (Pennington et al., 2018)).

G.2 SGD on FCN on Larger Dataset: Figure 6.

We report the training and test accuracy of FCN trained on a subset (16k training points) of CIFAR-10 using SGD with 20 ×\times 20 different (σω2,l)(\sigma_{\omega}^{2},l) configurations.

G.3 NNGP vs NTK prediction: Figure 7.

Here we compare the test performance of the NNGP and NTK with different (σω2,l)(\sigma_{\omega}^{2},l) configurations. In the chaotic phase, the generalizable depth-scale of the NNGP is captured by ξc∗=−1/log⁡(χc∗)\xi_{c^{*}}=-1/\log(\chi_{c^{*}}). In contrast, the generalizble depth-scale of the NTK is captured by ξ∗=−1/(log⁡(χc∗)−log⁡(χ1))\xi_{*}=-1/(\log(\chi_{c^{*}})-\log(\chi_{1})). Since χ1>1\chi_{1}>1 in the chaotic phase, ξc∗>ξ∗\xi_{c^{*}}>\xi_{*}. Thus for larger depth, the NNGP kernel performs better than the NTK. Corrections due to an additional average pooling layer is plotted in the third column of Figure .7

We demonstrate that our prediction for the generalizable depth-scales for the NTK (ξ∗\xi_{*}) and NNGP (ξc\xi_{c}) are robust across a variety of hyperparameters. We densely sweep over 9 different values of σb2∈[0.2,1.8]\sigma_{b}^{2}\in[0.2,1.8]. For each σb2\sigma_{b}^{2} we compute the NTK/NNGP test accuracy for 20 * 50 different configurations of (l, σω2\sigma_{\omega}^{2}) with l∈l\in and σω2∈[0.12,4.92]\sigma_{\omega}^{2}\in[0.1^{2},4.9^{2}]. The training set is a 88k subset of CIFAR-10.

G.5 Densely Sweeping Over the Regularization Strength σ𝜎\sigma: Figure 9

Similar to the above setup, we fixed σb2=1.6\sigma_{b}^{2}=1.6 and densely vary σ∈{0,10−6,…,100}\sigma\in\{0,10^{-6},\dots,10^{0}\}.