Polylogarithmic width suffices for gradient descent to achieve arbitrarily small test error with shallow ReLU networks
Ziwei Ji, Matus Telgarsky
Introduction
Despite the extensive empirical success of deep networks, their optimization and generalization properties are still not fully understood. Recently, the neural tangent kernel (NTK) has provided the following insight into the problem. In the infinite-width limit, the NTK converges to a limiting kernel which stays constant during training; on the other hand, when the width is large enough, the function learned by gradient descent follows the NTK (Jacot et al., 2018). This motivates the study of overparameterized networks trained by gradient descent, using properties of the NTK. In fact, parameters related to the NTK, such as the minimum eigenvalue of the limiting kernel, appear to affect optimization and generalization (Arora et al., 2019).
However, in addition to such NTK-dependent parameters, prior work also requires the width to depend polynomially on , or , where denotes the size of the training set, denotes the failure probability, and denotes the target error. These large widths far exceed what is used empirically, constituting a significant gap between theory and practice.
In this paper, we narrow this gap by showing that a two-layer ReLU network with hidden units trained by gradient descent achieves classification error on test data, meaning both optimization and generalization occur. Unlike prior work, the width is fully polylogarithmic in , , and ; the width will additionally depend on the separation margin of the limiting kernel, a quantity which is guaranteed positive (assuming no inputs are parallel), can distinguish between true labels and random labels, and can give a tight sample-complexity analysis in the infinite-width setting. The paper organization together with some details are described below.
gives a test error bound. Concretely, using the preceding gradient descent analysis, and standard Rademacher tools and exploiting how little the weights moved, we show that with samples and iterations, gradient descent finds a solution with test error (cf. Theorem 3.2 and Corollary 3.3). (As discussed in Remark 3.4, samples also suffice via a smoothness-based generalization bound, at the expense of large constant factors.)
considers stochastic gradient descent (SGD) with access to a standard stochastic online oracle. We prove that with width at least polylogarithmic and samples, SGD achieves an arbitrarily small test error (cf. Theorem 4.1).
discusses the separation margin, which is in general a positive number, but reflects the difficulty of the classification problem in the infinite-width limit. While this margin can degrade all the way down to for random labels, it can be much larger when there is a strong relationship between features and labels: for example, on the noisy 2-XOR data introduced in (Wei et al., 2018), we show that the margin is , and our SGD sample complexity is tight in the infinite-width case.
1 Related work
There has been a large literature studying gradient descent on overparameterized networks via the NTK. The most closely related work is (Nitanda and Suzuki, 2019), which shows that a two-layer network trained by gradient descent with the logistic loss can achieve a small test error, under the same assumption that the NTK with respect to the first layer can separate the data distribution. However, they analyze smooth activations, while we handle the ReLU. They require hidden units, data samples, and steps, while our result only needs polylogarithmic hidden units, data samples, and steps.
On deep networks, a variety of works have established low training error (Allen-Zhu et al., 2018b; Du et al., 2018a; Zou et al., 2018; Zou and Gu, 2019). Allen-Zhu et al. (2018c) show that SGD can minimize the regression loss for recurrent neural networks, and Allen-Zhu and Li (2019b) further prove a low generalization error. Allen-Zhu and Li (2019a) show that using the same number of training examples, a three-layer ResNet can learn a function class with a much lower test error than any kernel method. Cao and Gu (2019a) assume that the NTK with respect to the second layer of a two-layer network can separate the data distribution, and prove that gradient descent on a deep network can achieve test error with samples and hidden units. Cao and Gu (2019b) consider SGD with an online oracle and give a general result. Under the same assumption as in (Cao and Gu, 2019a), their result requires hidden units and sample complexity . By contrast, with the same online oracle, our result only needs polylogarithmic hidden units and sample complexity .
2 Notation
Note that in this paper, denotes the -th row of at step . We fix and only train , as in (Li and Liang, 2018; Du et al., 2018b; Arora et al., 2019; Nitanda and Suzuki, 2019). We consider the ReLU activation , though our analysis can be extended easily to Lipschitz continuous, positively homogeneous activations such as leaky ReLU.
For any , the gradient descent step is given by . Also define
Note that . This property generally holds due to homogeneity: for any and any ,
and thus .
Empirical risk minimization
In this section, we consider a fixed training set and empirical risk minimization. We first state our assumption on the separability of the NTK, and then give our main result and a proof sketch.
The key idea of the NTK is to do the first-order Taylor approximation:
The infinite-width limit of eq. 2.1 is formalized as Assumption 2.1, with an additional bound on the norm of the separator. A concrete construction of using Assumption 2.1 is given in eq. 2.2.
and particularly define for the training input .
As discussed in Section 5, the space is the reproducing kernel Hilbert space (RKHS) induced by the infinite-width NTK with respect to , and maps into . Assumption 2.1 supposes that the induced training set can be separated by some , with an additional bound on which is crucial in our analysis. It is also possible to give a dual characterization of the separation margin (cf. eq. 5.2), which also allows us to show that Assumption 2.1 always holds when there are no parallel inputs (cf. Proposition 5.1). However, it is often more convenient to construct directly; see Section 5 for some examples.
With Assumption 2.1, we state our main empirical risk result.
Under Assumption 2.1, given any risk target and any , let
Then for any and any constant step size , with probability over the random initialization,
Moreover for any and any ,
While the number of hidden units required by prior work all have a polynomial dependency on , or , Theorem 2.2 only requires . The required width has a polynomial dependency on , which is an adaptive quantity: while can be for random labels (cf. Proposition 5.2), it can be when there is a strong feature-label relationship, for example on the noisy 2-XOR data introduced in (Wei et al., 2018) (cf. Proposition 5.3). Moreover, we show in Proposition 5.4 that if we want \mathinner{\bigl{\{}\mathinner{\left(\nabla f_{i}(W_{0}),y_{i}\right)}\bigr{\}}}_{i=1}^{n} to be separable, which is the starting point of an NTK-style analysis, the width has to depend polynomially on .
In the rest of Section 2, we give a proof sketch of Theorem 2.2. The full proof is given in Appendix A.
In this subsection, we give some nice properties of random initialization.
Given an initialization , for any , define
Lemma 2.3 ensures that with high probability has a positive margin at initialization.
Under Assumption 2.1, given any and any , if , then with probability , it holds simultaneously for all that
For any , any , and any , define
Lemma 2.4 controls . It will help us show that has a good margin during the training process.
Under the condition of Lemma 2.3, for any , with probability , it holds simultaneously for all that
Finally, Lemma 2.5 controls the output of the network at initialization.
Given any , if , then with probability , it holds simultaneously for all that
2 Convergence analysis of gradient descent
We analyze gradient descent in this subsection. First, define
For any and any , , and thus . Therefore by the triangle inequality, .
The quantity first appeared in the perceptron analysis (Novikoff, 1962) for the ReLU loss, and has also been analyzed in prior work (Ji and Telgarsky, 2018; Cao and Gu, 2019a; Nitanda and Suzuki, 2019). In this work, specifically helps us prove the following result, which plays an important role in obtaining a width which only depends on .
For any and any , if , then
Consequently, if we use a constant step size for , then
The proof of Lemma 2.6 starts from the standard iteration guarantee:
Using Lemmas 2.3, 2.4, 2.5 and 2.6, we can prove Theorem 2.2. Below is a proof sketch; the full proof is given in Appendix A.
We first show that as long as for all , it holds that \widehat{\mathcal{R}}^{(t)}\mathinner{\bigl{(}W_{0}+\lambda\overline{U}\bigr{)}}\leq\epsilon/4. To see this, let us consider first. For any , Lemma 2.5 ensures that is bounded, while Lemma 2.3 ensures that \big{\langle}\nabla f_{i}(W_{0}),\overline{U}\big{\rangle} is concentrated around with a large width. As a result, with the chosen in Theorem 2.2, we can show that \big{\langle}\nabla f_{i}(W_{0}),W_{0}+\lambda\overline{U}\big{\rangle} is large, and is small due to the exponential tail of the logistic loss. To further handle , we use a standard NTK argument to control \big{\langle}\nabla f_{i}(W_{t})-\nabla f_{i}(W_{0}),W_{0}+\lambda\overline{U}\big{\rangle} under the condition that .
We then prove by contradiction that the above bound on holds for at least the first iterations. The key observation is that as long as , we can use it and Lemma 2.6 to control , and then just invoke .
The quantity has also been considered in prior work (Cao and Gu, 2019a; Nitanda and Suzuki, 2019), where it is bounded by using the Cauchy-Schwarz inequality, which introduces a factor. To make the required width depend only on , we also need an upper bound on which depends only on . Since the above analysis results in a factor, and in our case steps are needed, it is unclear how to get a width using the analysis in (Cao and Gu, 2019a; Nitanda and Suzuki, 2019). By contrast, using Lemma 2.6, we can show that , which only depends on .
The claims of Theorem 2.2 then follow directly from the above two steps and Lemma 2.6.
Generalization
To get a generalization bound, we naturally extend Assumption 2.1 to the following assumption.
for almost all sampled from the data distribution .
The above assumption is also made in (Nitanda and Suzuki, 2019) for smooth activations. (Cao and Gu, 2019a) make a similar separability assumption, but in the RKHS induced by the second layer ; by contrast, Assumption 3.1 is on separability in the RKHS induced by the first layer .
Here is our test error bound with Assumption 3.1.
Under Assumption 3.1, given any and any , let and be given as in Theorem 2.2:
Then for any and any constant step size , with probability over the random initialization and data sampling,
where denotes the step with the minimum empirical risk before .
Below is a direct corollary of Theorem 3.2.
Under Assumption 3.1, given any , using a constant step size no larger than and let
it holds with probability that , where denotes the step with the minimum empirical risk in the first steps.
To get Theorem 3.2, we use a Lipschitz-based Rademacher complexity bound. One can also use a smoothness-based Rademacher complexity bound (Srebro et al., 2010, Theorem 1) and get a sample complexity . However, the bound will become complicated and some large constant will be introduced. It is an interesting open question to give a clean analysis based on smoothness.
Stochastic gradient descent
There are some different formulations of SGD. In this section, we consider SGD with an online oracle. We randomly sample and , and fix during training. At step , a data example is sampled from the data distribution. We still let , and perform the following update
Still with Assumption 3.1, we show the following result.
Under Assumption 3.1, given any , using a constant step size and , it holds with probability that
Below is a proof sketch of Theorem 4.1; the complete proof is given in Appendix C. For any and , define
The first step is an extension of Lemma 2.6 to the SGD setting, with a similar proof.
With a constant step size , for any and any ,
With Lemma 4.2, we can also extend Theorem 2.2 to the SGD setting and get a bound on , using a similar proof. To further get a bound on the cumulative population risk , the key observation is that is a martingale. Using a martingale Bernstein bound, we prove the following lemma; applying it finishes the proof of Theorem 4.1.
Given any , with probability ,
On separability
In this section we give some discussion on Assumption 2.1, the separability of the NTK. The proofs are all given in Appendix D.
Given a training set , the linear kernel is defined as . The maximum margin achievable by a linear classifier is given by
where denotes the probability simplex and denotes the Hadamard product. In addition to the dual definition eq. 5.1, when there also exists a maximum margin classifier which gives a primal characterization of : it holds that and for all .
In this paper we consider another kernel, the infinite-width NTK with respect to the first layer:
Here and are defined at the beginning of Section 2. Similar to the dual definition of , the margin given by is defined as
We can also give a primal characterization of when it is positive.
The proof is given in Appendix D, and uses the Fenchel duality theory. Using the upper bound , we can see that satisfies Assumption 2.1 with . However, such an upper bound might be too loose, which leads to a bad rate. In fact, as shown later, in some cases we can construct directly which satisfies Assumption 2.1 with a large . For this reason, we choose to make Assumption 2.1 instead of assuming a positive .
However, we can use to show that Assumption 2.1 always holds when there are no parallel inputs. Oymak and Soltanolkotabi (2019, Corollary I.2) prove that if for any two feature vectors and , we have and for some , then the minimum eigenvalue of is at least . For arbitrary labels , since , we have the worst case bound . A direct improvement of this bound is , where denotes the number of support vectors, which could be much smaller than with real world data.
On the other hand, given any training set which may have a large margin, replacing with random labels would destroy the margin, which is what should be expected.
Although the above bounds all have a polynomial dependency on , they hold for arbitrary or random labels, and thus do not assume any relationship between the features and labels. Next we give some examples where there is a strong feature-label relationship, and thus a much larger margin can be proved.
and thus Assumption 2.1 holds with .
2 The noisy 2-XOR distribution
We consider the noisy 2-XOR distribution introduced in (Wei et al., 2018). It is the uniform distribution over the following points:
The factor ensures that , and above denotes the Cartesian product. Here the label only depends on the first two coordinates of the input .
Then can de defined as follows. It only depends on the first two coordinates of .
The following result shows that . Note that could be as large as , in which case is basically .
For any sampled from the noisy 2-XOR distribution and any , it holds that
We can prove two other interesting results for the noisy 2-XOR data.
The first step of an NTK analysis is to show that \mathinner{\bigl{\{}\mathinner{\left(\nabla f_{i}(W_{0}),y_{i}\right)}\bigr{\}}}_{i=1}^{n} is separable. Proposition 5.4 gives an example where \mathinner{\bigl{\{}\mathinner{\left(\nabla f_{i}(W_{0}),y_{i}\right)}\bigr{\}}}_{i=1}^{n} is nonseparable when the network is narrow.
For the noisy 2-XOR data, the separator given by eq. 5.3 has margin , and . As a result, if we want \mathinner{\bigl{\{}\mathinner{\left(\nabla f_{i}(W_{0}),y_{i}\right)}\bigr{\}}}_{i=1}^{n} to be separable, the width has to be . For a smaller width, gradient descent might still be able to solve the problem, but a beyond-NTK analysis would be needed.
A tight sample complexity upper bound for the infinite-width NTK.
(Wei et al., 2018) give a sample complexity lower bound for any NTK classifier on the noisy 2-XOR data. It turns out that could give a matching sample complexity upper bound for the NTK and SGD.
(Wei et al., 2018) consider the infinite-width NTK with respect to both layers. For the first layer, the infinite-width NTK is defined in Section 5, and the corresponding RKHS and RKHS mapping is defined in Section 2. For the second layer, the infinite width NTK is defined by
The corresponding RKHS and inner product are given by
Open problems
In this paper, we analyze gradient descent on a two-layer network in the NTK regime, where the weights stay close to the initialization. It is an interesting open question if gradient descent learns something beyond the NTK, after the iterates move far enough from the initial weights. It is also interesting to extend our analysis to other architectures, such as multi-layer networks, convolutional networks, and residual networks. Finally, in this paper we only discuss binary classification; it is interesting to see if it is possible to get similar results for other tasks, such as regression.
The authors are grateful for support from the NSF under grant IIS-1750051, and from NVIDIA via a GPU grant.
References
Appendix A Omitted proofs from Section 2
By Assumption 2.1, given any ,
is the empirical mean of i.i.d. r.v.’s supported on with mean . Therefore by Hoeffding’s inequality, with probability ,
Applying a union bound finishes the proof. ∎
Given any fixed and ,
because is a standard Gaussian r.v. and the density of standard Gaussian has maximum . Since is the empirical mean of Bernoulli r.v.’s, by Hoeffding’s inequality, with probability ,
Applying a union bound finishes the proof. ∎
To prove Lemma 2.5, we need the following technical result.
and by further using the -Lipschitz continuity of , we have
Given , let . By Lemma A.1, is sub-Gaussian with variance proxy , and with probability at least over ,
On the other hand, by Jensen’s inequality,
As a result, with probability , it holds that . By a union bound, with probability over , for all , we have .
For any such that the above event holds, and for any , the r.v. is sub-Gaussian with variance proxy . By Hoeffding’s inequality, with probability over ,
By a union bound, with probability over , for all , we have .
The probability that the above events all happen is at least , over and . ∎
The second-order term of eq. A.1 can be bounded as follows
because , and , and . Combining eqs. A.1, A.2 and A.3 gives
The required width ensures that with probability , Lemmas 2.3, 2.4 and 2.5 hold with and .
Let denote the first step such that there exists with . Therefore for any and any , it holds that . In addition, we let .
We will split the left hand side into three terms and control them individually:
The first term of eq. A.4 can be controlled using Lemma 2.5:
The second term of eq. A.4 can be written as
Let S_{c}\mathrel{\mathop{\ordinarycolon}}=\mathinner{\left\{s\ {}\middle|\ {}\mathds{1}\mathinner{\bigl{[}\left\langle w_{s,t},x_{i}\right\rangle>0\bigr{]}}-\mathds{1}\mathinner{\bigl{[}\left\langle w_{s,0},x_{i}\right\rangle>0\bigr{]}}\neq 0,1\leq s\leq m\right\}}. Note that implies
where in the last step we use the condition that .
The third term of eq. A.4 can be bounded as follows: by Lemma 2.3,
where we use . Therefore,
Putting eqs. A.5, A.6 and A.7 into eq. A.4, we have
for the given in the statement of Theorem 2.2. Consequently, for any , it holds that \widehat{\mathcal{R}}^{(t)}\mathinner{\bigl{(}\overline{W}\bigr{)}}\leq\epsilon/4.
Let . The next claim is that . To see this, note that Lemma 2.6 ensures
Suppose , then we have , and thus . As a result, using and the definition of ,
Furthermore, by the triangle inequality, for any
which contradicts the definition of . Therefore .
Now we are ready to prove the claims of Theorem 2.2. The bound on follow by repeating the steps in eq. A.8. The risk guarantee follows from Lemma 2.6:
Appendix B Omitted proofs from Section 3
The proof of Theorem 3.2 is based on Rademacher complexity. Given a sample (where ) and a function class , the Rademacher complexity of on is defined as
We will use the following general result.
(Shalev-Shwartz and Ben-David, 2014, Theorem 26.5) If , then with probability ,
(Shalev-Shwartz and Ben-David, 2014, Lemma 26.9) .
To prove Theorem 3.2, we need one more Rademacher complexity bound. Given a fixed initialization , consider the following classes:
Given a feature sample , the following Lemma B.3 controls the Rademacher complexity of . A similar version was given in (Liang, 2016, Theorem 43), and the proof is similar to the proof of (Bartlett and Mendelson, 2002, Theorem 18) which also pushes the supremum through and handles each hidden unit separately.
.
Note that for any , the mapping is -Lipschitz, and thus Lemma B.2 gives
Invoking the Rademacher complexity of linear classifiers (Shalev-Shwartz and Ben-David, 2014, Lemma 26.10) then gives
Now we are ready to prove the main generalization result Theorem 3.2.
On the other hand, Theorem 2.2 ensures that under the conditions of Theorem 3.2, for any fixed dataset, with probability over the random initialization, we have
As a result, invoking eq. B.1 with , with probability over the random initialization and data sampling,
Invoking finishes the proof. ∎
Appendix C Omitted proofs from Section 4
Recall that , we have
and the second-order term of eq. C.1 can be bounded as follows
With Lemma 4.2, we give the following result, which is an extension of Theorem 2.2 to the SGD setting.
Under Assumption 3.1, given any , any , and any positive integer , let
For any and any constant step size , if , then with probability ,
We first sample data examples , and then feed to SGD at step . We only consider the first steps.
The proof is similar to the proof of Theorem 2.2. Let denote the first step before such that there exists some with . If such a step does not exist, let .
Let , in exactly the same way as in Theorem 2.2, we can show that with probability , for any ,
Now consider . Using Lemma 4.2, in the same way as the proof of Theorem 2.2 (replacing with , etc.), we can show that . Then invoking Lemma 4.2 again, we get
Next we prove Lemma 4.3. We need the following martingale Bernstein bound.
(Beygelzimer et al., 2011, Theorem 1) Let denote a martingale with and be the trivial -algebra. Let denote the corresponding martingale difference sequence, and let
denote the sequence of conditional variance. If a.s., then for any , with probability at least ,
For any , let denote , and denote . Note that the quantity is a martingale w.r.t. the filtration . The martingale difference sequence is given by , which satisfies
Invoking Lemma C.2 with eqs. C.4 and LABEL:eq:sgd_tmp2 gives that with probability ,
Suppose the condition of Lemma C.1 holds. Then we have for , with probability ,
Further invoking Lemma 4.3 gives that with probability ,
Since , we get
For the condition of Lemma C.1 to hold, it is enough to let
Appendix D Omitted proofs from Section 5
with optimal primal-dual solutions . Moreover
By strong duality, the inequality holds with equality. It follows that
Now let us look at the dual optimization problem. It is clear that
and thus . Since , we have that . In addition,
and thus has margin . Moreover, we have
and thus . Therefore, satisfies all requirements of Proposition 5.1. ∎
Let denote the uniform probability vector . Note that
Since for any , by Markov’s inequality with probability , it holds that , and thus . ∎
By symmetry, we only need to consider an where . Let denote , and similarly define . We have
For any nonzero , we have , and . Therefore
Let denote the density function of the standard Gaussian distribution, and for , let denote the probability that a standard Gaussian random variable lies in the interval :
Since is a Gaussian variable with standard deviation , we have
Plugging eqs. D.4 and D.5 into eq. D.3 gives:
For , it holds that , and thus
To prove Proposition 5.4, we need the following technical lemma.
Given and that are independent where , we have
First note that for which is independent of ,
Still let denote the density of , and let denote the probability that . We have
We now give the proof of Proposition 5.4 using Lemma D.1.
By symmetry, we only need to consider the following training set:
The factor is omitted also because we only discuss the loss.
For any , let denote the event that
We will show that if , then is true for all with probability , and Proposition 5.4 follows from the fact that the XOR data is not linearly separable.
Since is or or or , event will happen as long as
Note that while . As a result, due to Lemma D.1,