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 nn, 1/δ1/\delta or 1/ϵ1/\epsilon, where nn denotes the size of the training set, δ\delta denotes the failure probability, and ϵ\epsilon 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 Ω(ln⁡(n/δ)+ln⁡(1/ϵ)2)\Omega(\ln(n/\delta)+\ln(1/\epsilon)^{2}) hidden units trained by gradient descent achieves classification error ϵ\epsilon on test data, meaning both optimization and generalization occur. Unlike prior work, the width is fully polylogarithmic in nn, 1/δ1/\delta, and 1/ϵ1/\epsilon; 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 Ω~(1/ϵ2)\widetilde{\Omega}(1/\epsilon^{2}) samples and Θ~(1/ϵ)\widetilde{\Theta}(1/\epsilon) iterations, gradient descent finds a solution with ϵ\epsilon test error (cf. Theorem 3.2 and Corollary 3.3). (As discussed in Remark 3.4, Ω~(1/ϵ)\widetilde{\Omega}(1/\epsilon) 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 Θ~(1/ϵ)\widetilde{\Theta}(1/\epsilon) 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 O(1/n)O(1/\sqrt{n}) 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 Ω(1/ln⁡(n))\Omega(1/\ln(n)), 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 Ω(1/ϵ2)\Omega(1/\epsilon^{2}) hidden units, Ω~(1/ϵ4)\widetilde{\Omega}(1/\epsilon^{4}) data samples, and O(1/ϵ2)O(1/\epsilon^{2}) steps, while our result only needs polylogarithmic hidden units, Ω~(1/ϵ2)\widetilde{\Omega}(1/\epsilon^{2}) data samples, and O~(1/ϵ)\widetilde{O}(1/\epsilon) 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 ϵ\epsilon test error with Ω(1/ϵ4)\Omega(1/\epsilon^{4}) samples and Ω(1/ϵ14)\Omega(1/\epsilon^{14}) 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 Ω(1/ϵ14)\Omega(1/\epsilon^{14}) hidden units and sample complexity O~(1/ϵ2)\widetilde{O}(1/\epsilon^{2}). By contrast, with the same online oracle, our result only needs polylogarithmic hidden units and sample complexity O~(1/ϵ)\widetilde{O}(1/\epsilon).

2 Notation

Note that in this paper, ws,tw_{s,t} denotes the ss-th row of WW at step tt. We fix aa and only train WW, as in (Li and Liang, 2018; Du et al., 2018b; Arora et al., 2019; Nitanda and Suzuki, 2019). We consider the ReLU activation σ(z):=max⁡{0,z}\sigma(z)\mathrel{\mathop{\ordinarycolon}}=\max\mathinner{\left\{0,z\right\}}, though our analysis can be extended easily to Lipschitz continuous, positively homogeneous activations such as leaky ReLU.

For any t≥0t\geq 0, the gradient descent step is given by Wt+1:=Wt−ηt∇R^(Wt)W_{t+1}\mathrel{\mathop{\ordinarycolon}}=W_{t}-\eta_{t}\nabla\widehat{\mathcal{R}}(W_{t}). Also define

Note that fi(t)(Wt)=fi(Wt)f_{i}^{(t)}(W_{t})=f_{i}(W_{t}). This property generally holds due to homogeneity: for any WW and any 1≤s≤m1\leq s\leq m,

and thus ⟨∇fi(W),W⟩=fi(W)\left\langle\nabla f_{i}(W),W\right\rangle=f_{i}(W).

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 (2,∞)(2,\infty) norm of the separator. A concrete construction of U‾\overline{U} using Assumption 2.1 is given in eq. 2.2.

and particularly define ϕi:=ϕxi\phi_{i}\mathrel{\mathop{\ordinarycolon}}=\phi_{x_{i}} for the training input xix_{i}.

As discussed in Section 5, the space H\mathcal{H} is the reproducing kernel Hilbert space (RKHS) induced by the infinite-width NTK with respect to WW, and ϕx\phi_{x} maps xx into H\mathcal{H}. Assumption 2.1 supposes that the induced training set {(ϕi,yi)}i=1n\{(\phi_{i},y_{i})\}_{i=1}^{n} can be separated by some vˉ∈H\bar{v}\in\mathcal{H}, with an additional bound on  ⁣∥vˉ(z)∥2\mathinner{\!\left\lVert\bar{v}(z)\right\rVert}_{2} 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 vˉ\bar{v} 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 ϵ∈(0,1)\epsilon\in(0,1) and any δ∈(0,1/3)\delta\in(0,1/3), let

Then for any m≥Mm\geq M and any constant step size η≤1\eta\leq 1, with probability 1−3δ1-3\delta over the random initialization,

Moreover for any 0≤t<T0\leq t<T and any 1≤s≤m1\leq s\leq m,

While the number of hidden units required by prior work all have a polynomial dependency on nn, 1/δ1/\delta or 1/ϵ1/\epsilon, Theorem 2.2 only requires m=Ω(ln⁡(n/δ)+ln⁡(1/ϵ)2)m=\Omega\mathinner{\left(\ln(n/\delta)+\ln(1/\epsilon)^{2}\right)}. The required width has a polynomial dependency on 1/γ1/\gamma, which is an adaptive quantity: while 1/γ1/\gamma can be poly⁡(n)\operatorname{poly}(n) for random labels (cf. Proposition 5.2), it can be polylog⁡(n)\operatorname{polylog}(n) 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 1/γ1/\gamma.

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 (W0,a)(W_{0},a), for any 1≤s≤m1\leq s\leq m, define

Lemma 2.3 ensures that with high probability U‾\overline{U} has a positive margin at initialization.

Under Assumption 2.1, given any δ∈(0,1)\delta\in(0,1) and any ϵ1∈(0,γ)\epsilon_{1}\in(0,\gamma), if m≥(2ln⁡(n/δ))/ϵ12m\geq\mathinner{\left(2\ln(n/\delta)\right)}/\epsilon_{1}^{2}, then with probability 1−δ1-\delta, it holds simultaneously for all 1≤i≤n1\leq i\leq n that

For any WW, any ϵ2>0\epsilon_{2}>0, and any 1≤i≤n1\leq i\leq n, define

Lemma 2.4 controls αi(W0,ϵ2)\alpha_{i}(W_{0},\epsilon_{2}). It will help us show that U‾\overline{U} has a good margin during the training process.

Under the condition of Lemma 2.3, for any ϵ2>0\epsilon_{2}>0, with probability 1−δ1-\delta, it holds simultaneously for all 1≤i≤n1\leq i\leq n that

Finally, Lemma 2.5 controls the output of the network at initialization.

Given any δ∈(0,1)\delta\in(0,1), if m≥25ln⁡(2n/δ)m\geq 25\ln(2n/\delta), then with probability 1−δ1-\delta, it holds simultaneously for all 1≤i≤n1\leq i\leq n that

2 Convergence analysis of gradient descent

We analyze gradient descent in this subsection. First, define

For any WW and any 1≤s≤m1\leq s\leq m,  ⁣∥∂fi/∂ws∥2≤1/m\mathinner{\!\left\lVert\partial f_{i}/\partial w_{s}\right\rVert}_{2}\leq 1/\sqrt{m}, and thus  ⁣∥∇fi(W)∥F≤1\mathinner{\!\left\lVert\nabla f_{i}(W)\right\rVert}_{F}\leq 1. Therefore by the triangle inequality,  ⁣∥∇R^(W)∥F≤Q^(W)\mathinner{\!\left\lVert\nabla\widehat{\mathcal{R}}(W)\right\rVert}_{F}\leq\widehat{\mathcal{Q}}(W).

The quantity Q^\widehat{\mathcal{Q}} 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, Q^\widehat{\mathcal{Q}} specifically helps us prove the following result, which plays an important role in obtaining a width which only depends on polylog⁡(1/ϵ)\operatorname{polylog}(1/\epsilon).

For any t≥0t\geq 0 and any W‾\overline{W}, if ηt≤1\eta_{t}\leq 1, then

Consequently, if we use a constant step size η≤1\eta\leq 1 for 0≤τ<t0\leq\tau<t, 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 ∥ws,t−ws,0∥2≤4λ/(γm)\|w_{s,t}-w_{s,0}\|_{2}\leq 4\lambda/(\gamma\sqrt{m}) for all 1≤s≤m1\leq s\leq m, it holds that \widehat{\mathcal{R}}^{(t)}\mathinner{\bigl{(}W_{0}+\lambda\overline{U}\bigr{)}}\leq\epsilon/4. To see this, let us consider R^(0)\widehat{\mathcal{R}}^{(0)} first. For any 1≤i≤n1\leq i\leq n, Lemma 2.5 ensures that ∣⟨∇fi(W0),W0⟩∣|\langle\nabla f_{i}(W_{0}),W_{0}\rangle| is bounded, while Lemma 2.3 ensures that \big{\langle}\nabla f_{i}(W_{0}),\overline{U}\big{\rangle} is concentrated around γ\gamma with a large width. As a result, with the chosen λ\lambda 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 R^(0)(W0+λU‾)\widehat{\mathcal{R}}^{(0)}(W_{0}+\lambda\overline{U}) is small due to the exponential tail of the logistic loss. To further handle R^(t)\widehat{\mathcal{R}}^{(t)}, 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 ∥ws,t−ws,0∥2≤4λ/(γm)\|w_{s,t}-w_{s,0}\|_{2}\leq 4\lambda/(\gamma\sqrt{m}).

We then prove by contradiction that the above bound on ∥ws,t−ws,0∥2\|w_{s,t}-w_{s,0}\|_{2} holds for at least the first TT iterations. The key observation is that as long as R^(t)(W0+λU‾)≤ϵ/4\widehat{\mathcal{R}}^{(t)}(W_{0}+\lambda\overline{U})\leq\epsilon/4, we can use it and Lemma 2.6 to control ∑τ<tQ^(Wτ)\sum_{\tau<t}\widehat{\mathcal{Q}}(W_{\tau}), and then just invoke ∥ws,t−ws,0∥2≤η∑τ<tQ^(Wτ)/m\|w_{s,t}-w_{s,0}\|_{2}\leq\eta\sum_{\tau<t}\widehat{\mathcal{Q}}(W_{\tau})/\sqrt{m}.

The quantity ∑τ<tQ^(Wτ)\sum_{\tau<t}\widehat{\mathcal{Q}}(W_{\tau}) has also been considered in prior work (Cao and Gu, 2019a; Nitanda and Suzuki, 2019), where it is bounded by t∑τ<tQ^(Wτ)2\sqrt{t}\sqrt{\sum_{\tau<t}\widehat{\mathcal{Q}}(W_{\tau})^{2}} using the Cauchy-Schwarz inequality, which introduces a t\sqrt{t} factor. To make the required width depend only on polylog⁡(1/ϵ)\operatorname{polylog}(1/\epsilon), we also need an upper bound on ∑τ<tQ^(Wτ)\sum_{\tau<t}\widehat{\mathcal{Q}}(W_{\tau}) which depends only on polylog⁡(1/ϵ)\operatorname{polylog}(1/\epsilon). Since the above analysis results in a t\sqrt{t} factor, and in our case Ω(1/ϵ)\Omega(1/\epsilon) steps are needed, it is unclear how to get a polylog⁡(1/ϵ)\operatorname{polylog}(1/\epsilon) width using the analysis in (Cao and Gu, 2019a; Nitanda and Suzuki, 2019). By contrast, using Lemma 2.6, we can show that ∑τ<tQ^(Wτ)≤4λ/γ\sum_{\tau<t}\widehat{\mathcal{Q}}(W_{\tau})\leq 4\lambda/\gamma, which only depends on ln⁡(1/ϵ)\ln(1/\epsilon).

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 (x,y)(x,y) sampled from the data distribution D\mathcal{D}.

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 aa; by contrast, Assumption 3.1 is on separability in the RKHS induced by the first layer WW.

Here is our test error bound with Assumption 3.1.

Under Assumption 3.1, given any ϵ∈(0,1)\epsilon\in(0,1) and any δ∈(0,1/4)\delta\in(0,1/4), let λ\lambda and MM be given as in Theorem 2.2:

Then for any m≥Mm\geq M and any constant step size η≤1\eta\leq 1, with probability 1−4δ1-4\delta over the random initialization and data sampling,

where kk denotes the step with the minimum empirical risk before ⌈\nicefrac2λ2ηϵ⌉\lceil\nicefrac{{2\lambda^{2}}}{{\eta\epsilon}}\rceil.

Below is a direct corollary of Theorem 3.2.

Under Assumption 3.1, given any ϵ,δ∈(0,1)\epsilon,\delta\in(0,1), using a constant step size no larger than 11 and let

it holds with probability 1−δ1-\delta that P(x,y)∼D(yf(x;Wk,a)≤0)≤ϵP_{(x,y)\sim\mathcal{D}}\mathinner{\left(yf(x;W_{k},a)\leq 0\right)}\leq\epsilon, where kk denotes the step with the minimum empirical risk in the first Θ~(\nicefrac1γ2ϵ)\widetilde{\Theta}(\nicefrac{{1}}{{\gamma^{2}\epsilon}}) 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 O~(\nicefrac1γ4ϵ)\widetilde{O}(\nicefrac{{1}}{{\gamma^{4}\epsilon}}). 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 W0W_{0} and aa, and fix aa during training. At step ii, a data example (xi,yi)(x_{i},y_{i}) is sampled from the data distribution. We still let fi(W):=f(xi;W,a)f_{i}(W)\mathrel{\mathop{\ordinarycolon}}=f(x_{i};W,a), and perform the following update

Still with Assumption 3.1, we show the following result.

Under Assumption 3.1, given any ϵ,δ∈(0,1)\epsilon,\delta\in(0,1), using a constant step size and m=Ω(\nicefrac(ln⁡(1/δ)+ln⁡(1/ϵ)2)γ8)m=\Omega\mathinner{\left(\nicefrac{{\mathinner{\left(\ln(1/\delta)+\ln(1/\epsilon)^{2}\right)}}}{{\gamma^{8}}}\right)}, it holds with probability 1−δ1-\delta that

Below is a proof sketch of Theorem 4.1; the complete proof is given in Appendix C. For any ii and WW, define

The first step is an extension of Lemma 2.6 to the SGD setting, with a similar proof.

With a constant step size η≤1\eta\leq 1, for any W‾\overline{W} and any i≥0i\geq 0,

With Lemma 4.2, we can also extend Theorem 2.2 to the SGD setting and get a bound on ∑i<nQi(Wi)\sum_{i<n}\mathcal{Q}_{i}(W_{i}), using a similar proof. To further get a bound on the cumulative population risk ∑i<nQ(Wi)\sum_{i<n}\mathcal{Q}(W_{i}), the key observation is that ∑i<n(Q(Wi)−Qi(Wi))\sum_{i<n}\mathinner{\left(\mathcal{Q}(W_{i})-\mathcal{Q}_{i}(W_{i})\right)} is a martingale. Using a martingale Bernstein bound, we prove the following lemma; applying it finishes the proof of Theorem 4.1.

Given any δ∈(0,1)\delta\in(0,1), with probability 1−δ1-\delta,

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 {(xi,yi)}i=1n\mathinner{\left\{(x_{i},y_{i})\right\}}_{i=1}^{n}, the linear kernel is defined as K0(xi,xj):=⟨xi,xj⟩K_{0}(x_{i},x_{j})\mathrel{\mathop{\ordinarycolon}}=\left\langle x_{i},x_{j}\right\rangle. The maximum margin achievable by a linear classifier is given by

where Δn\Delta_{n} denotes the probability simplex and ⊙\odot denotes the Hadamard product. In addition to the dual definition eq. 5.1, when γ0>0\gamma_{0}>0 there also exists a maximum margin classifier uˉ\bar{u} which gives a primal characterization of γ0\gamma_{0}: it holds that ∥uˉ∥2=1\|\bar{u}\|_{2}=1 and yi⟨uˉ,xi⟩≥γ0y_{i}\left\langle\bar{u},x_{i}\right\rangle\geq\gamma_{0} for all ii.

In this paper we consider another kernel, the infinite-width NTK with respect to the first layer:

Here ϕ\phi and H\mathcal{H} are defined at the beginning of Section 2. Similar to the dual definition of γ0\gamma_{0}, the margin given by K1K_{1} is defined as

We can also give a primal characterization of γ1\gamma_{1} when it is positive.

The proof is given in Appendix D, and uses the Fenchel duality theory. Using the upper bound  ⁣∥v^(z)∥2≤1/γ1\mathinner{\!\left\lVert\hat{v}(z)\right\rVert}_{2}\leq 1/\gamma_{1}, we can see that γ1v^\gamma_{1}\hat{v} satisfies Assumption 2.1 with γ≥γ12\gamma\geq\gamma_{1}^{2}. However, such an upper bound  ⁣∥v^(z)∥2≤1/γ1\mathinner{\!\left\lVert\hat{v}(z)\right\rVert}_{2}\leq 1/\gamma_{1} might be too loose, which leads to a bad rate. In fact, as shown later, in some cases we can construct vˉ\bar{v} directly which satisfies Assumption 2.1 with a large γ\gamma. For this reason, we choose to make Assumption 2.1 instead of assuming a positive γ1\gamma_{1}.

However, we can use γ1\gamma_{1} 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 xix_{i} and xjx_{j}, we have ∥xi−xj∥2≥θ\|x_{i}-x_{j}\|_{2}\geq\theta and ∥xi+xj∥2≥θ\|x_{i}+x_{j}\|_{2}\geq\theta for some θ>0\theta>0, then the minimum eigenvalue of K1K_{1} is at least θ/(100n2)\theta/(100n^{2}). For arbitrary labels y∈{−1,+1}ny\in\{-1,+1\}^{n}, since  ⁣∥q⊙y∥2≥1/n\mathinner{\!\left\lVert q\odot y\right\rVert}_{2}\geq 1/\sqrt{n}, we have the worst case bound γ12≥\nicefracθ100n3\gamma_{1}^{2}\geq\nicefrac{{\theta}}{{100n^{3}}}. A direct improvement of this bound is \nicefracθ100nS3\nicefrac{{\theta}}{{100n_{S}^{3}}}, where nSn_{S} denotes the number of support vectors, which could be much smaller than nn with real world data.

On the other hand, given any training set {(xi,yi)}i=1n\mathinner{\left\{(x_{i},y_{i})\right\}}_{i=1}^{n} which may have a large margin, replacing yy with random labels would destroy the margin, which is what should be expected.

Although the above bounds all have a polynomial dependency on nn, 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 γ=γ0/2\gamma=\gamma_{0}/2.

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 2d2^{d} points:

The factor \nicefrac1d−1\nicefrac{{1}}{{\sqrt{d-1}}} ensures that ∥x∥2=1\|x\|_{2}=1, and ×\times above denotes the Cartesian product. Here the label yy only depends on the first two coordinates of the input xx.

Then vˉ\bar{v} can de defined as follows. It only depends on the first two coordinates of zz.

The following result shows that γ=Ω(1/d)\gamma=\Omega(1/d). Note that nn could be as large as 2d2^{d}, in which case γ\gamma is basically O(1/ln⁡(n))O\mathinner{\left(1/\ln(n)\right)}.

For any (x,y)(x,y) sampled from the noisy 2-XOR distribution and any d≥3d\geq 3, 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 vˉ\bar{v} given by eq. 5.3 has margin γ=Ω(1/d)\gamma=\Omega(1/d), and 1/γ=O(d)1/\gamma=O(d). 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 Ω(1/γ)\Omega(1/\sqrt{\gamma}). 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 d2d^{2} sample complexity lower bound for any NTK classifier on the noisy 2-XOR data. It turns out that γ\gamma 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 K1K_{1} is defined in Section 5, and the corresponding RKHS H\mathcal{H} and RKHS mapping ϕ\phi is defined in Section 2. For the second layer, the infinite width NTK is defined by

The corresponding RKHS K\mathcal{K} and inner product ⟨w1,w2⟩K\langle w_{1},w_{2}\rangle_{\mathcal{K}} 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 1≤i≤n1\leq i\leq n,

is the empirical mean of i.i.d. r.v.’s supported on [−1,+1][-1,+1] with mean μ\mu. Therefore by Hoeffding’s inequality, with probability 1−\nicefracδn1-\nicefrac{{\delta}}{{n}},

Applying a union bound finishes the proof. ∎

Given any fixed ϵ2\epsilon_{2} and 1≤i≤n1\leq i\leq n,

because ⟨w,xi⟩\left\langle w,x_{i}\right\rangle is a standard Gaussian r.v. and the density of standard Gaussian has maximum 1/2π1/\sqrt{2\pi}. Since αi(W0,ϵ2)\alpha_{i}(W_{0},\epsilon_{2}) is the empirical mean of Bernoulli r.v.’s, by Hoeffding’s inequality, with probability 1−\nicefracδn1-\nicefrac{{\delta}}{{n}},

Applying a union bound finishes the proof. ∎

To prove Lemma 2.5, we need the following technical result.

and by further using the 11-Lipschitz continuity of σ\sigma, we have

Given 1≤i≤n1\leq i\leq n, let hi=σ(W0xi)/mh_{i}=\sigma(W_{0}x_{i})/\sqrt{m}. By Lemma A.1, ∥hi∥2\|h_{i}\|_{2} is sub-Gaussian with variance proxy 1/m1/m, and with probability at least 1−\nicefracδ2n1-\nicefrac{{\delta}}{{2n}} over W0W_{0},

On the other hand, by Jensen’s inequality,

As a result, with probability 1−\nicefracδ2n1-\nicefrac{{\delta}}{{2n}}, it holds that ∥hi∥2≤1\|h_{i}\|_{2}\leq 1. By a union bound, with probability 1−\nicefracδ21-\nicefrac{{\delta}}{{2}} over W0W_{0}, for all 1≤i≤n1\leq i\leq n, we have ∥hi∥2≤1\|h_{i}\|_{2}\leq 1.

For any W0W_{0} such that the above event holds, and for any 1≤i≤n1\leq i\leq n, the r.v. ⟨hi,a⟩\left\langle h_{i},a\right\rangle is sub-Gaussian with variance proxy ∥hi∥22≤1\|h_{i}\|_{2}^{2}\leq 1. By Hoeffding’s inequality, with probability 1−\nicefracδ2n1-\nicefrac{{\delta}}{{2n}} over aa,

By a union bound, with probability 1−\nicefracδ21-\nicefrac{{\delta}}{{2}} over aa, for all 1≤i≤n1\leq i\leq n, we have  ⁣∣f(xi;W0,a)∣≤2ln⁡(4n/δ)\mathinner{\!\left\lvert f(x_{i};W_{0},a)\right\rvert}\leq\sqrt{2\ln\mathinner{\left(4n/\delta\right)}}.

The probability that the above events all happen is at least (1−\nicefracδ2)(1−\nicefracδ2)≥1−δ(1-\nicefrac{{\delta}}{{2}})(1-\nicefrac{{\delta}}{{2}})\geq 1-\delta, over W0W_{0} and aa. ∎

The second-order term of eq. A.1 can be bounded as follows

because  ⁣∥∇R^(Wt)∥F≤Q^(Wt)\mathinner{\!\left\lVert\nabla\widehat{\mathcal{R}}(W_{t})\right\rVert}_{F}\leq\widehat{\mathcal{Q}}(W_{t}), and ηt,Q^(Wt)≤1\eta_{t},\widehat{\mathcal{Q}}(W_{t})\leq 1, and Q^(Wt)≤R^(Wt)\widehat{\mathcal{Q}}(W_{t})\leq\widehat{\mathcal{R}}(W_{t}). Combining eqs. A.1, A.2 and A.3 gives

The required width ensures that with probability 1−3δ1-3\delta, Lemmas 2.3, 2.4 and 2.5 hold with ϵ1=γ2/8\epsilon_{1}=\gamma^{2}/8 and ϵ2=4λ/(γm)\epsilon_{2}=4\lambda/(\gamma\sqrt{m}).

Let t1t_{1} denote the first step such that there exists 1≤s≤m1\leq s\leq m with  ⁣∥ws,t1−ws,0∥2>4λ/(γm)\mathinner{\!\left\lVert w_{s,t_{1}}-w_{s,0}\right\rVert}_{2}>4\lambda/(\gamma\sqrt{m}). Therefore for any 0≤t<t10\leq t<t_{1} and any 1≤s≤m1\leq s\leq m, it holds that  ⁣∥ws,t−ws,0∥2≤4λ/(γm)\mathinner{\!\left\lVert w_{s,t}-w_{s,0}\right\rVert}_{2}\leq 4\lambda/(\gamma\sqrt{m}). In addition, we let W‾:=W0+λU‾\overline{W}\mathrel{\mathop{\ordinarycolon}}=W_{0}+\lambda\overline{U}.

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 s∈Scs\in S_{c} implies

where in the last step we use the condition that m≥4096λ2/γ6m\geq 4096\lambda^{2}/\gamma^{6}.

The third term of eq. A.4 can be bounded as follows: by Lemma 2.3,

where we use m≥4096λ2/γ6m\geq 4096\lambda^{2}/\gamma^{6}. Therefore,

Putting eqs. A.5, A.6 and A.7 into eq. A.4, we have

for the λ\lambda given in the statement of Theorem 2.2. Consequently, for any 0≤t<t10\leq t<t_{1}, it holds that \widehat{\mathcal{R}}^{(t)}\mathinner{\bigl{(}\overline{W}\bigr{)}}\leq\epsilon/4.

Let T:=⌈\nicefrac2λ2ηϵ⌉T\mathrel{\mathop{\ordinarycolon}}=\lceil\nicefrac{{2\lambda^{2}}}{{\eta\epsilon}}\rceil. The next claim is that t1≥Tt_{1}\geq T. To see this, note that Lemma 2.6 ensures

Suppose t1<Tt_{1}<T, then we have t1≤\nicefrac2λ2ηϵt_{1}\leq\nicefrac{{2\lambda^{2}}}{{\eta\epsilon}}, and thus  ⁣∥Wt1−W‾∥F2≤2λ2\mathinner{\!\left\lVert W_{t_{1}}-\overline{W}\right\rVert}_{F}^{2}\leq 2\lambda^{2}. As a result, using ∥U‾∥F≤1\|\overline{U}\|_{F}\leq 1 and the definition of W‾\overline{W},

Furthermore, by the triangle inequality, for any 1≤s≤m1\leq s\leq m

which contradicts the definition of t1t_{1}. Therefore t1≥Tt_{1}\geq T.

Now we are ready to prove the claims of Theorem 2.2. The bound on  ⁣∥ws,t−ws,0∥2\mathinner{\!\left\lVert w_{s,t}-w_{s,0}\right\rVert}_{2} 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 S=(z1,…,zn)S=(z_{1},\ldots,z_{n}) (where zi=(xi,yi)z_{i}=(x_{i},y_{i})) and a function class H\mathcal{H}, the Rademacher complexity of H\mathcal{H} on SS is defined as

We will use the following general result.

(Shalev-Shwartz and Ben-David, 2014, Theorem 26.5) If h(z)∈[a,b]h(z)\in[a,b], then with probability 1−δ1-\delta,

(Shalev-Shwartz and Ben-David, 2014, Lemma 26.9) Rad(g∘F∘X)≤KRad(F∘X)\textup{Rad}\mathinner{\left(g\circ\mathcal{F}\circ X\right)}\leq K\textup{Rad}\mathinner{\left(\mathcal{F}\circ X\right)}.

To prove Theorem 3.2, we need one more Rademacher complexity bound. Given a fixed initialization (W0,a)(W_{0},a), consider the following classes:

Given a feature sample XX, the following Lemma B.3 controls the Rademacher complexity of Fρ∘X\mathcal{F}_{\rho}\circ X. 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.

Rad(Fρ∘X)≤ρm/n\textup{Rad}\mathinner{\left(\mathcal{F}_{\rho}\circ X\right)}\leq\rho\sqrt{m/n}.

Note that for any 1≤s≤m1\leq s\leq m, the mapping z↦asσ(z)z\mapsto a_{s}\sigma(z) is 11-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 1−3δ1-3\delta over the random initialization, we have

As a result, invoking eq. B.1 with ρ=4λ/(γm)\rho=4\lambda/(\gamma\sqrt{m}), with probability 1−4δ1-4\delta over the random initialization and data sampling,

Invoking P(x,y)∼D(yf(x;W,a)≤0)≤2Q(W)P_{(x,y)\sim\mathcal{D}}\mathinner{\left(yf(x;W,a)\leq 0\right)}\leq 2\mathcal{Q}(W) finishes the proof. ∎

Appendix C Omitted proofs from Section 4

Recall that  ⁣∥∇ft(Wt)∥F≤1\mathinner{\!\left\lVert\nabla f_{t}(W_{t})\right\rVert}_{F}\leq 1, 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 ϵ∈(0,1)\epsilon\in(0,1), any δ∈(0,1/3)\delta\in(0,1/3), and any positive integer n0n_{0}, let

For any m≥Mm\geq M and any constant step size η≤1\eta\leq 1, if n0≥n:=⌈\nicefrac2λ2ηϵ⌉n_{0}\geq n\mathrel{\mathop{\ordinarycolon}}=\lceil\nicefrac{{2\lambda^{2}}}{{\eta\epsilon}}\rceil, then with probability 1−3δ1-3\delta,

We first sample n0n_{0} data examples (x0,y0),…,(xn0−1,yn0−1)(x_{0},y_{0}),\ldots,(x_{n_{0}-1},y_{n_{0}-1}), and then feed (xi,yi)(x_{i},y_{i}) to SGD at step ii. We only consider the first n0n_{0} steps.

The proof is similar to the proof of Theorem 2.2. Let n1n_{1} denote the first step before n0n_{0} such that there exists some 1≤s≤m1\leq s\leq m with  ⁣∥ws,n1−ws,0∥2>4λ/(γm)\mathinner{\!\left\lVert w_{s,n_{1}}-w_{s,0}\right\rVert}_{2}>4\lambda/(\gamma\sqrt{m}). If such a step does not exist, let n1=n0n_{1}=n_{0}.

Let W‾:=W0+λU‾\overline{W}\mathrel{\mathop{\ordinarycolon}}=W_{0}+\lambda\overline{U}, in exactly the same way as in Theorem 2.2, we can show that with probability 1−3δ1-3\delta, for any 0≤i<n10\leq i<n_{1},

Now consider n:=⌈\nicefrac2λ2ηϵ⌉n\mathrel{\mathop{\ordinarycolon}}=\lceil\nicefrac{{2\lambda^{2}}}{{\eta\epsilon}}\rceil. Using Lemma 4.2, in the same way as the proof of Theorem 2.2 (replacing Q^(Wτ)\widehat{\mathcal{Q}}(W_{\tau}) with Qi(Wi)\mathcal{Q}_{i}(W_{i}), etc.), we can show that n≤n1n\leq n_{1}. 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 (Mt,Ft)t≥0(M_{t},\mathcal{F}_{t})_{t\geq 0} denote a martingale with M0=0M_{0}=0 and F0\mathcal{F}_{0} be the trivial σ\sigma-algebra. Let (Δt)t≥1(\Delta_{t})_{t\geq 1} denote the corresponding martingale difference sequence, and let

denote the sequence of conditional variance. If Δt≤R\Delta_{t}\leq R a.s., then for any δ∈(0,1)\delta\in(0,1), with probability at least 1−δ1-\delta,

For any i≥0i\geq 0, let ziz_{i} denote (xi,yi)(x_{i},y_{i}), and z0,iz_{0,i} denote (z0,…,zi)(z_{0},\ldots,z_{i}). Note that the quantity ∑t<i(Q(Wt)−Qt(Wt))\sum_{t<i}\mathinner{\left(\mathcal{Q}(W_{t})-\mathcal{Q}_{t}(W_{t})\right)} is a martingale w.r.t. the filtration σ(z0,i−1)\sigma(z_{0,i-1}). The martingale difference sequence is given by Q(Wt)−Qt(Wt)\mathcal{Q}(W_{t})-\mathcal{Q}_{t}(W_{t}), which satisfies

Invoking Lemma C.2 with eqs. C.4 and LABEL:eq:sgd_tmp2 gives that with probability 1−δ1-\delta,

Suppose the condition of Lemma C.1 holds. Then we have for n=⌈\nicefrac2λ2ηϵ⌉n=\lceil\nicefrac{{2\lambda^{2}}}{{\eta\epsilon}}\rceil, with probability 1−3δ1-3\delta,

Further invoking Lemma 4.3 gives that with probability 1−4δ1-4\delta,

Since P(x,y)∼D(yf(x;W,a)≤0)≤2Q(W)P_{(x,y)\sim\mathcal{D}}\mathinner{\left(yf(x;W,a)\leq 0\right)}\leq 2\mathcal{Q}(W), 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 (wˉ,qˉ)(\bar{w},\bar{q}). 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 f∗(A∗qˉ)=γ12/2f^{*}(A^{*}\bar{q})=\gamma_{1}^{2}/2. Since wˉ=A∗qˉ\bar{w}=A^{*}\bar{q}, we have that  ⁣∥wˉ∥H=γ1\mathinner{\!\left\lVert\bar{w}\right\rVert}_{\mathcal{H}}=\gamma_{1}. In addition,

and thus −wˉ-\bar{w} has margin γ12\gamma_{1}^{2}. Moreover, we have

and thus  ⁣∥wˉ(z)∥2≤1\mathinner{\!\left\lVert\bar{w}(z)\right\rVert}_{2}\leq 1. Therefore, v^=−wˉ/γ1\hat{v}=-\bar{w}/\gamma_{1} satisfies all requirements of Proposition 5.1. ∎

Let q^\hat{q} denote the uniform probability vector (\nicefrac1n,…,\nicefrac1n)(\nicefrac{{1}}{{n}},\ldots,\nicefrac{{1}}{{n}}). Note that

Since 0≤(q^⊙ϵ)⊤K1(q^⊙ϵ)≤10\leq\mathinner{\left(\hat{q}\odot\epsilon\right)}^{\top}K_{1}\mathinner{\left(\hat{q}\odot\epsilon\right)}\leq 1 for any ϵ\epsilon, by Markov’s inequality with probability 0.90.9, it holds that (q^⊙ϵ)⊤K1(q^⊙ϵ)≤1/(20n)\mathinner{\left(\hat{q}\odot\epsilon\right)}^{\top}K_{1}\mathinner{\left(\hat{q}\odot\epsilon\right)}\leq 1/(20n), and thus γ1≤1/20n\gamma_{1}\leq 1/\sqrt{20n}. ∎

By symmetry, we only need to consider an (x,y)(x,y) where (x1,x2,y)=(\nicefrac1d−1,0,1)(x_{1},x_{2},y)=(\nicefrac{{1}}{{\sqrt{d-1}}},0,1). Let zp,qz_{p,q} denote (zp,zp+1,…,zq)(z_{p},z_{p+1},\ldots,z_{q}), and similarly define xp,qx_{p,q}. We have

For any nonzero p∈A1p\in A_{1}, we have −p∈A3-p\in A_{3}, and ⟨vˉ(p),x1,2⟩=1/d−1\left\langle\bar{v}(p),x_{1,2}\right\rangle=1/\sqrt{d-1}. Therefore

Let φ\varphi denote the density function of the standard Gaussian distribution, and for c>0c>0, let U(c)U(c) denote the probability that a standard Gaussian random variable lies in the interval [−c,c][-c,c]:

Since ⟨q,x3,d⟩\left\langle q,x_{3,d}\right\rangle is a Gaussian variable with standard deviation \nicefrac(d−2)(d−1)\sqrt{\nicefrac{{(d-2)}}{{(d-1)}}}, we have

Plugging eqs. D.4 and D.5 into eq. D.3 gives:

For t∈[−1,+1]t\in[-1,+1], it holds that φ(t)≥12πe\varphi(t)\geq 1\sqrt{2\pi e}, and thus

To prove Proposition 5.4, we need the following technical lemma.

Given z1∼N(0,1)z_{1}\sim\mathcal{N}(0,1) and z2∼N(0,b2)z_{2}\sim\mathcal{N}(0,b^{2}) that are independent where b>1b>1, we have

First note that for z3∼N(0,1)z_{3}\sim\mathcal{N}(0,1) which is independent of z1z_{1},

Still let φ\varphi denote the density of N(0,1)\mathcal{N}(0,1), and let U(c)U(c) denote the probability that z3∈[−c,c]z_{3}\in[-c,c]. 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 1/d−11/\sqrt{d-1} factor is omitted also because we only discuss the 0/10/1 loss.

For any ss, let AsA_{s} denote the event that

We will show that if m≤d−2/4m\leq\sqrt{d-2}/4, then AsA_{s} is true for all 1≤s≤m1\leq s\leq m with probability 1/21/2, and Proposition 5.4 follows from the fact that the XOR data is not linearly separable.

Since ((xi)1,(xi)2)\mathinner{\left((x_{i})_{1},(x_{i})_{2}\right)} is (1,0)(1,0) or (0,1)(0,1) or (−1,0)(-1,0) or (0,−1)(0,-1), event AsA_{s} will happen as long as

Note that (ws)1,(ws)2∼N(0,1)(w_{s})_{1},(w_{s})_{2}\sim\mathcal{N}(0,1) while ∑j=3d(ws)j∼N(0,d−2)\sum_{j=3}^{d}(w_{s})_{j}\sim\mathcal{N}(0,d-2). As a result, due to Lemma D.1,