Finite Depth and Width Corrections to the Neural Tangent Kernel
Boris Hanin, Mihai Nica
Introduction
Modern neural networks typically overparameterized: they have many more parameters than the size of the datasets on which they are trained. That some setting of parameters in such networks can interpolate the data is therefore not surprising. But it is a priori unexpected that not only can such interpolating parameter values can be found by stochastic gradient descent (SGD) on the highly non-convex empirical risk but that the resulting network function not only interpolates but also extrapolates to unseen data. In an overparameterized neural network individual parameters can be difficult to interpret, and one way to understand training is to rewrite the SGD updates
of trainable parameters with a loss and learning rate as kernel gradient descent updates for the values of the function computed by the network:
Relation (1) is valid to first order in It translates between two ways of thinking about the difficulty of neural network optimization:
The function space view where the loss , which is a simple function of the network mapping , is minimized over the manifold of all functions representable by the architecture of using gradient descent with respect to a potentially complicated Riemannian metric on
Moreover, the joint statistical effects of depth and width on in finite size networks remain unclear, and the purpose of this article is to shed light on the simultaneous effects of depth and width on for finite but large widths and any depth . Our results apply to fully connected ReLU networks at initialization for which we will show:
In contrast to the regime in which the depth is fixed but the width is large, is not approximately deterministic at initialization so long as is bounded away from . Specifically, for a fixed input the normalized on-diagonal second moment of satisfies
Thus, when is bounded away from , even when both are large, the standard deviation of is at least as large as its mean, showing that its distribution at initialization is not close to a delta function. See Theorem 1.
Moreover, when is the square loss, the average of the SGD update to from a batch of size one containing satisfies
where is the input dimension. Therefore, if the NTK will have the potential to evolve in a data-dependent way. Moreover, if is comparable to and then it is possible that this evolution will have a well-defined expansion in See Theorem 2.
In both statements above, means is bounded above and below by universal constants. We emphasize that our results hold at finite and the implicit constants in both and in the error terms are independent of Moreover, our precise results, stated in §2 below, hold for networks with variable layer widths. We have denoted network width by only for the sake of exposition. The appropriate generalization of to networks with varying layer widths is the parameter
which in light of the estimates in (1) and (2) plays the role of an inverse temperature.
A number of articles have followed up on the original NTK work . Related in spirit to our results is the article , which uses Feynman diagrams to study finite width corrections to general correlations functions (and in particular the NTK). The most complete results obtained in are for deep linear networks but a number of estimates hold general non-linear networks as well. The results there, like in essentially all previous work, fix the depth and let the layer widths tend to infinity. The results here and in , however, do not treat as a constant, suggesting that the expansions (e.g. in ) can be promoted to expansions. Also, the sum-over-path approach to studying correlation functions in randomly initialized ReLU nets was previously taken up for the foward pass in and for the backward pass in and .
2. Implications and Future Work
Taken together (1) and (2) above (as well as Theorems 1 and 2) show that in fully connected ReLU nets that are both deep and wide the neural tangent kernel is genuinely stochastic and enjoys a non-trivial evolution during training. This suggests that in the overparameterized limit with , the kernel may learn data-dependent features. Moreover, our results show that the fluctuations of both and its time derivative are exponential in the inverse temperature
It would be interesting to obtain an exact description of its statistics at initialization and to describe the law of its trajectory during training. Assuming this trajectory turns out to be data-dependent, our results suggest that the double descent curve that trades off complexity vs. generalization error may display significantly different behaviors depending on the mode of network overparameterization.
However, it is also important to point out that the results in show that, at least for fully connected ReLU nets, gradient-based training is not numerically stable unless is relatively small (but not necessarily zero). Thus, we conjecture that there may exist a “weak feature learning” NTK regime in which network depth and width are both large but . In such a regime, the network will be stable enough to train but flexible enough to learn data-dependent features. In the language of one might say this regime displays weak lazy training in which the model can still be described by a stochastic positive definite kernel whose fluctuations can interact with data.
Finally, it is an interesting question to what extent our results hold for non-linearities other than ReLU and for network architectures other than fully connected (e.g. convolutional and residual). Typical ConvNets, for instance, are significantly wider than they are deep, and we leave it to future work to adapt the techniques from the present article to these more general settings.
Formal Statement of Results
The three assumptions in (4) hold for vitually all standard network initialization schemes. The on-diagonal NTK is
We emphasize that although we have initialized the biases to zero, they are not removed them from the list of trainable parameters. Our first result is the following:
times a multiplicative error , where means is bounded above and below by universal constants times In particular, if all the hidden layer widths are equal (i.e. , for ), we have
This result shows that in the deep and wide double scaling limit
the NTK does not converge to a constant in probability. This is contrast to the wide and shallow regime and is fixed.
times a multiplicative error of size , where as in Theorem 1, In particular, if all the hidden layer widths are equal (i.e. , for ), we find
The remainder of this article is structured as follows. First, in §3 we introduce some notation about paths and edges in the computation graph of . This notation will be used in the proofs of Theorems 1 and 2, which are outlined in §4 and particularly in §4.1 where give an in-depth but informal explanation of our strategy for computing moments of and its time derivative. Then, §5-§7 give the detailed argument. The computations in §5 explain how to handle the contribution to and coming only from the weights of the network. They are the most technical and we give them in full detail. Then, the discussion in §6 and §7 show how to adapt the method developed in §5 to treat the contribution of biases and mixed bias-weight terms in and . Since the arguments are simpler in these cases, we omit some details and focus only on highlighting the salient differences.
Notation
If each edge in the computational graph of is assigned a weight , then associated to a path is a collection of weights:
Next, for an edge in the computational graph of we will write
for the layer of In the course of proving Theorems 1 and 2, it will be useful to associate to every an unordered multi-set of edges
to be the unordered multiset of edges in the complete directed bi-paritite graph oriented from to For every define its left and right endpoints to be
where are unordered multi-sets.
for the set of all possible edge multisets realized by paths in On a number of occasions, we will also write
We will moreover say that for a path an edge in the computational graph of belongs to (written ) if
Finally, for an edge in the computational graph of , we set
for the normalized and unnormalized weights on the edge corresponding to (see (3)).
Overview of Proof of Theorems 1 and 2
and have suppressed the dependence on Similarly, we have
and have used that the loss on the batch is given by for some target value To prove Theorem 1 we must estimate the following quantities:
To prove Theorem 2, we must control in addition
We prove Proposition 3 in §5 below. The proof already contains all the ideas necessary to treat the remaining moments. In §6 and §7 we explain how to modify the proof of Proposition 3 to prove the following two Propositions:
up to a multiplicative error of When is small, this expression is bounded above and below by a constant times
Thus, since Propositions 3 and 4 also give
where the weight of a path was defined in (10) and includes both the product of the weights along and the condition that every neuron in is open at . The path begins at some neuron in the input layer of and passes through a neuron in every subsequent layer until ending up at the unique neuron in the output layer (see (7)). Being a product over edge weights in a given path, the derivative of with respect to a weight on an edge of the computational graph of is:
There is a subtle point here that also involves indicator functions of the events that neurons along are open at However, with probability , the derivative with respect to of these indicator functions is identically at The details are in Lemma 11.
Obtain an exact formula for the expectation in (18):
Observe that the dependence of on is only up to a multiplicative constant:
The precise relation is (24). This shows that, up to universal constants,
This is captured precisely by the terms defined in (27),(28).
Notice that depends only on the un-ordered multiset of edges determined by (see (14)). We therefore change variables in the sum from the previous step to find
where a loop in occurs when the four paths interact. More precisely, a loop occurs whenever all four paths pass through the same neuron in some layer (see Figures 1 and 2).
Finally, we use Proposition 10 to obtain for this expectation estimates above and below that match up multiplicative constants.
Proof of Proposition 3
We begin with the well-known formula for the output of a ReLU net with biases set to and a linear final layer with one neuron:
The weight of a path was defined in (10) and includes both the product of the weights along and the condition that every neuron in is open at . As explained in §3, the inner sum in (19) is over paths in the computational graph of that start at neuron in the input layer and end at the output neuron and the random variables are the normalized weights on the edge of between layer and layer (see (9)). Differentiating this formula gives sum-over-path expressions for the derivatives of with respect to both and its trainable parameters. For the NTK and its first SGD update, the result is the following:
where the sum is over collections of two paths in the computation graph of and edges that lie on both paths. Similarly, almost surely,
Lemma 7 is proved in §5.2. The expression (20) is simple to evaluate due to the delta function in We obtain:
where in the second-to-last equality we used that the number of paths in the comutational graph of from a given neuron in the input to the output neuron equals and in the last equality we used that This proves the first equality in Theorem 1.
It therefore remains to evaluate (21) and (22). Since they are so similar, we will continue to discuss them in parallel. To start, notice that the expression appearing in (21) and (22) satisfies
For the remainder of the proof we will write
The advantage of is that it does not depend on Observe that for every , we have that either , , or . Thus, by symmetry, the sum over in (21) and (22) takes only four distinct values, represented by the following possibilities:
keeping track of which paths begin at the same neuron in the input layer to Hence, since
To evaluate let us write
for the indicator function of the event that paths pass through the same edge between layers in the computational graph of . Observe that
To simplify and observe that depends only on only via the unordered edge multi-set (i.e. only which edges are covered matters; not their labelling)
defined in Definition 3. Hence, we find that for
The counts in and have a convenient representation in terms of
Informally, the event indicates the presence of a “collision” of the four paths in before the earlier of the layers , while gives a “collision” between layers ; see Section 4.1 for the intuition behind calling these collisions. We also write
Finally, for , we will define
That is, a loop is created at layer if the four edges in all begin at occupy the same vertex in layer but occupy two different vertices in layer We have the following Lemma.
Suppose for some For each
We prove Lemma 8 in §5.3 below. Assuming it for now, observe that
Observe that every unordered multi-set four edge multiset can be obtained by starting from some , considering its unordered edge multi-set and doubling all its edges. This map from to is surjective but not injective. The sizes of the fibers is computed by the following Lemma.
Fix . The number of so that is where as in (35),
Since the number of in with specified equals we find that so that for each we have
Here, is the expectation with respect to the probability measure on obtained by taking independent, each drawn from the products of the measure on and the uniform measure on
We are now in a position to complete the proof of Theorems 1 and 2. To do this, we will evaluate the expectations above to leading order in with the help of the following elementary result which is proven as Lemma 18 in .
Let be independent events with probabilities and be independent events with probabilities such that
Denote by the indicator that the event happens, , and by the indicator that happens, . Further, fix for every some as well as . Define
Then, if for every , we have:
where by convention In contrast, if for every , we have:
Since , we may also write
Putting this together with (42) and noting that
where in the last inequality we used that for Since we conclude
When combined with (23) this gives the lower bound in Proposition 3. The matching upper bound is obtained from (46) in the same way using the opposite inequality from Proposition 10.
This completes the proof of Proposition 3, modulo the proofs of Lemmas 6-9, which we supply below.
With probability either there exists so that or, for every we have
Lemma 11 shows that for our fixed , with probability the derivative of each in (19) vanishes. Hence, almost surely, for any edge in the computational graph of
This proves the formulas for To derive the result for we write
where the loss on a single batch containing only is We therefore find
Using (47) and again applying Lemma 11, we find that with probability
To complete the proof of Lemma 6 it therefore remains to check that this last term has mean To do this, recall that the output layer of is assumed to be linear and that the distribution of each weight is symmetric around (and hence has vanishing odd moments). Thus, the expectation over the weights in layer has either or weights in it and so vanishes.
2. Proof of Lemma 7
Lemma 7 is almost a corollary of of Theorem 3 in and Proposition 2 in . The difference is that, in , the biases in were assumed to have a non-degenerate distribution, whereas here we’ve set them to zero. The non-degeneracy assumption is not really necessary, so we repeat here the proof from with the necessary modifications.
If then for any configuration of weights since the network biases all vanish. Will therefore suppose that Let us first show (20). We have from Lemma 6 that
To compute the inner expectation, write for the sigma algebra generated by the weight in layers up to and including . Let us also define the events:
where we recall from (2) that are the post-activations in layer Supposing first that is not in layer , the expectation becomes
Thus, the expectation in (48) becomes times
Note that given the pre-activations of different neurons in layer are independent. Hence,
Recall that by assumption, the weight matrix in layer is equal in distribution to This replacement leaves the product unchanged but changes to . On the event (which occurs whenever ) we have that with probability since we assumed that the distribution of each weight has a density relative to Lebesgue measure. Hence, symmetrizing over , we find that
Similarly, if is in layer then we automatically find that and , giving an expectation of . Proceeding in this way yields
which is precisely (20). The proofs of (21) and (22) are similar. We have
As before let us first assume that edges are not in layer . Then,
Again symmetrizing with respect to and using that the pre-activation of different neurons are independent given the activations in the previous layer we find that, on the event ,
where is the event that and are not in layer . Proceeding in this way one layer at a time completes the proofs of (21) and (22).
3. Proof of Lemma 8
For each there exists unique so that
We will say that two layers belong to the same loop of if exists so that
We proceed layer by layer to count the number of satisfying and To do this, suppose we are given and we have . Then is some permutation of with Moreover, for there is a unique edge (with multiplicity ) in whose left endpoint is Therefore, determines when In contrast, suppose If then consists of a single edge with multiplicity which again determines . In short, determines for all belonging to the same loop of as Therefore, the initial condition determines for all and the conditions determine in the loops of containing the layers of
4. Proof of Lemma 9
The proof of Lemma 9 is essentially identical to the proof of Lemma 8. In fact it is slightly simpler since there are no distinguished edges to consider. We omit the details.
Proof of Proposition 4
to be the event that the pre-activations of the neurons are positive.
where are defined in §3. Further, almost surely,
The proof of this result is a small modification of the proof of Lemma 6 and hence is omitted. Taking expectations, we therefore obtain the following analog to Lemma 7.
where for we have
The proof is identical to the argument used in §5.2 to establish Lemma 7, so we omit the details. The relation (51) is easy to simplify:
Since the inner sum in (54) is independent of by symmetry, we find
where for the second estimate we applied Lemma 9 and have written . Thus, as in the derivation of (42), we find that
Putting this together with (54) completes the proof of Lemma 14. ∎
as claimed in the statement Proposition 4.
Proof of Proposition 5
Here, for a neuron and we’ve denoted by the set of four tuples of paths in the computational graph of where start from and start at neurons respectively. The analog of Lemmas 7 and 13 (with essentially the same proof), gives that the expectation in the previous line equals
which, up to a multiplicative constant equals
which is independent of Thus, we find
where if we recall that is the indicator function of the event that paths pass through the same edge in the computational graph of at layer (see (29)).
As in the proof of Proposition 3, observe that depends only on the unordered multiset of edges in . Thus, we find that
Applying Proposition 10 as in the end of the proof of Propositions 3 and 4 we conclude
plus a term that has mean Therefore, as in Lemma 7, we find
where is as in (29), the sum is over unordered edge multisets (see (14)), and we’ve set
As in Lemma 8, the counting term satisfies
This completes the proof of Proposition 5.