Beyond Linearization: On Quadratic and Higher-Order Approximation of Wide Neural Networks

Yu Bai, Jason D. Lee

Introduction

Deep Learning has made remarkable impact on a variety of artificial intelligence applications such as computer vision, reinforcement learning, and natural language processing. Though immensely successful, theoretical understanding of deep learning lags behind. It is not understood how non-linear neural networks can be efficiently trained to approximate complex decision boundaries with a relatively few number of training samples.

There has been a recent surge of research on connecting neural networks trained via gradient descent with the neural tangent kernel (NTK) (Jacot et al., 2018; Du et al., 2018a, b; Chizat and Bach, 2018b; Allen-Zhu et al., 2018a; Arora et al., 2019a, b). This line of analysis proceeds by coupling the training dynamics of the nonlinear network with the training dynamics of its linearization in a local neighborhood of the initialization, and then analyzing the expressiveness and generalization of the network via the corresponding properties of its linearized model.

Though powerful, NTK is not yet a completely satisfying theory for explaining the success of deep learning in practice. In theory, the expressive power of the linearized model is roughly the same as, and thus limited to, that of the corresponding random feature space (Allen-Zhu et al., 2018a; Wei et al., 2019) or the Reproducing Kernel Hilbert Space (RKHS) (Bietti and Mairal, 2019). While these spaces can approximate any regular (e.g. bounded Lipschitz) function up to arbitrary accuracy, the norm of the approximators can be exponentially large in the feature dimension for certain non-smooth but very simple functions such as a single ReLU (Yehudai and Shamir, 2019). Using NTK analyses, the sample complexity bound for learning these functions can be poor whereas experimental evidence suggests that the sample complexity is mild (Livni et al., 2014). In practice, kernel machines with the NTK have been experimentally demonstrated to yield competitive results on large-scale tasks such as image classification on CIFAR-10; yet there is still a non-neglible performance gap between NTK and full training on the same convolutional architecture (Arora et al., 2019a; Lee et al., 2019). It is an increasingly compelling question whether we can establish theories for training neural networks beyond the NTK regime.

In this paper, we study the optimization and generalization of over-parametrized two-layer neural networks via relating to their higher-order approximations, a principled generalization of the NTK. Our theory starts from the fact that a two-layer neural network fW0+W(x)f_{{\mathbf{W}}_{0}+{\mathbf{W}}}({\mathbf{x}}) (with smooth activation) can be Taylor expanded with respect to the weight matrix W{\mathbf{W}} as

Above, fW0f_{{\mathbf{W}}_{0}} does not depend on W{\mathbf{W}}, and f(1)f^{(1)} corresponds to the NTK model, which is the dominant W{\mathbf{W}}-dependent term when {wr}{\left\{{\mathbf{w}}_{r}\right\}} are small and leads to the coupling between the gradient dynamics for training neural net and its NTK f(1)f^{(1)}.

Our key observation is that the dominance of f(1)f^{(1)} is deduced from comparing the upper bounds—rather than the actual values—of fW0,W(k)(x)f^{(k)}_{{\mathbf{W}}_{0},{\mathbf{W}}}({\mathbf{x}}). It is a priori possible that there exists a subset of W{\mathbf{W}}’s in which the dominating term is not f(1)f^{(1)} but some other f(k)f^{(k)}, k≥2k\geq 2. If we were able to train in that set, the gradient dynamics would be coupled with the dynamics on f(k)f^{(k)} rather than f(1)f^{(1)} and thus could be very different. That learning is coupled with f(k)f^{(k)} could further offer possibilities for expressing certain functions with parameters of lower complexities, or generalizing better, as f(k)f^{(k)} is no longer a linearized model. In this paper, we build on this perspective and identify concrete regimes in which neural net learning is coupled with higher-order f(k)f^{(k)}’s rather than its linearization.

The contribution of this paper can be summarized as follows.

We demonstrate that after randomization, the linear NTK f(1)f^{(1)} is no longer the dominant term, and so the gradient dynamics of the neural net is no longer coupled with NTK. Through a simple sign randomization, the training loss of an over-parametrized two-layer neural network can be coupled with that of a quadratic model (Section 3). We prove that the randomized neural net loss exhibits a nice optimization landscape in that every second-order stationary point has training loss not much higher than the best quadratic model, making it amenable to efficient minimization (Section 4).

We establish results on the generalization and expressive power of such randomized neural nets (Section 5). These results lead to sample complexity bounds for learning certain simple functions that matches the NTK without distributional assumptions and are advantageous when mild isotropic assumptions on the feature are present. In particular, using randomized networks, the sample complexity bound for learning polynomials (and their linear combination) on (relatively) uniform base distributions is O(d)O(d) lower than using NTK.

We show that the randomization technique can be generalized to find neural nets that are dominated by the kk-th order term in their Taylor series (k>2k>2) which we term as higher-order NTKs. These models also have expressive power similar as the linear NTK, and potentially even better generalization and sample complexity (Section 6 & Appendix D).

We review prior work on the optimization, generalization, and expressivity of neural networks.

Neal (1996) first proposed the connection between infinite-width networks and kernel methods. Later work (Daniely et al., 2016; Williams, 1997; Lee et al., 2018; Novak et al., 2019; Matthews et al., 2018) extended this connection to various settings including deep networks and deep convolutional networks. These works established that gradient descent on only the output layer weights is well-approximated by a kernel method for large width.

More recently, several groups discovered the connection between gradient descent on all the parameters and the neural tangent kernel (Jacot et al., 2018). Li and Liang (2018); Du et al. (2018b) utilized the coupling of the gradient dynamics to prove that gradient descent finds global minimizers of the training loss of two-layer networks, and Du et al. (2018a); Allen-Zhu et al. (2018b); Zou et al. (2018) generalized this to deep residual and convolutional networks. Using the NTK coupling, Arora et al. (2019b); Cao and Gu (2019a, b) proved generalization error bounds that match the kernel method.

Despite the close theoretical connection between NTK and training deep networks, Arora et al. (2019a); Lee et al. (2019); Chizat and Bach (2018b) empirically found a significant performance gap between NTK and actual training. This gap has been theoretically studied in Wei et al. (2019); Allen-Zhu and Li (2019); Yehudai and Shamir (2019); Ghorbani et al. (2019a) which established that NTK has provably higher generalization error than training the neural net for specific data distributions and architectures.

The idea of randomization is initiated by Allen-Zhu et al. (2018a), who use randomization to provably learn a three-layer network; however it is unclear how the sample complexity of their algorithm compares against the NTK. Inspired by their work, we study the potential gains of coupling with a non-linear approximation over the linear NTK — we compare the performance of a quadratic approximation model with the linear NTK on two-layer networks and find that under mild data assumptions the quadratic approximation reduces sample complexity under mild data assumptions.

It is believed that the success of SGD is largely due to its algorithmic regularization effects. A large body of work Li et al. (2017); Nacson et al. (2019); Gunasekar et al. (2018b, a, 2017); Woodworth et al. (2019) shows that asymptotically gradient descent converges to a max-margin solution with a strong regularization effect, unlike the NTK regularizationAs a concrete example, Woodworth et al. (2019) showed that for matrix completion the NTK solution estimates zero on all unobserved entries and the max-margin solution corresponds to the minimum nuclear norm solution..

For two-layer networks, a series of works used the mean field method to establish the evolution of the network parameters via a Wasserstein gradient flow (Mei et al., 2018b; Chizat and Bach, 2018a; Wei et al., 2018; Rotskoff and Vanden-Eijnden, 2018; Sirignano and Spiliopoulos, 2018). In the mean field regime, the parameters move significantly from their initialization, unlike NTK regime, however it is unclear if the dynamics converge to solutions of low training loss.

Finally, Li et al. (2019) showed how a combination of large learning rate and injected noise amplifies the regularization from the noise and outperforms the NTK of the corresponding architecture.

Many prior works have tried to establish favorable landscape properties such as every local minimum is a global minimum (Ge et al., 2017; Du and Lee, 2018; Soltanolkotabi et al., 2018; Hardt and Ma, 2016; Freeman and Bruna, 2016; Nguyen and Hein, 2017a, b; Haeffele and Vidal, 2015; Venturi et al., 2018). Combining with existing advances in gradient descent avoiding saddle-points (Ge et al., 2015; Lee et al., 2016; Jin et al., 2017), these show that gradient descent find the global minimum. Notably, Du and Lee (2018); Ge et al. (2017) show that gradient descent converges to solutions also of low test error, with lower sample complexity than their corresponding NTKs.

Recently, researchers have studied norm-based generalization based (Bartlett et al., 2017; Neyshabur et al., 2015; Golowich et al., 2017), tighter compression-based bounds (Arora et al., 2018), and PAC-Bayes bounds (Dziugaite and Roy, 2017; Neyshabur et al., 2017) that identify properties of the parameter that allow for efficient generalization.

Preliminaries

denote respectively the empirical risk and population risk for any predictor f:X→Yf:\mathcal{X}\to\mathcal{Y}.

We consider learning an over-parametrized two-layer neural network of the form

Throughout this paper we assume that the activation is second-order smooth in the following sense.

An example is the cubic ReLU σ(t)=relu3(t)=max⁡{t,0}3\sigma(t)={\rm relu}^{3}(t)=\max{\left\{t,0\right\}}^{3}. The reason for requiring σ\sigma to be higher-order smooth (and thus excluding ReLU) will be made clear in the subsequent textWe note that the only restrictive requirement in Assumption A is the Lipschitzness of σ′′\sigma^{\prime\prime}, which guarantees second-order smoothness of the objectives. The bounds on derivatives (and specifically their bound near zero) are merely for technical convenience and can be weakened without hurting the results..

1 Notation

Escaping NTK via randomization

To motivate our study, we now briefly review the NTK theory for over-parametrized neural nets and provide insights on how to go beyond the NTK regime.

Let W0{\mathbf{W}}_{0} denote the weights in a two-layer neural network at initialization and W{\mathbf{W}} denote its movement from W0{\mathbf{W}}_{0} (so that the current weight matrix is W0+W{\mathbf{W}}_{0}+{\mathbf{W}}.) The observation in NTK theory, or the theory of lazy training (Chizat and Bach, 2018b), is that for small W{\mathbf{W}} the neural network fW0+Wf_{{\mathbf{W}}_{0}+{\mathbf{W}}} can be Taylor expanded as

so that the network can be decomposed as the sum of the initial network fW0f_{{\mathbf{W}}_{0}}, the linearized model fWLf^{L}_{{\mathbf{W}}}, and higher order terms. Specifically (ignoring fW0f_{{\mathbf{W}}_{0}} for the moment), when mm is large and ∥wr∥2=O(m−1/2)\left\|{{\mathbf{w}}_{r}}\right\|_{2}=O(m^{-1/2}), we expect fWL=O(1)f^{L}_{{\mathbf{W}}}=O(1) and higher order terms to be om(1)o_{m}(1), which is indeed the regime when we train fW0+Wf_{{\mathbf{W}}_{0}+{\mathbf{W}}} via gradient descent. Therefore, the trajectory of training fW0+Wf_{{\mathbf{W}}_{0}+{\mathbf{W}}} is coupled with the trajectory of training fW0+fWLf_{{\mathbf{W}}_{0}}+f^{L}_{\mathbf{W}}, which is a convex problem and enjoys convergence guarantees (Du et al., 2018b).

Our goal is to find subsets of W{\mathbf{W}} so that the dominating term is not fLf^{L} but something else in the higher order part. The above expansion makes clear that this cannot be achieved through simple fixes such as tuning the leading scale 1/m1/\sqrt{m} or the learning rate — the domination of fLf^{L} appears to hold so long as the movements wr{\mathbf{w}}_{r} are small.

then the second-order Taylor expansion of fW0+WΣf_{{\mathbf{W}}_{0}+{\mathbf{W}}{\mathbf{\Sigma}}} can be written as

where we have defined in addition the quadratic part fWΣQf^{Q}_{{\mathbf{W}}{\mathbf{\Sigma}}}. Due to the existence of {Σrr}\{\Sigma_{rr}\}, each original weight wr{\mathbf{w}}_{r} now has an additional a scalar that is different in fLf^{L} and fQf^{Q}. Specifically, if we choose

More precisely, when ∥wr∥2≍m−1/4\left\|{{\mathbf{w}}_{r}}\right\|_{2}\asymp m^{-1/4}, the scalings of fLf^{L} and fQf^{Q} compare as follows:

so we expect fWΣL(x)=O(m−1/4)f^{L}_{{\mathbf{W}}{\mathbf{\Sigma}}}({\mathbf{x}})=O(m^{-1/4}) over a random draw of Σ{\mathbf{\Sigma}}.

Therefore, at the random weight matrix WΣ{\mathbf{W}}{\mathbf{\Sigma}}, fQf^{Q} dominates fLf^{L} and thus the network is coupled with its quadratic part rather than the linear NTK.

1 Learning randomized neural nets

The randomization technique leads to the following recipe for learning W{\mathbf{W}}: train W{\mathbf{W}} so that ∥wr∥2=O(m−1/4)\left\|{{\mathbf{w}}_{r}}\right\|_{2}=O(m^{-1/4}) and WΣ{\mathbf{W}}{\mathbf{\Sigma}} has in expectation low loss. We make this precise by formulating the problem as minimizing a randomized neural net risk.

where we have reparametrized the weight matrix into W0+W{\mathbf{W}}_{0}+{\mathbf{W}} so that learning starts at W=0{\mathbf{W}}={\mathbf{0}}.

Following our randomization recipe, we now formulate our problem as minimizing the expected risk

Our regularizer penalizes W{\mathbf{W}}, i.e. the distance from initialization, similar as in (Hu et al., 2019).Our specific choice of ∥⋅∥2,4\left\|{\cdot}\right\|_{2,4} norm is needed for measuring the average magnitude of fWQf^{Q}_{\mathbf{W}}, whereas the high (8-th) power is not essential and can be replaced by any (4+ε)(4+\varepsilon)-th power without affecting the result.

We initialize the parameters (a,W0)({\mathbf{a}},{\mathbf{W}}_{0}) randomly in the following way: set

Above, we set half of the aia_{i}’s as +1+1 and half as −1-1, and the weights w0,r{\mathbf{w}}_{0,r} are i.i.d. in the +1+1 half and copied exactly into the −1-1 half. Such an initialization is almost equivalent to i.i.d. random W0{\mathbf{W}}_{0}, but has the additional benefit that fW0(x)≡0f_{{\mathbf{W}}_{0}}({\mathbf{x}})\equiv 0 and also leads to simple expressivity arguments. Our initialization scale Bx−2B_{x}^{-2} is chosen so that for a random draw of w0{\mathbf{w}}_{0}, we have w0⊤x∼N(0,1){\mathbf{w}}_{0}^{\top}{\mathbf{x}}\sim\mathsf{N}(0,1), which is on average O(1)O(1)Our choice covers two commonly used scales in neural net analyses: Bx=1B_{x}=1, w0,r∼N(0,Id){\mathbf{w}}_{0,r}\sim\mathsf{N}(0,I_{d}) in e.g. (Arora et al., 2019b; Allen-Zhu et al., 2018a); Bx=dB_{x}=\sqrt{d}, w0,r∼N(0,Id/d){\mathbf{w}}_{0,r}\sim\mathsf{N}(0,I_{d}/d) in e.g. (Ghorbani et al., 2019b).. For technical convenience, we also assume henceforth that the realized {w0,r}{\left\{{\mathbf{w}}_{0,r}\right\}} satisfies the bound

This happens with probability at least 1−δ1-\delta under random initialization (see proof in Appendix A.3), and ensures that max⁡r∈[m]∣w0,r⊤x∣≤O~(d)\max_{r\in[m]}|{\mathbf{w}}_{0,r}^{\top}{\mathbf{x}}|\leq\widetilde{O}(\sqrt{d}) simultaneously for all x{\mathbf{x}}.

Optimization

In this section, we show that LλL_{\lambda} enjoys a nice optimization landscape.

As the randomized loss LL induces coupling of the neural net fW0+WΣf_{{\mathbf{W}}_{0}+{\mathbf{W}}{\mathbf{\Sigma}}} with the quadratic model fWQf^{Q}_{\mathbf{W}}, we expect its behavior to resemble the behavior of gradient descent on the following clean risk:

We now show that the clean risk LQL^{Q}, albeit non-convex, possesses a nice optimization landscape.

This result implies that, for W{\mathbf{W}} in a certain ball and large mm, every point of higher loss than W⋆{\mathbf{W}}_{\star} will have either a first-order or a second-order descent direction. In other words, every approximate second-order stationary point of LQL^{Q} is also an approximate global minimum. Our proof utilizes the fact that LQL^{Q} is similar to the loss function in matrix sensing / learning quadratic neural networks, and builds on recent understandings that the landscapes of these problems are often nice (Soltanolkotabi et al., 2018; Du and Lee, 2018; Allen-Zhu et al., 2018a). The proof is deferred to Appendix B.1.

2 Nice landscape of randomized neural net risk

Suppose there exists W⋆∈B2,4(Bw,⋆){\mathbf{W}}_{\star}\in{\sf B}_{2,4}(B_{w,\star}) such that LQ(W⋆)≤OPTL^{Q}({\mathbf{W}}_{\star})\leq{\sf OPT}, and that

for some fixed ε∈(0,1]\varepsilon\in(0,1] and Bw≥Bw,⋆B_{w}\geq B_{w,\star}, then for all W∈B2,4(Bw){\mathbf{W}}\in{\sf B}_{2,4}(B_{w}), we have

As an immediate corollary, we have a similar characterization of the regularized loss LλL_{\lambda}.

For any Bw≥Bw,⋆B_{w}\geq B_{w,\star}, under the conditions of Theorem 2, we have for all λ>0\lambda>0 and all W∈B2,4(Bw){\mathbf{W}}\in{\sf B}_{2,4}(B_{w}) that

Theorem 2 follows directly from Lemma 1 through the coupling between LL and LQL^{Q} (as well as their gradients and Hessians). Corollary 3 then follows by controlling in addition the effect of the regularizer. The full proof of Theorem 2 and Corollary 3 are deferred to Appendices B.4 and B.5.

We now present our main optimization result, which follows directly from Corollary 3.

Suppose there exists W⋆{\mathbf{W}}_{\star} such that

for some OPT>0{\sf OPT}>0. For any γ=Θ(1)\gamma=\Theta(1) and ε>0\varepsilon>0, we can choose λ\lambda suitably and m≥O~(poly(d,BxBw,⋆,ε−1))m\geq\widetilde{O}({\rm poly}(d,B_{x}B_{w,\star},\varepsilon^{-1})) such that the regularized loss LλL_{\lambda} satisfies the following: any second order stationary point W^\widehat{{\mathbf{W}}} has low loss and bounded norm:

Proof sketch. The proof of Theorem 4 consists of two stages: first “localize” any second-order stationary point into a (potentially very big) norm ball using the ∥⋅∥2,48\left\|{\cdot}\right\|_{2,4}^{8} regularizer, then use Corollary 3 in this ball to further deduce that LλL_{\lambda} is low and ∥W^∥2,4≤O(∥W⋆∥2,4)\left\|{\widehat{{\mathbf{W}}}}\right\|_{2,4}\leq O(\left\|{{\mathbf{W}}_{\star}}\right\|_{2,4}). The full proof is deferred to Appendix B.6.

Generalization and Expressivity

We now shift attention to studying the generalization and expressivity of the (randomized) neural net W^\widehat{{\mathbf{W}}} learned in Theorem 4.

As W^\widehat{{\mathbf{W}}} is always coupled (through randomization) with the quadratic model fW^Qf^{Q}_{\widehat{{\mathbf{W}}}}, we begin by studying the generalization of the quadratic model.

where σi∼iidUnif{±1}\sigma_{i}\stackrel{{\scriptstyle\rm iid}}{{\sim}}{\rm Unif}{\left\{\pm 1\right\}} are Rademacher variables.

Lemma 5 suggests a possibility for the quadratic model to generalize better than the NTK model: the Rademacher complexity of FQ(Bw){\mathcal{F}}^{Q}(B_{w}) depends on the “feature maps” 1n∑i=1nσiσ′′(w0,r⊤xi)xixi⊤\frac{1}{n}\sum_{i=1}^{n}\sigma_{i}\sigma^{\prime\prime}({\mathbf{w}}_{0,r}^{\top}{\mathbf{x}}_{i}){\mathbf{x}}_{i}{\mathbf{x}}_{i}^{\top} through their matrix operator norm. Compared with the (naive) Frobenius norm based generalization bounds, the operator norm is never worse and can be better when additional structure on x{\mathbf{x}} is present. The proof of Lemma 5 is deferred to Appendix C.1.

We now state our main generalization bound on the (randomized) neural net loss LL, which concretizes the above insight.

For any data-dependent W^\widehat{{\mathbf{W}}} such that ∥W^∥2,4≤Bw\left\|{\widehat{{\mathbf{W}}}}\right\|_{2,4}\leq B_{w}, we have

The generalization bound in Theorem 6 features two desirable properties:

For large mm (e.g. m≳n4m\gtrsim n^{4}), the bound scales at most logarithmically with the width mm, therefore allowing learning with small samples and extreme over-parametrization;

Theorem 6 follows directly from Lemma 5 and a matrix concentration Lemma. The proof is deferred to Appendix C.2.

2 Expressivity and Sample Complexity through Quadratic Models

In order to concretize our generalization result, we now study the expressive power of quadratic models through the concrete example of learning functions of the form ∑j≤kαj(βj⊤x)pj\sum_{j\leq k}\alpha_{j}({\bm{\beta}}_{j}^{\top}{\mathbf{x}})^{p_{j}}, i.e. sum of “one-directional” polynomials (for consistency and comparability with (Arora et al., 2019b).)

The proof of Theorem 7 is based on a reduction from expressing degree pp polynomials using quadratic models to expressing degree p−2p-2 polynomials using random feature models. The proof can be found in Appendix C.4.

We now illustrate our results in Theorem 6 and 7 in three concrete examples, in which we compare the sample complexity bounds of the randomized (quadratic) network and the linear NTK when mm is sufficiently large.

Learning a single polynomial. Suppose f⋆(x)=α(β⊤x)pf_{\star}({\mathbf{x}})=\alpha({\bm{\beta}}^{\top}{\mathbf{x}})^{p} satisfies L(f⋆)≤ϵL(f_{\star})\leq\epsilon, and we wish to find W^\widehat{{\mathbf{W}}} with O(ε)O(\varepsilon) test loss. By Theorem 7 we can choose W⋆{\mathbf{W}}_{\star} such that LQ(W⋆)≤OPT=2εL^{Q}({\mathbf{W}}_{\star})\leq{\sf OPT}=2\varepsilon, and by Theorem 4 we can find W^\widehat{{\mathbf{W}}} such that L(W^)≤Lλ(W^)≤3εL(\widehat{{\mathbf{W}}})\leq L_{\lambda}(\widehat{{\mathbf{W}}})\leq 3\varepsilon and ∥W^∥2,4=O(Bw,⋆)\|\widehat{{\mathbf{W}}}\|_{2,4}=O(B_{w,\star}). Take Bx=1B_{x}=1, and assume x{\mathbf{x}} is sufficiently isotropic so that Mx,op=O(1d)M_{x,{\rm op}}=O(\frac{1}{\sqrt{d}}), the sample complexity from Theorem 6 is

In contrast, the sample complexity for linear NTK (Arora et al., 2019b; Cao and Gu, 2019a) to reach ϵ\epsilon test loss is

We have nQ/nL=O~(p/d)n_{Q}/n_{L}=\widetilde{O}(p/d), a reduction by a dimension factor unless p≍dp\asymp d. We note that the above comparison is simply comparing upper bounds, since in general the lower bound on the sample complexity of linear NTK is unknown.

Learning a noisy 22-XOR. Wei et al. (2019) established a sample complexity lower bound of linear NTK of n≥nL=Ω(d2)n\geq n_{L}=\Omega(d^{2}) to achieve constant generalization error on the noisy 22-XOR problem, which allows for a rigorous comparison against the quadratic model.

The ground truth function in 22-XOR is f⋆(x)=x1x2=([(e1+e2)⊤x]2−[(e1−e2)⊤x]2)/4f_{\star}({\mathbf{x}})=x_{1}x_{2}=([({\mathbf{e}}_{1}+{\mathbf{e}}_{2})^{\top}{\mathbf{x}}]^{2}-[({\mathbf{e}}_{1}-{\mathbf{e}}_{2})^{\top}{\mathbf{x}}]^{2})/4, where x∈{±1}d{\mathbf{x}}\in\{\pm 1\}^{d}, and f⋆f_{\star} attains constant margin on the training distribution constructed in Wei et al. (2019). By Theorem 7, f⋆f_{\star} can be ε\varepsilon-approximated by fW⋆Qf^{Q}_{{\mathbf{W}}_{\star}} with Bw,⋆4≤O(1)B_{w,\star}^{4}\leq O(1). Thus by Theorem 6 the sample complexity for learning noisy 22-XOR through the randomized net W^\widehat{{\mathbf{W}}} is

This is O~(d)\widetilde{O}(d) better than the sample complexity lower bound of linear NTK and thus provably better.

through the randomized net W^\widehat{{\mathbf{W}}} is

This compares favorably against the sample complexity upper bound for linear NTK, which needs

Higher-order NTKs

In this section, we demonstrate that our idea of randomization for changing the dynamics of learning neural networks can be generalized systematically — through randomization we are able to obtain over-parametrized neural networks in which the kk-th order term dominates the Taylor series. Consider a two-layer neural network with 2m2m neurons and symmetric initialization (cf. (3))

where we have defined the kk-th order NTK

Note that f(0)(x)≡0f^{(0)}({\mathbf{x}})\equiv 0 due to the symmetric initialization, and f(1)(x)f^{(1)}({\mathbf{x}}) is the standard NTK. For an arbitrary W{\mathbf{W}} such that ∥w+,r∥2,∥w−,r∥2=om(1)\left\|{{\mathbf{w}}_{+,r}}\right\|_{2},\left\|{{\mathbf{w}}_{-,r}}\right\|_{2}=o_{m}(1), we expect that f(1)(x)f^{(1)}({\mathbf{x}}) is the dominating term in the expansion.

We now describe an approach to finding W{\mathbf{W}} so that

that is, the neural net is approximately the kk-th order NTK plus an error term that goes to zero as m→∞m\to\infty, thereby “escaping” the NTK regime. Our approach builds on the following randomization technique: let z+z_{+}, z−z_{-} be two random variables (distributions) such that

Set (w+,r,w−,r)=(z+,rw⋆,r,z−,rw⋆,r)({\mathbf{w}}_{+,r},{\mathbf{w}}_{-,r})=(z_{+,r}{\mathbf{w}}_{\star,r},z_{-,r}{\mathbf{w}}_{\star,r}), and take ∥w⋆,r∥2=O(m−1/2k)\left\|{{\mathbf{w}}_{\star,r}}\right\|_{2}=O(m^{-1/2k}), we have

Therefore, with high probability, all f(1),…,f(k−1)f^{(1)},\dots,f^{(k-1)} as well as the remainder term f−∑j≤kf(j)f-\sum_{j\leq k}f^{(j)} has order O(m−1/2k)O(m^{-1/2k}), and the kk-th order NTK f(k)f^{(k)} can express an O(1)O(1) function.

We establish the generalization of expressivity of f(k)f^{(k)} in Appendix D, which systematically extends our results on the quadratic model. We show that the sample complexity for learning degree ≥k\geq k polynomials through f(k)f^{(k)} compared with linear NTK can be better by a factor of dk−1d^{k-1} for large nn, when mild distributional assumptions on x{\mathbf{x}} such as approximate isotropy (constant condition number of the kthk^{th} moment tensor) is present.

Conclusion

In this paper we proposed and studied the optimization and generalization of over-parametrized neural networks through coupling with higher-order terms in their Taylor series. Through coupling with the quadratic model, we showed that the randomized two-layer neural net has a nice optimization landscape (every second-order stationary point has low loss) and is thus amenable to efficient minimization through escape-saddle style algorithms. These networks enjoy the same expressivity and generalization guarantees as linearized models but in addition can generalize better by a dimension factor when distributional assumptions are present. We extended the idea of randomization to show the existence of neural networks whose Taylor series is dominated by the kk-th order term.

We believe our work brings in a number of open questions, such as how to better utilize the expressivity of quadratic models, or whether the study of higher-order expansions can lead to a more satisfying theory for explaining the success of full training. We also note that the Taylor series is only one avenue to obtaining accurate approximations of nonlinear neural networks. It would be of interest to design other approximation schemes for neural networks that are coupled with the network in larger regions of the parameter space.

Acknowledgment

The authors would like to thank Wei Hu, Tengyu Ma, Song Mei, and Andrea Montanari for their insightful comments. JDL acknowledges support of the ARO under MURI Award W911NF-11-1-0303, the Sloan Research Fellowship, and NSF CCF #1900145. The majority of this work was done while YB was at Stanford University. The authors also thank the Simons Institute Summer 2019 program on the Foundations of Deep Learning, and the Institute of Advanced Studies Special Year on Optimization, Statistics, and Theoretical Machine Learning for hosting the authors.

References

Appendix A Technical tools

Suppose {Ar,i}r∈[m],i∈[n]{\left\{{\mathbf{A}}_{r,i}\right\}}_{r\in[m],i\in[n]} are fixed symmetric d×dd\times d matrices, and {σi}i∈[n]∼iidUnif{±1}{\left\{\sigma_{i}\right\}}_{i\in[n]}\stackrel{{\scriptstyle\rm iid}}{{\sim}}{\rm Unif}{\left\{\pm 1\right\}} are Rademacher variables. Letting

Applying the high-probability bound in (Tropp et al., 2015, Theorem 4.6.1) and the union bound, we get

Let V:=max⁡r∈[m]v(Yr)V\mathrel{\mathop{:}}=\max_{r\in[m]}v({\mathbf{Y}}_{r}), we have by integrating the above bound over tt that

A.2 Expressing polynomials with random features

and the infimum over aa is attainable whenever it is finite.

For the ReLU random feature kernel KK, let u:=x⊤x′/Bx2u\mathrel{\mathop{:}}={\mathbf{x}}^{\top}{\mathbf{x}}^{\prime}/B_{x}^{2} and N2(ρ)\mathsf{N}_{2}(\rho) denote a bivariate normal distribution with marginals N(0,1)\mathsf{N}(0,1) and correlation ρ∈\rho\in. We have that

where the constants {cp}{\left\{c_{p}\right\}} satisfy

we have K(x,x′)=⟨ϕ(x),ϕ(x′)⟩K({\mathbf{x}},{\mathbf{x}}^{\prime})=\left\langle\phi({\mathbf{x}}),\phi({\mathbf{x}}^{\prime})\right\rangle. With this feature map, the function f⋆(x)=α(β⊤x)pf_{\star}({\mathbf{x}})=\alpha(\beta^{\top}{\mathbf{x}})^{p} can be represented as

Thus by the feature map equivalence (11), we have f⋆∈HKf_{\star}\in\mathcal{H}_{K} and

Now apply the feature map equivalence (11) again with the random feature map

A.3 Proof of Equation (4)

Setting t=8(dlog⁡5+log⁡(m/δ))=O(d+log⁡(m/δ))t=\sqrt{8(d\log 5+\log(m/\delta))}=O(\sqrt{d+\log(m/\delta)}) ensures that the above probability does not exceed δ\delta as desired. ∎

Appendix B Proofs for Section 4

Computing the gradient of LQL^{Q}, we obtain

where the last step used Cauchy-Schwarz on {∥wr∥2}{\left\{\left\|{{\mathbf{w}}_{r}}\right\|_{2}\right\}} and {∥w⋆,r∥2}{\left\{\left\|{{\mathbf{w}}_{\star,r}}\right\|_{2}\right\}}.

Term I does not involve Σ′{\mathbf{\Sigma}}^{\prime} and can be deterministically bounded as

B.2 Coupling lemmas

for all x∈Sd−1(Bx){\mathbf{x}}\in S^{d-1}(B_{x}).

∣ΔWΣQ(x)∣≤O(Bx3∥W∥2,43m−1/4)|\Delta^{Q}_{{\mathbf{W}}{\mathbf{\Sigma}}}({\mathbf{x}})|\leq O(B_{x}^{3}\left\|{{\mathbf{W}}}\right\|_{2,4}^{3}m^{-1/4}) (almost surely for all Σ{\mathbf{\Sigma}}.)

Above, (i) follows from the assumption that ∣σ′(t)∣≤Ct2|\sigma^{\prime}(t)|\leq Ct^{2}, (ii) is Cauchy-Schwarz, (iii) uses the bound (4), and (iv) uses the power mean inequality on ∥wr∥2\left\|{{\mathbf{w}}_{r}}\right\|_{2}.

We have by the Lipschitzness of σ′′\sigma^{\prime\prime} that

where again (i) uses the power mean inequality on ∥wr∥2\left\|{{\mathbf{w}}_{r}}\right\|_{2}.

B.3 Closeness of landscapes

Differentiating LL and LQL^{Q} and taking the inner product with W{\mathbf{W}}, we get

Therefore, by expanding σ′((w0,r+Σrrwr)⊤x)\sigma^{\prime}(({\mathbf{w}}_{0,r}+\Sigma_{rr}{\mathbf{w}}_{r})^{\top}{\mathbf{x}}) and noticing that Σrr2≡1\Sigma_{rr}^{2}\equiv 1, we have

where (i) uses Cauchy-Schwarz and (ii) uses the bounds in Lemma 10 and 11. For term III we first note by the smoothness of σ′\sigma^{\prime} that

Substituting this bound into term III yields

Putting together the bounds for term I, II, III gives the desired result. ∎

Differentiating LL and LQL^{Q} twice on the direction W⋆Σ′{\mathbf{W}}_{\star}{\mathbf{\Sigma}}^{\prime}, we get

We first bound the terms I(L){\rm I}(L) and I(LQ){\rm I}(L^{Q}). We have

Using similar arguments on I(LQ){\rm I}(L^{Q}) gives the bound

We now shift attention to bounding II(L)−II(LQ){\rm II}(L)-{\rm II}(L^{Q}). First note that

Then we have, by applying the bounds in Lemma 10 and 11,

Combining all the bounds gives the desired result. ∎

B.4 Proof of Theorem 2

We apply Lemma 12, 13, and 14 to connect the neural net loss LL to the “clean risk” LQL^{Q}. First, by Lemma 12, we have for all the assumed W{\mathbf{W}} that

Therefore we have ∣L(W)−LQ(W)∣≤ε/6{\left|L({\mathbf{W}})-L^{Q}({\mathbf{W}})\right|}\leq\varepsilon/6 so long as

provided that the error term in Lemma 1 is bounded by ε/3\varepsilon/3, which happens when

Finally, we choose mm sufficiently large so that

which combined with (14) yields the desired result. By Lemma 13 and 14, it suffices to choose mm such that, to satisfy the closeness of directional gradients,

and to satisfy the closeness of Hessian quadratic forms,

Collecting the requirements on mm in (13), (15), (16), (17) and merging terms using ε≤1\varepsilon\leq 1 and Bw,⋆≤BwB_{w,\star}\leq B_{w}, the desired result holds whenever

B.5 Proof of Corollary 3

Recall that Lλ(W)=L(W)+λ∥W⋆∥2,48L_{\lambda}({\mathbf{W}})=L({\mathbf{W}})+\lambda\left\|{{\mathbf{W}}_{\star}}\right\|_{2,4}^{8}. By differentiating A↦∥A∥2,48{\mathbf{A}}\mapsto\left\|{{\mathbf{A}}}\right\|_{2,4}^{8} we get

where (i) used Cauchy-Schwarz and (ii) used the AM-GM inequality p3q≤αp4/4+27q4/(4α3)p^{3}q\leq\alpha p^{4}/4+27q^{4}/(4\alpha^{3}) for all p,qp,q and α>0\alpha>0. Substituting the above expressions into Aλ−A0A_{\lambda}-A_{0} yields

Choosing α=5/14\alpha=5/14 gives the desired result. ∎

B.6 Proof of Theorem 4

We begin by choosing the regularization strength as

where λ0\lambda_{0} is a constant to be determined. Let ε\varepsilon be an accuracy parameter also to be determined.

Now, applying the coupling Lemma 13, and combining with the fact that ⟨∇W(λ∥W∥2,48),W⟩=8λ∥W∥2,48\left\langle\nabla_{\mathbf{W}}(\lambda\left\|{{\mathbf{W}}}\right\|_{2,4}^{8}),{\mathbf{W}}\right\rangle=8\lambda\left\|{{\mathbf{W}}}\right\|_{2,4}^{8}, we have simultaneously for all W{\mathbf{W}} that

Therefore we see that any stationary point W{\mathbf{W}} has to satisfy

By Corollary 3, choosing m≥poly(λ0−1,d,Bw,⋆Bx,ε)m\geq{\rm poly}(\lambda_{0}^{-1},d,B_{w,\star}B_{x},\varepsilon), the coupling error is bounded by ε\varepsilon in B2,4(Bw,0){\sf B}_{2,4}(B_{w,0}), i.e. for all W∈B2,4(Bw,0){\mathbf{W}}\in{\sf B}_{2,4}(B_{w,0}) we have that

we get that CλBw,⋆8=2γOPT+εC\lambda B_{w,\star}^{8}=2\gamma{\sf OPT}+\varepsilon, and thus the bound (18) reads

For the second-order stationary point W^\widehat{{\mathbf{W}}}, the gradient term vanishes and the Hessian term is non-negative, so we get

for any γ=O(1)\gamma=O(1). This is the desired result. ∎

Appendix C Proofs for Section 5

where the last step used the power mean (or Cauchy-Schwarz) inequality on {∥wr∥2}{\left\{\left\|{{\mathbf{w}}_{r}}\right\|_{2}\right\}}. ∎

C.2 Proof of Theorem 6

We first relate the generalization of LL to that of LQL^{Q} through

By Lemma 12, we have simultaneously for all W∈B2,4(Bw){\mathbf{W}}\in{\sf B}_{2,4}(B_{w}) that

These bounds hold for all W∈B2,4(Bw){\mathbf{W}}\in{\sf B}_{2,4}(B_{w}) so apply to W^\widehat{{\mathbf{W}}}. Therefore it remains to bound LPQ(W^)−LQ(W^)L^{Q}_{P}(\widehat{{\mathbf{W}}})-L^{Q}(\widehat{{\mathbf{W}}}), i.e. the generalization of the quadratic model.

By symmetrization and applying Lemma 5, we have

We now focus on bounding the expected max operator norm above. First, we apply the matrix concentration Lemma 8 to deduce that

As ∣σ′′(t)∣≤Ct|\sigma^{\prime\prime}(t)|\leq Ct and w0,r⊤xi∼N(0,1){\mathbf{w}}_{0,r}^{\top}{\mathbf{x}}_{i}\sim\mathsf{N}(0,1) for all (r,i)(r,i), by standard expected max bound on sub-exponential variables we have

and substituting the above bound into (21) yields that

Combining the bound with the coupling error (19) and (20), we arrive at the desired result.

For Mx,opM_{x,{\rm op}} we have two versions of bounds:

We always have ∥∑i≤nxixi⊤/n∥op≤Bx2\left\|{\sum_{i\leq n}{\mathbf{x}}_{i}{\mathbf{x}}_{i}^{\top}/n}\right\|_{\rm op}\leq B_{x}^{2}, and thus Mx,op≤1M_{x,{\rm op}}\leq 1.

and that κ(Cov(x))≤κ\kappa(\text{Cov}({\mathbf{x}}))\leq\kappa, then we have ∥Cov(x))∥op≤κBx2/d\left\|{\text{Cov}({\mathbf{x}}))}\right\|_{\rm op}\leq\kappa B_{x}^{2}/d. Applying (Vershynin, 2018, Theorem 4.7.1), we get Mx,op≤κ/dM_{x,{\rm op}}\leq\kappa/\sqrt{d} whenever n≥O(K4d)n\geq O(K^{4}d).

C.3 Expressive power of infinitely wide quadratic models

Our proof builds on reducing the problem from representing (β⊤x)p({\bm{\beta}}^{\top}{\mathbf{x}})^{p} via quadratic networks to representing (β⊤x)p−2({\bm{\beta}}^{\top}{\mathbf{x}})^{p-2} through a random feature model. More precisely, we consider choosing

where aa is a real-valued random scalar that can depend on w0{\mathbf{w}}_{0}, and β{\bm{\beta}} is the fixed coefficient vector in f⋆f_{\star}. With this choice, the quadratic network reduces to

Therefore, to let the above express f⋆(x)=α(β⊤x)pf_{\star}({\mathbf{x}})=\alpha({\bm{\beta}}^{\top}{\mathbf{x}})^{p}, it suffices to choose aa such that

for all x{\mathbf{x}}. By Lemma 9, there exists a=a(w0)a=a({\mathbf{w}}_{0}) satisfying (23) and such that

Using this aa in (22), the quadratic network induced by (w+,w−)({\mathbf{w}}_{+},{\mathbf{w}}_{-}) has the desired expressivity, and further satisfies the expected 4th power norm bound

C.4 Proof of Theorem 7

We begin by stating and proving the result for k=1k=1 in Appendix C.4.1, i.e. when f⋆=α(β⊤x)pf_{\star}=\alpha({\bm{\beta}}^{\top}{\mathbf{x}})^{p} is a single “one-directional” polynomial. The main theorem then follows as a straightforward extension of the k=1k=1 case, which we prove in Appendix C.4.2.

where we recall (w+(w0),w−(w0))=(a+(w0)β,a−(w0)β)({\mathbf{w}}_{+}({\mathbf{w}}_{0}),{\mathbf{w}}_{-}({\mathbf{w}}_{0}))=(\sqrt{a_{+}({\mathbf{w}}_{0})}{\bm{\beta}},\sqrt{a_{-}({\mathbf{w}}_{0})}{\bm{\beta}}). We then have

As f⋆(x)=α(β⊤x)pf_{\star}({\mathbf{x}})=\alpha({\bm{\beta}}^{\top}{\mathbf{x}})^{p}, Lemma 15 guarantees that the coefficient a(w0)a({\mathbf{w}}_{0}) involved above satisfies that

By Markov inequality, we have with probability at least 1−δ/21-\delta/2 that

Let fm(x)=1m∑r≤m/2σ′′(w0,r⊤x)a(w0,r)f_{m}({\mathbf{x}})=\frac{1}{m}\sum_{r\leq m/2}\sigma^{\prime\prime}({\mathbf{w}}_{0,r}^{\top}{\mathbf{x}})a({\mathbf{w}}_{0,r}). We now show the concentration of fmf_{m} to f⋆,p−2(x):=α(β⊤x)p−2f_{\star,p-2}({\mathbf{x}})\mathrel{\mathop{:}}=\alpha(\beta^{\top}{\mathbf{x}})^{p-2} over the dataset {x1,…,xn}{\left\{{\mathbf{x}}_{1},\dots,{\mathbf{x}}_{n}\right\}}. We perform a truncation argument: let RR be a large radius (to be chosen) satisfying

Applying Chebyshev inequality and a union bound, we get

For any ε>0\varepsilon>0, by substituting in t=εBx−2∥β∥2−2/2t=\varepsilon B_{x}^{-2}\left\|{{\bm{\beta}}}\right\|_{2}^{-2}/2, we see that

Next, for any x{\mathbf{x}} we have the bound

Combining (35) and (37), we see that with probability at least 1−δ1-\delta,

To satisfy the requirements for mm and RR in (36) and (34), we first set R=O~(d)R=\widetilde{O}(\sqrt{d}) (with sufficiently large log factor) to satisfy (36) by standard Gaussian norm concentration (cf. Appendix A.3), and by (34) it suffices to set mm as

C.4.2 Proof of main theorem

such that with probability at least 1−δ/k1-\delta/k we have

(Note we have slightly abused notation, so that now {fW⋆(j)Q}j∈[k]{\left\{f^{Q}_{{\mathbf{W}}_{\star}^{(j)}}\right\}}_{j\in[k]} use a disjoint set of initial weights (a0(j),W0(j))({\mathbf{a}}_{0}^{(j)},{\mathbf{W}}_{0}^{(j)}).) Concatenating all the (W⋆(j),a0(j),W0(j))({\mathbf{W}}_{\star}^{(j)},{\mathbf{a}}_{0}^{(j)},{\mathbf{W}}_{0}^{(j)}) and applying a union bound, we have the following: so long as the width

which by the 1-Lipschitzness of the loss implies that

Further, as W⋆{\mathbf{W}}_{\star} is the concatenation of {W⋆(j)}j∈[k]{\left\{{\mathbf{W}}_{\star}^{(j)}\right\}}_{j\in[k]}, we have the norm bound

Appendix D Existence, generalization, and expressivity of higher-order NTKs

Recall that for analytic σ\sigma we have the expansion

For an arbitrary W{\mathbf{W}} such that ∥w+,r∥2,∥w−,r∥2=om(1)\left\|{{\mathbf{w}}_{+,r}}\right\|_{2},\left\|{{\mathbf{w}}_{-,r}}\right\|_{2}=o_{m}(1), we expect that f(1)(x)f^{(1)}({\mathbf{x}}) is the dominating term in the expansion.

We now describe an approach to finding W{\mathbf{W}} so that

that is, the neural net is approximately the kk-th order NTK plus an error term that goes to zero as m→∞m\to\infty, thereby “escaping” the NTK regime. Our approach builds on the following randomization technique: let z+z_{+}, z−z_{-} be two random variables (distributions) such that

Set (w+,r,w−,r)=(z+,rw⋆,r,z−,rw⋆,r)({\mathbf{w}}_{+,r},{\mathbf{w}}_{-,r})=(z_{+,r}{\mathbf{w}}_{\star,r},z_{-,r}{\mathbf{w}}_{\star,r}), and take ∥w⋆,r∥2=O(m−1/2k)\left\|{{\mathbf{w}}_{\star,r}}\right\|_{2}=O(m^{-1/2k}), we have

Therefore, with high probability, all f(1),…,f(k−1)f^{(1)},\dots,f^{(k-1)} as well as the remainder term f−∑j≤kf(j)f-\sum_{j\leq k}f^{(j)} has order O(m−1/2k)O(m^{-1/2k}), and the kk-th order NTK f(k)f^{(k)} can express an O(1)O(1) function.

We now turn to studying the generalization and expressivity of the kk-th order NTK f(k)f^{(k)}, Throughout this subsection, we assume (for convenience) that

As we have seen in Section 6, we have f(k)=O(1)f^{(k)}=O(1) by choosing wr∼O(m−1/2k){\mathbf{w}}_{r}\sim O(m^{-1/2k}), therefore we restrict attention on such W{\mathbf{W}}’s by considering the constraint set {W:∥W∥2,2k2k≤Bw2k}\{{\mathbf{W}}:\left\|{{\mathbf{W}}}\right\|_{2,2k}^{2k}\leq B_{w}^{2k}\} for some Bw=Om(1)B_{w}=O_{m}(1).

This subsection establishes the following results for the kk-th order NTK.

We bound the generalization of f(k)f^{(k)} through the tensor operator norm of a certain kk-tensor involving the features (Lemma 17). Consequently, the generalization of the kk-th order NTK for ∥W∥2,2k≤Bw\left\|{{\mathbf{W}}}\right\|_{2,2k}\leq B_{w}, when the base distribution of x{\mathbf{x}} is uniform on the sphere, scales as

(Theorem 19). Compared with the distribution-free bound BxkBwk/nB_{x}^{k}B_{w}^{k}/\sqrt{n}, the leading term is better by a factor of min⁡{dk−1,n}\sqrt{\min{\left\{d^{k-1},n\right\}}}. In particular, when n≥dk−1n\geq d^{k-1}, the generalization is better by a factor of dk−1\sqrt{d^{k-1}} than the distribution-free bound.

For the polynomial f⋆(x)=α(β⊤x)pf_{\star}({\mathbf{x}})=\alpha({\bm{\beta}}^{\top}{\mathbf{x}})^{p} with p≥kp\geq k (and p−kp-k is even or one), when mm is sufficiently large, there exists a W⋆{\mathbf{W}}_{\star} expressing f⋆f_{\star} such that

(Theorem 20). Substituting into the generalization bound yields the following generalization error for learning f⋆f_{\star}:

In particular, the leading multiplicative factor is the same for all kk (including the linear NTK with k=1k=1), but the sample complexity is lower by a factor of dk−1d^{k-1} when n≥dk−1n\geq d^{k-1}. This shows systematically the benefit of higher-order NTKs when distributional assumptions are present.

The nuclear norm ∥⋅∥∗\left\|{\cdot}\right\|_{*} is defined as the dual norm of the operator norm:

Specifically, for any rank-one tensor u⊗k{\mathbf{u}}^{\otimes k}, we have

i.e. its nuclear norm equals its operator norm (and also the Frobenius norm).

D.2.1 Generalization

We begin by stating a generalization bound for f(k)f^{(k)}, which depends on the operator norm of a kk-th order tensor feature, generalizing Lemma 5.

where σi∼iidUnif{±1}\sigma_{i}\stackrel{{\scriptstyle\rm iid}}{{\sim}}{\rm Unif}{\left\{\pm 1\right\}} are Rademacher variables.

where the last step used the power mean (or Cauchy-Schwarz) inequality on {∥wr∥2}{\left\{\left\|{{\mathbf{w}}_{r}}\right\|_{2}\right\}}. ∎∎

It is straightforward to see that the expected tensor operator norm can be bounded as

Substituting the above bound into Lemma 17 directly leads to the following generalization bound for f(k)f^{(k)}:

The proof of Lemma 18 is deferred to Appendix D.3.

D.2.2 Expressivity

The proof of Theorem 20 is deferred to Appendix D.4.

D.3 Proof of Lemma 18

We now perform a truncation argument to upper bound the above probability. Let M>0M>0 be a truncation level to be determined, we have by the Bernstein inequality that

where the O~(1)⋅Bx2kd−k\widetilde{O}(1)\cdot B_{x}^{2k}d^{-k} comes from computing the variance of

using that xi{\mathbf{x}}_{i} are uniform on the sphere (see, e.g. (Ghorbani et al., 2019b, Proof of Lemma 4)); MBxkMB_{x}^{k} is the bound on the variable ZiZ_{i}, and the O~(1)\widetilde{O}(1) comes from the fact that ∥w0,r∥2≤O~(dBx−1)\left\|{{\mathbf{w}}_{0,r}}\right\|_{2}\leq\widetilde{O}(\sqrt{d}B_{x}^{-1}) with high probability. Now, choosing

It remains to bound ∫t=0∞pt\int_{t=0}^{\infty}p_{t} to give an expectation bound on the desired tensor operator norm. This follows by adding up the following three bounds:

For the main branch “nt2/O~(Bx2kd−k)nt^{2}/\widetilde{O}(B_{x}^{2k}d^{-k})” we have

This follows by integrating the “1” branch for t≤O~(Bx2kd−(k−1)/n)t\leq\widetilde{O}(\sqrt{B_{x}^{2k}d^{-(k-1)}/n}) (which yields the right hand side) and integrating the other branch otherwise (the integral being upper bounded by O~(Bx2kd−k/n)\widetilde{O}(\sqrt{B_{x}^{2k}d^{-k}/n}), dominated by the right hand side).

The branch “(nt/Bxk)1/2(nt/B_{x}^{k})^{1/2}” is taken only when

On the other hand, the inequality (nt/Bxk)1/2>O~(d)(nt/B_{x}^{k})^{1/2}>\widetilde{O}(d) happens when

which is implied by the preceding condition so long as k≥3k\geq 3. Therefore, when this branch is taken, the O~(d)\widetilde{O}(d) can already be absorbed into the main term, so the contribution of this branch can be bounded as

Putting together the above three bounds, we obtain

D.4 Proof of Theorem 20

Our proof is analogous to that of Theorem 16, in which we first look at the case of infinitely many neurons and then use concentration to carry the result onto finitely many neurons.

We first consider expressing f⋆(x)=α(β⊤x)pf_{\star}({\mathbf{x}})=\alpha({\bm{\beta}}^{\top}{\mathbf{x}})^{p} with infinite-neuron version of f(k)f^{(k)}, that is, we wish to find random variables (w+,w−)({\mathbf{w}}_{+},{\mathbf{w}}_{-}) such that

for some real-valued random scalar aa (that depends on w0{\mathbf{w}}_{0}), we have

Using this aa, the kk-th order NTK defined by (w+,w−)({\mathbf{w}}_{+},{\mathbf{w}}_{-}) expresses f⋆f_{\star} and further satisfies the bound

where we recall (w+(w0),w−(w0))=(a+(w0)1/kβ,a−(w0)1/kβ)({\mathbf{w}}_{+}({\mathbf{w}}_{0}),{\mathbf{w}}_{-}({\mathbf{w}}_{0}))=(a_{+}({\mathbf{w}}_{0})^{1/k}{\bm{\beta}},a_{-}({\mathbf{w}}_{0})^{1/k}{\bm{\beta}}). We then have

As f⋆(x)=α(β⊤x)pf_{\star}({\mathbf{x}})=\alpha({\bm{\beta}}^{\top}{\mathbf{x}})^{p}, (32) guarantees that the coefficient a(w0)a({\mathbf{w}}_{0}) involved above satisfies that

By Markov inequality, we have with probability at least 1−δ/21-\delta/2 that

Let fm(x)=1m∑r≤mσk(w0,r⊤x)a(w0,r)f_{m}({\mathbf{x}})=\frac{1}{m}\sum_{r\leq m}\sigma_{k}({\mathbf{w}}_{0,r}^{\top}{\mathbf{x}})a({\mathbf{w}}_{0,r}). We now show the concentration of fmf_{m} to f⋆,p−k(x):=α(β⊤x)p−kf_{\star,p-k}({\mathbf{x}})\mathrel{\mathop{:}}=\alpha(\beta^{\top}{\mathbf{x}})^{p-k} over the dataset {x1,…,xn}{\left\{{\mathbf{x}}_{1},\dots,{\mathbf{x}}_{n}\right\}}. We perform a truncation argument: let RR be a large radius (to be chosen) satisfying

Applying Chebyshev inequality and a union bound, we get

For any ε>0\varepsilon>0, by substituting in t=εBx−k∥β∥2−k/2t=\varepsilon B_{x}^{-k}\left\|{{\bm{\beta}}}\right\|_{2}^{-k}/2, we see that

Next, for any x{\mathbf{x}} we have the bound

Combining (35) and (37), we see that with probability at least 1−δ1-\delta,

To satisfy the requirements for mm and RR in (36) and (34), we first set R=O~(d)R=\widetilde{O}(\sqrt{d}) (with sufficiently large log factor) to satisfy (36) by standard Gaussian norm concentration (cf. Appendix A.3), and by (34) it suffices to set mm as