Learning Over-Parametrized Two-Layer ReLU Neural Networks beyond NTK

Yuanzhi Li, Tengyu Ma, Hongyang R. Zhang

Introduction

Gradient-based optimization methods are the method of choice for learning neural networks. However, it has been challenging to understand their working on non-convex functions. Prior works prove that stochastic gradient descent provably convergences to an approximate local optimum Ge et al. (2015); Sun et al. (2015); Lee et al. (2017); Kleinberg et al. (2018). Remarkably, for many highly complex neural net models, gradient-based methods can also find high-quality solutions Sun (2019) and interpretable features Zeiler and Fergus (2014).

Recent studies made the connection between training wide neural networks and Neural Tangent Kernels (NTK) Jacot et al. (2018); Arora et al. (2019b); Cao and Gu (2019); Du et al. (2018c). The main idea is that training neural networks with gradient descent with a particular initialization is equivalent to using kernel methods. However, the NTK approach has not yet provided a fully satisfactory theory for explaining the success of neural networks. Empirically, there seems to be a non-negligible gap between the test performance of neural networks trained by SGD and that of the NTK Arora et al. (2019a); Li et al. (2019b). Recent works have suspected that the gap stems from that the NTK approach has difficulty dealing with non-trivial explicit regularizers or does not sufficiently leverage the implicit regularization of the algorithm Wei et al. (2019); Chizat and Bach (2018b); Li et al. (2019a); HaoChen et al. (2020).

In this work, we provide a new convergence analysis of the gradient descent dynamic on an over-parametrized two-layer ReLU neural network. We prove that for learning a certain two-layer target network with orthonormal ground truth weights, gradient descent is provably more accurate than any kernel method that uses polynomially large feature maps.

where aia_{i} is in [1κd,κd][\frac{1}{\kappa d},\frac{\kappa}{d}] for an absolute constant κ≥1\kappa\geq 1 and satisfies ∑i∈[d]ai=1\sum_{i\in[d]}a_{i}=1, and {wi⋆}i=1d\{w_{i}^{\star}\}_{i=1}^{d} forms an orthonormal basis. Equation (1.1) can also be written as the sum of 2d2d neurons with ReLU activation:

Let Z={(xj,yj)}j=1N\mathcal{Z}=\{(x_{j},y_{j})\}_{j=1}^{N} be a training dataset of NN i.i.d. samples from the Gaussian distribution with identity covariance and yj=f⋆(xj)y_{j}=f^{\star}(x_{j}) for any 1≤j≤N1\leq j\leq N.

We learn the target network f⋆f^{\star} using an over-parametrized two-layer ReLU network with m≥2dm\geq 2d neurons W={wi}i=1mW=\{w_{i}\}_{i=1}^{m}, given by:

Note that we have re-parametrized the output layer with the norm of the corresponding neuron, so that we only have one set of parameters WW. This is without loss of generality for learning f⋆f^{\star} because when ai≥0a_{i}\geq 0, ai⋅ReLU⁡(wi⊤x)a_{i}\cdot\operatorname*{ReLU}(w_{i}^{\top}x) is equal to ∥wi′∥⋅ReLU⁡(wi′⊤x)\|w_{i}^{\prime}\|\cdot\operatorname*{ReLU}({{w_{i}^{\prime}}^{\top}x}) where wi′=ai/∥wi∥⋅wiw_{i}^{\prime}=\sqrt{a_{i}/\|w_{i}\|}\cdot w_{i}. Given a training dataset Z={(xi,yi)}i=1N\mathcal{Z}=\{(x_{i},y_{i})\}_{i=1}^{N}, we learn the target network by minimizing the following empirical loss:

Let L(W)L(W) denote the population loss given by the expectation of L^(W)\hat{L}(W) over Z\mathcal{Z}.

Algorithm. We focus on truncated gradient descent with random initialization. Algorithm 1 describes the procedure.An interesting feature is that when a neuron becomes larger than a certain threshold, we no longer update the neuron. This is a variant of gradient clipping often used in training recurrent neural networks (e.g. Merity et al. (2017); Gehring et al. (2017); Peters et al. (2018)) — here we drop the gradients of the large weights instead of re-scaling them. The truncation allows us to upper bound the norm of every neuron. Our main result is to show that Algorithm 1 learns the target network accurately in polynomially many iterations.

Let Z\mathcal{Z} be a training dataset with N=poly⁡κ(d)N=\operatorname{poly}_{\kappa}(d) samples generated by the model described above.Let poly⁡(d)\operatorname{poly}(d) denote a polynomial of dd and poly⁡κ(d)\operatorname{poly}_{\kappa}(d) denote a polynomial whose degree may depend on κ\kappa. Let C(κ)C(\kappa) be a sufficiently large constant that only depends on κ\kappa. Let 0<Q<1/1000<Q<1/100 be a sufficiently small absolute constant that does not depend on κ\kappa. Let λ0\lambda_{0} be a sufficiently small value on the order of 1/poly⁡(d)1/\operatorname{poly}(d) and λ1≤λ0/O(poly⁡κ(d))\lambda_{1}\leq\lambda_{0}/O(\operatorname{poly}_{\kappa}(d)) be a sufficiently small value on the order of 1/poly⁡κ(d)1/\operatorname{poly}_{\kappa}(d). For a learning rate η<min⁡(λ02,O(1poly⁡κ(d)))\eta<\min\left(\lambda_{0}^{2},O(\frac{1}{\operatorname{poly}_{\kappa}(d)})\right), a network width m≥Ω(poly⁡(d)/poly⁡(λ1))m\geq\Omega(\operatorname{poly}(d)/\operatorname{poly}(\lambda_{1})), and truncation parameters λ0,λ1\lambda_{0},\lambda_{1}, let W^\hat{W} be the final network learned by Algorithm 1. With probability 1−1poly⁡(d)1-\frac{1}{\operatorname{poly}(d)} over the choice of the random initialization, we have that the population loss of W^\hat{W} satisfies

The intuition behind our main result is as follows. We build on a connection between the popluation L(W)L(W) and tensor decomposition for Gaussian inputs Ge et al. (2017, 2018). By expanding the population loss in the Hermite polynomial basis, the optimization problem becomes an infinite sum of tensor decompositions problems (cf. equation (2) in Section 2) To analyze the gradient descent dynamic on the infinite sum tensor decomposition objective, we first analyze the infinite-width case – when mm goes to infinity. We establish a conditional-symmetry condition on the population of neurons, which greatly simplifies the analysis. This is established using the fact that our input distribution and labeling function (the absolute value activation) are both symmetric. Our analysis uncovers a stage-wise convergence of the gradient descent dynamic as follows, which matches our observations in simulations.

First, Algorithm 1 minimizes the 0th and 2nd order tensor decompositions. Informally, the distribution of neurons is fitting to the 0th moment and the 2nd moment of {wi⋆}i=1d\left\{w_{i}^{\star}\right\}_{i=1}^{d}.

Second, Algorithm 1 minimizes the 4th and higher order tensor decompositions. Initially, there is a long plateau where the evolution is slow, but after a certain point gets faster. As a remark, this behavior has been observed for randomly initialized tensor power method Anandkumar et al. (2017). Because the solution to the 4th and higher order orthogonal tensor decomposition problems is unique, we can learn the ground truth weights {wi⋆}i=1d\left\{w_{i}^{\star}\right\}_{i=1}^{d}.

Then we show that the sampling error between the infinite-width case and the finite-width case is small. The finite-width case can be thought of as a finite sample of the infinite-width case. As the network width increases, the sampling error reduces. In Section 3 and 4, we will first present a proof overview. The full proof is given in Section A and B.

As a complement, we show that the generalization error bound of Theorem 1.1 cannot be achieved by kernel functions with polynomially large feature map. Hence, by minimizing the higher order tensor decomposition terms, the learned neural network is provably more accurate than kernel functions that simply fit the lower order terms. Our result is stated as follows.

Under either of the following two situations,

Comparing the above result with Theorem 1.1, we conclude that provided with polynomially many samples, Algorithm 1 can recover the target two-layer neural network more accurately than the feature map and kernel method described above. Section C shows how to prove Theorem 1.2.

2 Related Work

Neural tangent kernel (NTK). A sequence of recent work shows that the learning process of gradient descent on over-parametrized neural networks, under certain initializations, reduces to the learning process of the associated neural tangent kernel. See Jacot et al. (2018); Arora et al. (2019b); Cao and Gu (2019); Du et al. (2018c); Arora et al. (2019a); Allen-Zhu and Li (2019b); Allen-Zhu et al. (2019c, b); Li and Liang (2018); Zou et al. (2018); Du et al. (2018a); Daniely et al. (2016); Ghorbani et al. (2019); Li et al. (2019a); Hanin and Nica (2019); Yang (2019) and the references therein. For NTK based results, the learning process of gradient descent can be viewed as solving convex kernel regression. Our work analyzes a non-convex objective that involves an infinite sum of tensor decomposition problems. By analyzing the higher order tensor decompositions, we can achieve a smaller generalization error than kernel methods.

Allen-Zhu and Li (2019a, 2020a) show that over-parametrized neural networks can learn certain concept class more efficient than any kernel method. Their work assumes the target network satisfies a certain “information gap” assumption between the first and second layer, while our target network does not require such gaps. Allen-Zhu et al. (2019a); Bai and Lee (2019) go beyond NTK by studying quadratic approximations of neural networks. Our work further analyzes higher-order tensor decompositions that are present in the Taylor expansion of the loss objective.

Two-layer neural networks given Gaussian inputs. There is a large body of work on learning two-layer neural networks over the last few years, such as Kawaguchi (2016); Soudry and Carmon (2016); Xie et al. (2016); Soltanolkotabi et al. (2017); Tian (2017); Brutzkus and Globerson (2017); Boob and Lan (2017); Vempala and Wilmes (2018); Oymak and Soltanolkotabi (2019); Bakshi et al. (2018); Yehudai and Shamir (2019); Zhang et al. (2018); Li and Liang (2017); Li and Dou (2020); Allen-Zhu and Li (2020b). Our work is particularly related to those that learn a two-layer neural network given Gaussian inputs. Li and Yuan (2017); Zhong et al. (2017) consider learning two-layer networks with ReLU activations with a warm start tensor initialization, as opposed to from a random initialization. Du et al. (2017) consider learning a target function consisting of a single ReLU activation. Brutzkus and Globerson (2017); Tian (2017) study the case where the weight vector for each neuron has disjoint support. Apart from the gradient descent algorithm, the method of moments has also been shown to be an effective strategy with provable guarantees (e.g. Bakshi et al. (2018); Ge et al. (2018)).

The closest work to ours is Ge et al. (2017) that consider a similar concept class. However, their work requires designing a complicated loss function, which is different from the mean squared loss. The learner network also uses a low-degree activation function as opposed to the ReLU activation. These are introduced to address the challenge of analyzing non-convex optimization for tensor decomposition with multiple components as variables, because prior works mostly focus on the non-convex formulation that optimizes over a single component (e.g., see Ge and Ma (2017)). Ge et al. (2017) have stated the question of analyzing the gradient descent dynamic for minimizing the sum of second and fourth order tensor decompositions as a challenging open question. Our analysis not only applies to this setting, but also allows for more even order tensor decompositions. Apart from ReLU activations, quadratic activations have been studied in Li et al. (2018); Oymak and Soltanolkotabi (2019); Soltanolkotabi et al. (2017).

Infinite-width neural networks. Previous work such as Mei et al. (2018); Chizat and Bach (2018a) show that as the hidden layer width goes to infinity, gradient descent approaches the Wasserstein gradient flow. Mei et al. (2018) use tools from partial differential equations to prove the global convergence of the gradient descent. Both of these results do not provide explicit convergence rates. Wei et al. (2018) show that under a certain regularity assumption on the activation function, the Wasserstein gradient flow converges in polynomial iterations for infinite-width neural networks..

Organizations. The rest of the paper is organized as follows. In Section 2, we reduce our setting to learning a sum of tensor decomposition problems. In Section 3, we describe an overview of the analysis for the infinite-width case. In Section 4, we show how to connect the above case to the gradient descent dynamic on the empirical loss for polynomially-wide networks. Finally we validate our theoretical insight on simulations in Section 5. In Section A, we provide the proof of the infinite-width case. In Section B, we provide an error analysis of the infinite-width case and complete the proof of Theorem 1.1. In Section C, we present the proof of Theorem 1.2.

Preliminaries

Recall that the ground-truth weights {wi∗}i=1d\left\{w_{i}^{*}\right\}_{i=1}^{d} forms an orthonormal basis. Since the input distribution x∼N(0,Id×d)x\sim\mathcal{N}(0,I_{d\times d}) and the initialization {wi}i=1m∼N(0,1d⋅Id⁡d×d)\left\{w_{i}\right\}_{i=1}^{m}\sim\mathcal{N}\left(0,\frac{1}{d}\cdot\operatorname{Id}_{d\times d}\right) are both rotation invariant, without loss of generality we can assume that wi∗=eiw_{i}^{*}=e_{i}, for all 1≤i≤d1\leq i\leq d.

We can average out the randomness in xx by applying Theorem 2.1 of Ge et al. (2017) on the loss function L(W)L(W), by expanding the activations function in the Hermite basis O’Donnell (2014).

where ck=2[(k−3)!!]2π⋅k!c_{k}=\frac{2[(k-3)!!]^{2}}{\pi\cdot k!} is the Hermite coefficients of the absolute value function for any k≥0k\geq 0. We remark that the population loss is a infinite sum of orthogonal tensor decomposition problems! For example, the -th order tensor decomposition concerns the l2l_{2}-norm of the weights. More generally, the kk-order tensor decomposition concerns the kk-th moment of the weights.

Correspondingly, the population loss of fP(x)f_{\mathcal{P}}(x) is given as

Gradient descent update. It has been shown in prior works that gradient descent in the (natural) parameter space corresponds to Wasserstein gradient descent in the distributional space. However, we found that the Wasserstein gradient perspective is not particularly helpful for us to analyze our algorithms and therefore we work with the update in the parameter space. The distribution P\mathcal{P} can be viewed as a collection of infinitesimal neurons. The gradient of each neuron vv is given by computing the gradient of the objective L(W)L(W) w.r.t a particle vv assuming the rest of the particles follow the distribution P\mathcal{P}. Let ∇vL∞(P)\nabla_{v}L_{\infty}(\mathcal{P}) denote the gradient of vv. We have that

where b0=4c0,b1=2c1b_{0}=4c_{0},b_{1}=2c_{1}, and for any j≥2j\geq 2, b2j=(4j)×c2j=Θ(1j2)b_{2j}=(4j)\times c_{2j}=\Theta\left(\frac{1}{j^{2}}\right) and b2j′=(4j−4)×c2jb_{2j^{\prime}}=(4j-4)\times c_{2j}. We use ∇vL∞\nabla_{v}L_{\infty} and ∇v\nabla_{v} as a shorthand for ∇vL∞(P)\nabla_{v}L_{\infty}(\mathcal{P}). Based on equation (2.5), we can further decompose ∇vL∞(P)\nabla_{v}L_{\infty}(\mathcal{P}) into the sum of ∇2j,vL∞(P)\nabla_{2j,v}L_{\infty}(\mathcal{P}) for j≥0j\geq 0, where the 2j2j-th gradient refers to the gradient of the 2j2j-th tensor decomposition. As a result, given a neural network with neuron distribution P(t)\mathcal{P}^{(t)}, the neuron distribution after a truncated gradient descent step, denoted by P(t+1)\mathcal{P}^{(t+1)}, satisfies that

Finite-width case. We briefly describe the connection between the above infinite-width case and the finite-width case. Intuitively, we can think of the finite-width case as sampling mm neurons randomly from the neuron population P\mathcal{P} in the infinite-width case. There are two sources of sampling error that arise from the above process: (i) the error of the gradients between the finite neuron distribution and the infinite neuron distribution; (ii) the error between the empirical loss and the population loss. Because of gradient truncation, the norm of every neuron is bounded by 1/λ1/\lambda. Therefore, the sampling error reduces as mm and NN increases, as shown in the following claim.

With probability at least 1−δ1-\delta over the randomness of {wi}i=1m\{w_{i}\}_{i=1}^{m} and the training dataset Z\mathcal{Z}, for every w∈Ww\in W, we have that:

Claim 2.1 can be proved by standard concentration inequalities such as the Chernoff bound.

Overview of the Infinite-Width Case

We begin by studying Algorithm 1 for minimizing the population loss using an infinite-width neural network. The infinite-width case plays a central role in our analysis. First, the infinite-width case allows us to simplify the gradient update rule through a conditional-symmetry condition that we describe below. Second, the finite-width case can be reduced to the infinite-width case by bounding the sampling error of the two cases — we describe the reduction in the next section.

A natural starting point for the infinite-width case is to simply set the network width mm to infinity in Theorem 1.1. However, this will include negligible outliers such as those with large norms in the Gaussian distribution. Therefore, we focus on a truncated probability measure P(0)\mathcal{P}^{(0)} of N(0,Id⁡d×d)\mathcal{N}(0,\operatorname{Id}_{d\times d}) by enforcing a certain bounded condition. The precise definition of P(0)\mathcal{P}^{(0)} is presented in Definition A.1 of Appendix A. For the purpose of providing an overview of the analysis, it suffices to think of P(0)\mathcal{P}^{(0)} as a Gaussian-like distribution that satisfies the following property.

Provided with P(0)\mathcal{P}^{(0)} as initialization, we are ready to state the main result of the infinite-width case as follows.

In the setting of Theorem 1.1, let the number of samples NN go to infinity. Starting from the initialization W(0)W^{(0)} as the neuron distribution P(0)\mathcal{P}^{(0)}, let W^\hat{W} be the final output network by Algorithm 1. The population loss of W^\hat{W} satisfies L(W^)≤O(1/d1+Q)L(\hat{W})\leq O(1/d^{1+Q}).

In the rest of this section, we present an overview of the proof of Theorem 3.1 and provide pointers to the proof details to be found in Section A. First, we provide a simplifying formula for the gradient of L∞(P)L_{\infty}(\mathcal{P}). We describe an overview of the two stages of Algorithm 1 in Section 3.1 and 3.2, respectively.

Suppose the update rule of P(t)\mathcal{P}^{(t)} is given in equation (2.6). If P(t)\mathcal{P}^{(t)} is conditionally-symmetric, then P(t+1)\mathcal{P}^{(t+1)} is also conditionally-symmetric.

To see that Claim 3.1 is true, we first observe that the 1st order tensor decomposition is always zero when P(t)\mathcal{P}^{(t)} is conditionally symmetric. For the even order tensor decompositions, we observe that for every neuron vv in P(t)\mathcal{P}^{(t)} and every 1≤j≤d1\leq j\leq d, subject to v−jv_{-j} being fixed, ∇vL∞(P)\nabla_{v}L_{\infty}(\mathcal{P}) is a polynomial of vjv_{j} that only involves odd degree monomials. Therefore, as long as P(t)\mathcal{P}^{(t)} is conditionally-symmetric, then P(t+1)\mathcal{P}^{(t+1)} is still conditionally-symmetric. Since P(0)\mathcal{P}^{(0)} is conditionally-symmetric by definition, we conclude that the neuron distribution is conditionally-symmetric throughout Algorithm 1. Based on this claim, we simplify equation (2.5) as follows.

Suppose that P=P(t)\mathcal{P}=\mathcal{P}^{(t)} is conditionally-symmetric. For any j≥0j\geq 0, let ∇2j,v\nabla_{2j,v} be a shorthand for the gradient of the 2j-th tensor ∇2j,vL∞(P)\nabla_{2j,v}L_{\infty}(\mathcal{P}). For any 1≤i≤d1\leq i\leq d, let [∇2j,v]i[\nabla_{2j,v}]_{i} be the ii-th coordinate of ∇2j,v\nabla_{2j,v}. We have that [∇2j,v]i[\nabla_{2j,v}]_{i} is equal to the following for each value of jj:

The proof of Claim 3.2 is by applying Claim 3.1 to equation (2.5), which zeroes out the coordinates in ww that has an odd order before taking the expectation of ww in P\mathcal{P}. For the 2nd order gradient [∇2,v]i[\nabla_{2,v}]_{i}, we have that

Similar arguments apply to the gradient of higher order tensor decompositions. Claim 3.1 and 3.2 together implies that for the infinite-width case, the gradient descent update is given by equation (3.2) and (3.3).

We show that Algorithm 1 minimizes the 0th and 2nd order tensor decompositions of the objective L∞L_{\infty} to zero first.

First, we show that the gradient of the 4th and higher order tensor decompositions is dominated by ∇0,v\nabla_{0,v} and ∇2,v\nabla_{2,v}. We observe that for v∼P(0)v\sim\mathcal{P}^{(0)}, the ii-th coordinate of ∇0,v\nabla_{0,v} and ∇2,v\nabla_{2,v} satisfies that

This is because P(0)\mathcal{P}^{(0)} is a suitable truncation of N(0,Id⁡d×d/d)\mathcal{N}(0,\operatorname{Id}_{d\times d}/d). We further have that

After the 0th and 2nd order tensor decompositions are minimized to a small enough value, the gradient of higher order tensor decompositions begins to dominate the update. In Lemma A.3, we show that for a small fraction of neurons, their norms become much larger than an average neuron — a phenomenon that we term as “winning the lottery ticket”. The main intuition is as follows.

In Proposition A.10, we show that the gradient of most neurons vv except a small fraction can be approximated by a signal term from the 4th order gradient plus an O(1/d2)O(1/d^{2}) error term:

where Ct(κ)C_{t}(\kappa) is a function that only depends on κ\kappa but grows slowly with tt. To see that equation (3.5) is true, except a small set of neurons with probability mass at most 1/dα1/d^{\alpha} where α\alpha will be specified later, any other neuron ww satisfies ∥w∥∞2≤αlog⁡d/d\|w\|_{\infty}^{2}\leq{\alpha\log d}/{d}. For the small set of neurons, since we stop updating a neuron when its norm grows larger than 1/λ01/\lambda_{0}, the norm of any of these neurons is less than 1/λ0{1}/{\lambda_{0}}. Thus, provided with a sufficiently large α\alpha, the contribution of these neurons to the gradient is negligible. Combined together, we prove equation (3.5) in Proposition A.10.

Next, we reduce the dynamic to tensor power method. Based on equation (3.5), we observe that the update of viv_{i} is approximately vi(t+1)≈vi(t)+η⋅ai⋅(vi(t))3v_{i}^{(t+1)}\approx v_{i}^{(t)}+\eta\cdot a_{i}\cdot(v_{i}^{(t)})^{3}, which is analogous to performing power method over a fourth order tensor decomposition problem. Hence, for larger initializations of viv_{i}, viv_{i} also grows faster. Based on the intuition, we introduce the set of “basis-like” neurons Si,good\mathcal{S}_{i,good} in the population P\mathcal{P}, which are defined more precisely in Lemma A.3. Intuitively, Si,good\mathcal{S}_{i,good} includes any neuron vv that satisfies [vi(0)]2≥C2log⁡d/d[v_{i}^{(0)}]^{2}\geq{C^{2}\log d}/{d}, which has probability measure at least 1/dC2{1}/{d^{C^{2}}} by standard anti-concentration inequalities. Following equation (3.5), we show that the neurons in Si,good\mathcal{S}_{i,good} keeps growing until they become roughly equal to ei/(λ0poly⁡(d))e_{i}/(\lambda_{0}\operatorname{poly}(d)).

As shown in Lemma A.3, Algorithm 1 goes through a long plateau of Oκ(d2/(ηpoly⁡log⁡(d))O_{\kappa}({d^{2}}/({\eta}\operatorname{poly}\log(d)) iterations, until the neurons of Si,good\mathcal{S}_{i,good} are sufficiently large. Intuitively, the scaling of d2d^{2} in the number of iterations arises from the 1/d21/d^{2} increment in equation (3.5). This concludes Stage 1. The update of these basis-like neurons will be the focus of Stage 2.

2 Dynamic during Stage 2

In the second stage, we reduce the gradient truncation parameter in Algorithm 1 from λ0=Θ(1/poly⁡(d))\lambda_{0}=\Theta(1/\operatorname{poly}(d)) to a smaller value λ1=Θ(1/poly⁡κ(d))\lambda_{1}=\Theta(1/\operatorname{poly}_{\kappa}(d)). This allows the neurons that are close to basis vectors to fit the target network more accurately.

In Lemma A.5, we show that after Θ(dlog⁡d/η)\Theta(d\log d/\eta) iterations, the population loss reduces to less than o(1/(dlog⁡0.01d))o(1/(d\log^{0.01}d)). The proof of Lemma A.5 involves analyzing the 0th and 2nd order tensor decompositions, similar to Stage 1.1.

At the end of Stage 2.1, the weights of the learner neural network form a “warm start” initialization, meaning that its population loss is less than o(1/d)o(1/d) Li and Yuan (2017); Zhong et al. (2017). The final substage will show that the population loss can be further reduced from o(1/(dlog⁡0.01d))o(1/(d\log^{0.01}d)) to O(1/d1+Q)O(1/d^{1+Q}), where QQ is a fixed constant defined in Theorem 1.1.

In Lemma A.6, we show that the population loss further reduces to O(1/d1+Q)O({1}/d^{1+Q}) after Θ(d1+10Q/η)\Theta(d^{1+10Q}/\eta) iterations. We describe an informal argument by contrasting the gradient update of neurons in Si,good\mathcal{S}_{i,good} and the rest of the neurons for a particular coordinate i∈[d]i\in[d].

For any neuron v∈Si,goodv\in\mathcal{S}_{i,good}, in Claim A.10, we show that the ii-th coordinate of vv approximately follows the following update (cf. equation (A.45)):

where ctc_{t} is a function that grows with tt but bounded above by O(dQ)O(d^{Q}) and C(κ)C(\kappa) is a function that only depends on κ\kappa. For any neuron v∉Si,goodv\notin\mathcal{S}_{i,good}, in Claim A.10, we show that viv_{i} follows a similar update but its corresponding value of ctc_{t} is much smaller than that of neurons in Si,good\mathcal{S}_{i,good}. Thus, basis-like neurons grow faster than the rest of neurons by an additive factor that scales with ct/d2c_{t}/d^{2}.

Once Lemma A.6 is finished, Algorithm 1 has learned an accurate approximation of f⋆(⋅)f^{\star}(\cdot) and we can conclude the proof of Theorem 3.1. We show that the population loss has also become less than O(d1+Q)O(d^{1+Q}) (cf. equation (A.10)). Thus, we have finished the analysis of Algorithm 1 for L∞(P)L_{\infty}(\mathcal{P}). We provide the proof details of Theorem 3.1 in Section A.

Overview of the Finite-Width Case

Based on the analysis of the infinite-width case, we reduce the finite-width case to the infinite-width case. By applying Claim 2.1 with P=P(t)\mathcal{P}=\mathcal{P}^{(t)}, when {wi(t)}i=1m\{w_{i}^{(t)}\}_{i=1}^{m} are i.i.d. samples from P(t)\mathcal{P}^{(t)}, the empirical loss and its gradient are tightly concentrated around the population loss and its gradient. Furthermore, as we increase the number of neurons mm and the number of samples NN, the sampling error reduces. Therefore, the goal of our reduction is to show that the sampling error remains small throughout the iterations of Algorithm 1. We describe our reduction informally and leave the details to Section B.

Combined together, we show in Lemma B.1 that ξw(t)\xi_{w}^{(t)} indeed remains polynomially small. For Stage 2, we analyze the propagation of ξw(t)\xi_{w}^{(t)} in Lemma B.2 and B.3 using similar arguments.

Combining the above three lemmas on error propagation and Theorem 3.1, we complete the proof of Theorem 1.1 in Section B.

Simulations

We provide simulations to complement our theoretical result. We consider a setting where wi⋆=eiw_{i}^{\star}=e_{i} and ai=1/da_{i}=1/d, for 1≤i≤d1\leq i\leq d. The input is drawn from the Gaussian distribution. For the iith order tensor, we measure the corresponding tensor decomposition loss from the population loss L(W)L(W).

We validate the insight of our analysis, which shows that the convergence of gradient descent has several stages. We use the labeling function of equation (1.1) and a learner network with absolute value activation functions as in Section 3 and Section A. First, the 0th and 2nd order tensor decomposition losses converge to zero quickly. Second, the 4th and higher order tensor decomposition losses converge to zero followed by a long plateau. Figure 2 shows the result. Here we use d=30d=30 and m=100>2dm=100>2d. The number of samples is 10410^{4}.

We can see that initially, the 0th and 2nd order tensor decompositions have higher loss than the 4th and higher order tensor decompositions. Then, both the 0th and the 2nd order losses decrease significantly from the initial value and converge to below 10−110^{-1} very quickly. Moreover, after a quick warm up period, the 0th order loss always stays smaller than the 2nd order loss, as our theory predicts. This is followed by a long plateau, which corresponds to Stage 1.2 of our analysis. During this stage, the 4th and higher order losses dominate dynamic, where a small fraction of neurons converge to basis-like neurons. Eventually, the learner neural network accumulates enough basis-like neurons from the 4th and higher tensors in the network. The 4th and higher order losses become less than 10−210^{-2}. The 0th and 2nd order losses further reduce to closer to zero. Our theory provides an in-depth explanation of these phenomena.

It has been observed that for properly parametrized gradient descent, gradient descent can get stuck starting from a random initialization Ge et al. (2017); Du et al. (2018b). We show that this is because the higher order losses remain large even though the 0th order loss has become small. We consider the same setting as the previous experiment but use m=2dm=2d. Figure 2 shows the result. We can see that the 0th order loss still reduces to less than 10−210^{-2}. However, the 2nd, 4th and 6th order losses are still larger than 10−110^{-1} even after 10510^{5} iterations.

Conclusions and Discussions

In this work, we have shown that for learning a certain target network with absolute value activation, a truncated gradient descent algorithm can provably converge in polynomially many iterations starting from a random initialization. The learned network is more accurate compared to any kernel method that uses polynomially large feature mappings.

We describe several interesting questions for future work. First, it would be interesting to extend our result to a setting where the target network uses ReLU activation, i.e. f⋆(x)=a⊤ReLU⁡(Wx)f^{\star}(x)=a^{\top}\operatorname*{ReLU}(Wx). We note that there is a straightforward reduction from the above setting to our setting by simply solving a linear regression. After applying the reduction, we could then apply our result. The challenge of directly analyzing gradient descent for learning f⋆(x)=a⊤ReLU⁡(Wx)f^{\star}(x)=a^{\top}\operatorname*{ReLU}(Wx) is that the 1st order tensor decomposition in the Hermite expansion of f⋆(x)f^{\star}(x) breaks the conditionally-symmetric property. Second, it would be interesting to extend our result to settings where W⋆W^{\star} is not necessarily orthonormal. The challenge is to analyze the gradient descent dynamic beyond orthogonal tensors. We leave this question for future research.

The work is in part supported by SDSI and SAIL. T. M is also supported in part by Lam Research and Google Faculty Award.

References

The appendix provides complete proofs to Theorem 1.1 and 1.2.

In Section A, we describe the proof of Theorem 3.1 for the infinite-width case. This section comprises the bulk of the appendix.

In Section B, we describe the proof of Theorem 1.1 by reducing the finite-width case to the infinite-width case.

In Section C, we prove Theorem 1.2 using ideas from the work of Allen-Zhu and Li [2019a].

Appendix A Proof of the Infinite-Width Case

We provide the proof of Theorem 3.1, which shows that running truncated gradient descent on an infinite-width network can recover the target network with population loss at most O(d1+Q)O(d^{1+Q}), where QQ is a sufficiently small constant defined in Theorem 3.1. Recall from Section 3 that our analysis begins by setting up the random initialization and then proceeds in two stages. We fill in the proof details left from Section 3. The rest of this section is organized as follows.

Initialization: We set up the random initialization used by Algorithm 1.

Stage 1: We fill in the proof details of the dynamic during Stage 1, which subsumes Stage 1.1 and Stage 1.2 described in Section 3.1. This stage runs for Θ(d2ηC(κ)log⁡d)\Theta(\frac{d^{2}}{\eta C(\kappa)\log d}) iterations.

Stage 2: We fill in the proof details of the dynamic during Stage 2, which subsumes Stage 2.1 and Stage 2.2 described in Section 3.2. This stage runs for Θ(d1+10Qη)\Theta(\frac{d^{1+10Q}}{\eta}) iterations.

Recall that for the infinite-width case, our initialization of the neuron distribution is a probability measure truncated from a Gaussian distribution with identity covariance. We formally define the truncation and the initialization, denoted by P(0)\mathcal{P}^{(0)}, as follows.

The maximum entry of ww is bounded: ∥w∥∞≤poly⁡log⁡(d)d\|w\|_{\infty}\leq\frac{\operatorname{poly}\log(d)}{\sqrt{d}}.

Both ∥w∥22\|w\|_{2}^{2} and ∑i=1daid⋅wi2\sum_{i=1}^{d}a_{i}d\cdot w_{i}^{2} are in the range

There are at most O(log⁡0.01(d))O(\log^{0.01}(d)) coordinates i∈[d]i\in[d] of ww such that wi2≥log⁡ddw_{i}^{2}\geq\frac{\log d}{d}.

We define P(0)\mathcal{P}^{(0)} as the probability measure of N(0,Id⁡d×d/d)\mathcal{N}(0,\operatorname{Id}_{d\times d}/d) conditional on the support set Sg\mathcal{S}_{g}.

Remark. For our purpose of proving the finite-width case later in Section B, it suffices to consider P(0)\mathcal{P}^{(0)} as the initialization as opposed to N(0,Id⁡d×d/d)\mathcal{N}(0,\operatorname{Id}_{d\times d}/d). This is because when Algorithm 1 samples m=poly⁡κ(d)m=\operatorname{poly}_{\kappa}(d) neurons from N(0,Id⁡d×d/d)\mathcal{N}(0,\operatorname{Id}_{d\times d}/d), with high probability all the mm samples are in the set Sg\mathcal{S}_{g}. To see this, by standard concentration inequalities for the Gaussian distribution, we can show that the set Sg\mathcal{S}_{g} has probability measure at least μ(Sg)≥1−1dΩ(1)\mu(\mathcal{S}_{g})\geq 1-\frac{1}{d^{\Omega(1)}}. Thus by union bound, with high probability all mm samples are in Sg\mathcal{S}_{g}.

As stated in Section 3, we are going to heavily use the conditionally-symmetric property (cf. Definition 3.1). We observe that the initialization P(0)\mathcal{P}^{(0)} is indeed conditionally-symmetric. This is because N(0,Id⁡d×d/d)\mathcal{N}\left(0,\operatorname{Id}_{d\times d}/d\right) satisfies the conditionally-symmetric property and our truncation in Definition A.1 only involves conditions on the square of the coordinates of ww. Hence the truncation of N(0,Id⁡d×d/d)\mathcal{N}(0,\operatorname{Id}_{d\times d}/d) to Sg\mathcal{S}_{g} preserves the conditionally-symmetric condition.

Notations for gradients. Before describing the analysis, we introduce several notations first. Recall from Claim 3.2 that the gradient of a neuron vv in the distribution P\mathcal{P} can be simplified given the conditionally-symmetric property. For each coordinate 1≤i≤d1\leq i\leq d, the gradient of neuron vv satisfies that [∇v]i=∑j≥0[∇2j,v]i[\nabla_{v}]_{i}=\sum_{j\geq 0}\left[\nabla_{2j,v}\right]_{i}, where ∇v=∇vL∞(P)\nabla_{v}=\nabla_{v}L_{\infty}(\mathcal{P}), ∇2j,v=∇2j,vL∞(P)\nabla_{2j,v}=\nabla_{2j,v}L_{\infty}(\mathcal{P}) denotes the gradient of vv for the 2j2j-th loss, and [∇v]i[\nabla_{v}]_{i} denotes the ii-th coordinate of ∇v\nabla_{v}. Let B1,2j=b2j+b2j′B_{1,2j}=b_{2j}+b_{2j}^{\prime} and B2,2j=b2j′B_{2,2j}=b_{2j}^{\prime}, where b2jb_{2j} and b2j′b_{2j}^{\prime} are the Hermite coefficients of the 2j2j-th loss given in Section 2. For a vector w∈Sgw\in\mathcal{S}_{g}, let w(0)w^{(0)} denote a neuron with initialization ww in the initialization P(0)\mathcal{P}^{(0)}. Let P(t)\mathcal{P}^{(t)} denote the tt-th iterate of P(0)\mathcal{P}^{(0)} following the update rule of equation (2.6).

Recall from Section 3.1 that the goal of Stage 1 is to show that a small fraction of neurons becomes basis-like, i.e. close to a basis eie_{i} times a scaling factor of poly⁡(d)\operatorname{poly}(d) at the end of Θκ(d2/ηlog⁡d)\Theta_{\kappa}(d^{2}/\eta\log d) iterations for some i∈[d]i\in[d]. To facilitate the analysis, we maintain an inductive hypothesis throughout Stage 1 that provides an upper bound on the norm of a typical neuron during the update. We first introduce the set of neurons that will not become basis-like by the end of Stage 1.

Let C0C_{0} be a large enough constant. Let c0=C0log⁡dc_{0}=C_{0}\log d and S\mathcal{S} be the set of all vectors ww in Sg\mathcal{S}_{g} such that

where wˉ=w/∥w∥\bar{w}=w/\|w\| denotes ww being normalized to norm 11.

Based on the above definition, we introduce the following inductive hypothesis that shows the neurons in S\mathcal{S} remain “small and dense” (i.e. not basis-like) throughout Stage 1. This stage runs for Θ(d2ηlog⁡d)\Theta(\frac{d^{2}}{\eta\log d}) iterations. We use κ1\kappa_{1} to denote a value that is less than O(exp⁡(poly⁡(κ)))O(\exp(\operatorname{poly}(\kappa))).

In the setting of Theorem 3.1, let T2=Θ(d2ηc0exp⁡(poly⁡(κ)))T_{2}=\Theta(\frac{d^{2}}{\eta c_{0}\exp(\operatorname{poly}(\kappa))}). There exists an increasing sequence {ct}t=1T2\left\{c_{t}\right\}_{t=1}^{T_{2}} where ct≤exp⁡(poly⁡(κ))log⁡dc_{t}\leq\exp(\operatorname{poly}(\kappa))\log d such that for every w∈Sw\in\mathcal{S} and every t≤T2t\leq T_{2}, the tt-th iterate of the neuron w(t)w^{(t)} with initialization w(0)=ww^{(0)}=w satisfies that

Furthermore, for every coordinate i∈[d]i\in[d], we have that in expectation,

Equation (A.2) and (A.3), which we also refer to as inductive hypothesis H1\mathcal{H}_{1}, show that the norm of any neuron in S\mathcal{S} will not grow beyond Oκ(log⁡d/d)O_{\kappa}(\log d/d). Hence they will not become basis-like during Stage 1.

The set S\mathcal{S} contains most neurons in P(0)\mathcal{P}^{(0)} because by standard anti-concentration inequalities, the measure of the set S\mathcal{S} is at least 1−d−O(C0)1-d^{-O(C_{0})}. Hence, 1−μ(S)1-\mu(\mathcal{S}) is at most d−O(C0)d^{-O(C_{0})}. Based on this fact, we state a simple claim on the norm of neurons that are not in S\mathcal{S} that will be used later:

To see that equation (A.4) is true, recall that the truncation of Algorithm 1 ensures that ∥w∥2≤1/λ0\|w\|^{2}\leq 1/\lambda_{0}. Combined with the fact that 1−μ(S)≤d−O(C)1-\mu(\mathcal{S})\leq d^{-O(C)} and λ0=Θ(1/poly⁡(d))\lambda_{0}=\Theta(1/\operatorname{poly}(d)), we have that equation (A.4) holds for a sufficiently large constant C0C_{0}. This finishes our introduction of the inductive hypothesis H1\mathcal{H}_{1}. The proof of Proposition A.1 can be found in Section A.2.2.

Given the inductive hypothesis H1\mathcal{H}_{1}, we can state the formal result that corresponds to Stage 1.1 in Section 3.1. For a neuron distribution P\mathcal{P}, let us first introduce the following notations, which corresponds to the population loss of the 0th and 2nd order tensor decompositions.

Based on the above notations, we show the following convergence result at the end of Stage 1.1.

In the setting of Theorem 3.1, suppose that Proposition A.1 holds. Let T1=Θ(poly⁡(κ1)dlog⁡dη)T_{1}=\Theta\left(\frac{\operatorname{poly}(\kappa_{1})d\log d}{\eta}\right). Then, for every t≥T1t\geq T_{1}, we have that Δ(t),δ+(t),δ−(t)\Delta^{(t)},\delta_{+}^{(t)},\delta_{-}^{(t)} are all less than ctpoly⁡(κ1)d2\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}, where ctc_{t} is given in Proposition A.1.

The above result implies that after T1T_{1} iterations, the 0th and 2nd order losses remain smaller than ctpoly⁡(κ1)/d2c_{t}\operatorname{poly}(\kappa_{1})/d^{2}. The proof of Lemma A.2 can be found in Section A.2.1.

Once Stage 1.1 is finished, recall from Section 3 that the higher order gradients begin to dominate the dynamic. Hence Algorithm 1 enters Stage 1.2. We introduce the following notations in order to state the formal result. Let T2′=T2−d2ηpoly⁡log⁡(d)T_{2}^{\prime}=T_{2}-\frac{d^{2}}{\eta\operatorname{poly}\log(d)}. For every 1≤i≤d1\leq i\leq d, let Γi=12B1,4(ai2d)(ηT2′)\Gamma_{i}=\frac{1}{2B_{1,4}(a_{i}^{2}d)(\eta T_{2}^{\prime})}. Let ρ=poly⁡(κ1)⋅log⁡dd\rho=\frac{\operatorname{poly}(\kappa_{1})\cdot\log d}{d}. Here, by our assumption, we know that ai2=Θ(1/d2)a_{i}^{2}=\Theta({1}/{d^{2}}). Since T2′=Θ(d2/(ηlog⁡d))T_{2}^{\prime}=\Theta({d^{2}}/(\eta\log d)), we can see that Γi=Θ(log⁡d/d)\Gamma_{i}=\Theta({\log d}/{d}). Consider a coordinate i∈[d]i\in[d]. We define the set of good neurons whose ii-th coordinate is larger than Γi+ρ\Gamma_{i}+\rho as

Then we define the set of bad neurons that have two large coordinates as

The following lemma shows that, among other statements, the neurons in Si,good\mathcal{S}_{i,good} will win the lottery and become basis-like at the end of Stage 1.2 in the sense described below.

In the setting of Theorem 3.1, suppose that Proposition A.1 holds. At iteration T2T_{2} (recall that T2T_{2} is defined in Proposition A.1), the following holds for Si,good\mathcal{S}_{i,good} and Si,bad\mathcal{S}_{i,bad}:

For every i∈[d]i\in[d] and every v∈Si,goodv\in\mathcal{S}_{i,good}, we have that

For every 1≤i≤d1\leq i\leq d and every v∈Sgv\in\mathcal{S}_{g}, if there exists j≠ij\neq i such that ∣vi(T2)∣|v_{i}^{(T_{2})}| and ∣vj(T2)∣|v_{j}^{(T_{2})}| are both greater than 2(log⁡d)2d\frac{2(\log d)^{2}}{\sqrt{d}}, then the neuron vv is in the union of Si,bad\mathcal{S}_{i,bad} and Sj,bad\mathcal{S}_{j,bad}.

For every i∈[d]i\in[d], the probability measure of Si,good\mathcal{S}_{i,good} and Si,bad\mathcal{S}_{i,bad} satisfies that

In the above result, the set Si,good\mathcal{S}_{i,good} contains neurons that become approximately a large scaling of the basis eie_{i} after T2T_{2} iterations, a phenomenon that we term as winning the lottery ticket. The norm of these neurons become much larger than those in S\mathcal{S}, whose norm is bounded by Oκ(log⁡d/d)O_{\kappa}(\log d/d). The set Si,bad\mathcal{S}_{i,bad} contains neurons whose coordinate ii might be large in the end, but not close to a basis. The final statement in this lemma shows that the probability measure of bad neurons is small compared to good neurons. Lemma A.3 is proved in Section A.3. This concludes Stage 1.

The second stage begins by reducing the gradient truncation parameter from λ0=Θ(1poly⁡(d))\lambda_{0}=\Theta(\frac{1}{\operatorname{poly}(d)}) to λ1=Θ(1poly⁡κ(d))\lambda_{1}=\Theta(\frac{1}{\operatorname{poly}_{\kappa}(d)}).As a remark, the rational for this technical twist is that the neurons do not grow too large Stage 1. This is useful for the error analysis later in the finite-width case. Recall from Section 3.2 that the goal of Stage 2 is to allow basis-like neurons to grow until they fit the target network with population loss at most O(d1+Q)O(d^{1+Q}).

The first substage of the analysis shows that the population loss reduces below O(1dlog⁡0.01d)O(\frac{1}{d\log^{0.01}d}), after T3=Θ(dlog⁡d/η)T_{3}=\Theta({d\log d}/{\eta}) many iterations.

The second substage of the analysis shows that the population loss further reduces below O(1/d1+Q)O\left({1}/{d^{1+Q}}\right), after T4=Θ(d1+10Q/η)T_{4}=\Theta({d^{1+10Q}}/{\eta}) many iterations.

To facilitate the analysis, we introduce an inductive hypothesis throughout Stage 2 that describes the behavior of the good and bad neurons. Let us introduce several notations first. Let the union of the bad neurons for all coordinates be given by

The set of potential neurons for coordinate i∈[d]i\in[d] is given by

We remark that these are the set of neurons whose coordinate ii can become larger than O(poly⁡log⁡(d)d)O(\frac{\operatorname{poly}\log(d)}{\sqrt{d}}) at the end of Stage 1 (cf. Section A.2.1). The set of good neurons Si,good\mathcal{S}_{i,good} is a subset of Si,pot\mathcal{S}_{i,pot}. Let the union of the potential neurons for all coordinates be given by

We maintain the following running hypothesis that, among other things, specifies the behavior of the potential, good, and bad neurons in detail.

In the setting of Theorem 3.1, there exists a monotonically increasing sequence ctt=T2T4\mathcal{c_{t}}_{t=T_{2}}^{T_{4}} such that cT2=poly⁡(log⁡d)≤ct≤dO(Q)≤d1/10c_{T_{2}}=\operatorname{poly}(\log d)\leq c_{t}\leq d^{O(Q)}\leq d^{1/10} and for every T2<t≤T4T_{2}<t\leq T_{4}, the following list of properties holds for the neuron distribution P(t)\mathcal{P}^{(t)}:

For every v∈Sgv\in\mathcal{S}_{g}, we have that ∥v(t)∥22≤1/λ1\|v^{(t)}\|_{2}^{2}\leq 1/\lambda_{1}. As a result, gradient truncation never happens during this stage.

For every v∉Spotv\notin\mathcal{S}_{pot}, we have that

For every i∈[d]i\in[d], every v∈Si,pot\Sbadv\in\mathcal{S}_{i,pot}\backslash\mathcal{S}_{bad}, and j≠ij\not=i, we have that

The probability mass of the set of bad neurons satisfies that

For every i∈[d]i\in[d] and every v∈Si,goodv\in\mathcal{S}_{i,good}, we have that ∥vi(t)∥22≥1λ0poly⁡(d)\|v_{i}^{(t)}\|_{2}^{2}\geq\frac{1}{\lambda_{0}\operatorname{poly}(d)}.

For every i∈[d]i\in[d], the following claims regarding the set of potential neurons and bad neurons hold:

where κ2\kappa_{2} denotes exp⁡(poly⁡(κ1))\exp(\operatorname{poly}(\kappa_{1})) and κ1\kappa_{1} denotes exp⁡(poly⁡(κ))\exp(\operatorname{poly}(\kappa)).

We remark that in the above inductive hypothesis, equation (A.5) and (A.6) show similar conditions as equation (A.2) provided in Proposition A.1. For the rest of the section, we refer to the conclusion of Proposition A.4 as inductive hypothesis H2\mathcal{H}_{2}. The proof of Proposition A.4 can be found in Section A.4.1.

Given the inductive hypothesis, we can state the formal result that corresponds to Stage 2.1 in Section 3.2. We introduce the notation Δ(t)=2b0(∑i=1d(γi(t)+βi(t))−∑i=1dai)\Delta^{(t)}=2b_{0}\left(\sum_{i=1}^{d}(\gamma_{i}^{(t)}+\beta_{i}^{(t)})-\sum_{i=1}^{d}a_{i}\right) that measures the average error of the neurons across all coordinates at iteration tt. We show that by the end of T3=T2+Θ(dlog⁡d/η)T_{3}=T_{2}+\Theta(d\log d/\eta) iterations, we have obtained a warm start neuron distribution for Δ(t)\Delta^{(t)}, {β1(t),…,βd(t)}\left\{\beta_{1}^{(t)},\dots,\beta_{d}^{(t)}\right\}, and {γ1(t),…,γd(t)}\left\{\gamma_{1}^{(t)},\dots,\gamma_{d}^{(t)}\right\}. We state the result below.

In the setting of Theorem 3.1, suppose Proposition A.4 holds. There exists an iteration T3=T2+Θ(dlog⁡d/η)T_{3}=T_{2}+\Theta(d\log d/\eta) such that at iteration T3T_{3}, the following holds:

The above result implies that the set of potential neurons has fit the ii-th coordinate of the target network with error less than o(1/d)o(1/d). The 0th order loss has also been reduced below o(1/d)o(1/d). The proof of Lemma A.5 can be found in Appendix A.3.

In the end, we describe the formal result that corresponds to Stage 2.2 in Section 3.2. We construct a potential function to show that βi(t)+γi(t)\beta_{i}^{(t)}+\gamma_{i}^{(t)} converges to aia_{i} when t≥T3t\geq T_{3}. After running for T4=T3+Θ(d1+10Qη)T_{4}=T_{3}+\Theta(\frac{d^{1+10Q}}{\eta}) many iterations, we show that a certain set of potential neurons has converged to aia_{i} with error at most O(1/d2+Q)O(1/d^{2+Q}), for every 1≤i≤d1\leq i\leq d.

The result is shown in Lemma A.6 below. We introduce the following notations for defining the potential function at iteration tt:

where C1,C2C_{1},C_{2} denote two sufficiently large constants. Consider the following functions (recall that Δ+\Delta_{+} and Δ−\Delta_{-} have been defined in Stage 1):

Let β+(t)=1Cmax⁡i∈[d]{βi(t)}\beta_{+}^{(t)}=\frac{1}{C}\max_{i\in[d]}\{\beta_{i}^{(t)}\}. Let Φ(t)=max⁡{Φ+(t),Φ−(t),β+(t)}\Phi^{(t)}=\max\{\Phi_{+}^{(t)},\Phi_{-}^{(t)},\beta_{+}^{(t)}\} be our potential function. Lemma A.5 implies that by the end of t=T3t=T_{3} iterations, we have that δ−(t),δ+(t),β+(t),Δ+(t),Δ−(t)\delta_{-}^{(t)},\delta_{+}^{(t)},\beta_{+}^{(t)},\Delta_{+}^{(t)},\Delta_{-}^{(t)} are all less than O(1/dlog⁡0.01d)O\left({1}/{d\log^{0.01}d}\right). Hence Φ(T3)≤O(1/(dlog⁡0.01d))\Phi^{(T_{3})}\leq O(1/(d\log^{0.01}d)). The result below shows that after iteration T3T_{3}, Φ(t)\Phi^{(t)} further decreases whenever Φ(t)\Phi^{(t)} is at least O(poly⁡(κ2)ct)d2)O(\frac{\operatorname{poly}(\kappa_{2})c_{t})}{d^{2}}).

In the setting of Theorem 3.1, suppose that Proposition A.4 holds. Let C1C_{1} be a fixed constant. For any T3<t≤T4T_{3}<t\leq T_{4}, as long as Φ(t)≥poly⁡(κ2)ctd2\Phi^{(t)}\geq\frac{\operatorname{poly}(\kappa_{2})c_{t}}{d^{2}} (recalling that ctc_{t} is defined in Proposition A.4) we have that

By combining the results of Stage 1 and Stage 2, we are ready to prove Theorem 3.1.

When Proposition A.1 and A.4 hold, using the induction hypothesis in equation (A.7), we have that for the infinite-width case, the population loss L∞(P(t))L_{\infty}(\mathcal{P}^{(t)}) satisfies:

where the first term comes from the 0th order loss and the second term comes from 2nd and higher order losses. This claim also implies that

At the beginning of Stage 2.2, by Lemma A.5, we know that Φ(T3)≤1/(dlog⁡0.01d)\Phi^{(T_{3})}\leq 1/(d\log^{0.01}d). During Stage 2.2, by Lemma A.6, as long as Φ(t)≥Oκ(ct/d2)\Phi^{(t)}\geq O_{\kappa}(c_{t}/d^{2}), Φt+1≤Φ(t)≤Φ(t)(1−O(Φ(t)))\Phi^{t+1}\leq\Phi^{(t)}\leq\Phi^{(t)}(1-O(\Phi^{(t)})). Hence, after at most d1+O(Q)/ηd^{1+O(Q)}/\eta iterations (or T4−T3T_{4}-T_{3} more precisely), Φ(T4)\Phi^{(T_{4})} reduces to below O(d1+Q)O(d^{1+Q}). Applying this result to equation (A.11), we conclude that L∞(P(T4))≤O(1/d1+Q)L_{\infty}(\mathcal{P}^{(T_{4})})\leq O(1/d^{1+Q}).

A.1 Stage 1.1: Proof of Convergence for 0th and 2nd Order Tensors

This section provides the proof of Lemma A.2 is organized as follows.

In Proposition A.7, we first show that the gradients from 4th and higher order tensor decompositions are small compared to that of the 0th and 2nd order tensor decompositions.

The above shows that the dynamic is mainly dominated by the 0th and 2nd losses initially. In Proposition A.8 and Proposition A.9, we show the gradient update of the 0th and 2nd order. Based on these, we show the proof Lemma A.2 at the end of this subsection.

We first show that the 4th and higher order tensor gradients do not have much contribution to the gradient, for all the neurons in S\mathcal{S}. We introduce the following notations for convenience. For a neuron distribution P\mathcal{P}, let the following denote the gradient of vv involving only other neurons ww.

Recall that ∇2j,v\nabla_{2j,v} is the gradient of vv for the 2j-th tensor (cf. equation (3.3)). Let

The following result provides an upper bound on the higher order gradients.

In the setting of Lemma A.2, suppose Proposition A.1 holds. Then there exists an absolute constant C>0C>0 such that for every i∈[d]i\in[d] and v∈Sgv\in\mathcal{S}_{g}, at the tt-th iteration for t≤T2t\leq T_{2}, the neuron vv from distribution P(t)\mathcal{P}^{(t)} satisfies that

Moreover, the gradient from the network satisfies

As a corollary, for every v∈S⊆Sgv\in\mathcal{S}\subseteq\mathcal{S}_{g}, we have that

Let us focus on ∇4,v\nabla_{4,v} first. We now bound each term in [∇4,v]i[\nabla_{4,v}]_{i} in equation (3.3) separately.

For the signal term in the gradient, ai⟨ei,v⟩⟨ei,vˉ⟩2a_{i}\langle e_{i},v\rangle\langle e_{i},\bar{v}\rangle^{2}, because ai≤κ/da_{i}\leq\kappa/d, we have

Another term in the gradient is (again, using the fact that for ww with w(0)∈Sw^{(0)}\in\mathcal{S}, ∥w∥∞2≤ctd\|w\|_{\infty}^{2}\leq\frac{c_{t}}{d}):

The last term in the gradient is given by:

Combining Eq (A.13), Eq (A.14), Eq (A.15) and Eq (A.16), we obtain that

For Δ2j,v\Delta_{2j,v}, with j≥3j\geq 3, we can apply the same calculation as above, and show that

Since ∇≥4,v=∑j≥2∇2j,v\nabla_{\geq 4,v}=\sum_{j\geq 2}\nabla_{2j,v} and ∑jb2j=O(1)\sum_{j}b_{2j}=O(1) we complete the proof. ∎

Based on the above result, we describe the dynamic of the 0th order tensor in the following proposition.

∣Δ(t)∣≤ctpoly⁡(κ1)d2|\Delta^{(t)}|\leq\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}.

If Δ(t)>0\Delta^{(t)}>0, then Δ(t)≤δ−(t)(1−116κ2d)\Delta^{(t)}\leq\delta_{-}^{(t)}\left(1-\frac{1}{16\kappa^{2}d}\right). If Δ(t)<0\Delta^{(t)}<0, then ∣Δ(t)∣≤(1−15κ15d)δ+(t)|\Delta^{(t)}|\leq\left(1-\frac{1}{5\kappa_{1}^{5}d}\right)\delta_{+}^{(t)}.

Moreover, when Δ(t)≥max⁡{8κ2δ+(t),ctpoly⁡(κ1)d2}\Delta^{(t)}\geq\max\left\{8\kappa^{2}\delta_{+}^{(t)},\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}\right\} , it holds that

where S+(t)\mathcal{S}^{(t)}_{+} is the set of all i∈[d]i\in[d] with δi(t)≥0\delta_{i}^{(t)}\geq 0 and ∣S+(t)∣\left|\mathcal{S}^{(t)}_{+}\right| denotes its cardinality.

Consider the iteration tt, we have that for every ii and every v∈Sgv\in\mathcal{S}_{g}, the update of vi(t)v_{i}^{(t)} is given as:

Hence, using Proposition A.7 and inequality (A.4), it holds that

This implies that for every sufficiently small η≤λ02\eta\leq\lambda_{0}^{2} and Λ≤λ02\Lambda\leq\lambda_{0}^{2}, it holds:

Let us consider two cases when abs⁡Δ(t)=Ω(κ18ct/d2)\operatorname*{abs}{\Delta^{(t)}}=\Omega(\kappa_{1}^{8}c_{t}/d^{2}).

Now, consider a value ρ=14κ2d\rho=\frac{1}{4\kappa^{2}d}, when Δ(t)≥δ−(t)(1−ρ2)\Delta^{(t)}\geq\delta_{-}^{(t)}(1-\frac{\rho}{2}), it holds that (1−ρ2)Δ(t)≥δ−(t)(1−ρ)(1-\frac{\rho}{2})\Delta^{(t)}\geq\delta_{-}^{(t)}(1-\rho). Therefore, when Δ(t)≥δ−(t)(1−ρ2)\Delta^{(t)}\geq\delta_{-}^{(t)}(1-\frac{\rho}{2}), Eq (A.21) implies that

Summing up all δi(t)\delta_{i}^{(t)} with δi(t)≤0\delta_{i}^{(t)}\leq 0, this implies that

Combine the above inequality with inequality (A.20), we have that (using Δ(t)≥0\Delta^{(t)}\geq 0 so that Δ+(t)≥Δ−(t)\Delta_{+}^{(t)}\geq\Delta_{-}^{(t)}):

Therefore we conclude that when Δ(t)≥Ω(ctκ8d2)\Delta^{(t)}\geq\Omega\left(\frac{c_{t}\kappa^{8}}{d^{2}}\right) and Δ(t)≥δ−(t)(1−ρ2)\Delta^{(t)}\geq\delta_{-}^{(t)}(1-\frac{\rho}{2}), it must holds that

Here the second inequality comes from Eq (A.23). This implies that Δ(t+1)\Delta^{(t+1)} will decrease faster than δ−(t+1)\delta_{-}^{(t+1)} at the next iteration. Hence, when Δ(t)≥Ω(ctκ8d2)\Delta^{(t)}\geq\Omega\left(\frac{c_{t}\kappa^{8}}{d^{2}}\right), then Δ(t)≥δ−(t)(1−ρ2)\Delta^{(t)}\geq\delta_{-}^{(t)}(1-\frac{\rho}{2}) can never happen. Hence, by our choice of ρ\rho, we conclude that as long as Δ(t)≥Ω(ctκ8d2)\Delta^{(t)}\geq\Omega\left(\frac{c_{t}\kappa^{8}}{d^{2}}\right), then

On the other hand, even when Δ(t)≤δ−(t)(1−ρ2)\Delta^{(t)}\leq\delta_{-}^{(t)}(1-\frac{\rho}{2}) but Δ(t)=Ω(κ18ctd2)\Delta^{(t)}=\Omega\left(\frac{\kappa_{1}^{8}c_{t}}{d^{2}}\right) , we still have that for every ii with δi(t)≤0\delta_{i}^{(t)}\leq 0, by Eq (A.22):

Hence, as long as Δ(t)=Ω(κ18ctd2)\Delta^{(t)}=\Omega\left(\frac{\kappa_{1}^{8}c_{t}}{d^{2}}\right), we will always have

Combining the above with equation (A.20), we have that

Here, we are using the fact that Δ−(t)≤Δ+(t)≤δ+(t)∣S+(t)∣\Delta^{(t)}_{-}\leq\Delta^{(t)}_{+}\leq\delta_{+}^{(t)}|\mathcal{S}^{(t)}_{+}|. Now, this implies that when Δ(t)≥max⁡{8κ2δ+(t),Ω(κ18ctd2)}\Delta^{(t)}\geq\max\left\{8\kappa^{2}\delta_{+}^{(t)},\Omega\left(\frac{\kappa_{1}^{8}c_{t}}{d^{2}}\right)\right\} , it also holds that

The proof follows by a similar argument to Case 1. ∎

Based on the above result, next we describe the dynamic of the 2nd order tensor.

Moreover, when Δ(t)>0\Delta^{(t)}>0, we have the following improved bound for δ+(t+1)\delta_{+}^{(t+1)}:

By the update rule, we can obtain (in Eq (A.19)) that

which proves the condition. On the other hand when δ−(t)≤ctpoly⁡(κ1)d2\delta_{-}^{(t)}\leq\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}, we directly completes the proof by choosing a larger poly in ctpoly⁡(κ1)d3\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{3}}. We can apply the same argument for δ+\delta_{+}, and the improved bound for the case when Δ(t)>0\Delta^{(t)}>0. ∎

A.1.1 Proof of the Main Lemma

Now we are ready to show the final convergence lemma. We first provide the following claim that shows on average, each coordinate of the neuron distribution lies in a bounded range. This also proves the first equation of (A.3) in the inductive hypothesis H1\mathcal{H}_{1}.

Hence, we only need to consider t∈Tt\in\mathcal{T}, for these iterations, by Proposition A.9 we know that

On the other hand by Proposition A.8, we have that when Δ(t)≥max⁡{8κ2δ+(t),ctpoly⁡(κ1)d2}\Delta^{(t)}\geq\max\left\{8\kappa^{2}\delta_{+}^{(t)},\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}\right\} , it holds that:

Now, let us define γ0=δ+0≤2κd\gamma_{0}=\delta_{+}^{0}\leq\frac{2\kappa}{d}, with γt+1=γt(1−η14κd)+ηctpoly⁡(κ1)d3\gamma_{t+1}=\gamma_{t}\left(1-\eta\frac{1}{4\kappa d}\right)+\eta\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{3}} for every t∈Tt\in\mathcal{T} and γt+1=γt+ηctpoly⁡(κ1)d3\gamma_{t+1}=\gamma_{t}+\eta\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{3}} otherwise. We know that as long as Δ(t)≥ctpoly⁡(κ1)d2\Delta^{(t)}\geq\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}, we have:

Clearly, by Proposition A.9 once δ+(t)\delta_{+}^{(t)} or δ−(t)≤ctpoly⁡(κ1)d2\delta_{-}^{(t)}\leq\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}, they will stay within the interval for the next iterations. By Proposition A.8, after both δ+(t)\delta_{+}^{(t)} and δ−(t)≤ctpoly⁡(κ1)d2\delta_{-}^{(t)}\leq\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}, we know Δ(t)\Delta^{(t)} will be within the interval as well.

Hence, we just need to consider the first time that δ+(t)\delta_{+}^{(t)} and δ−(t)\delta_{-}^{(t)} goes outside the interval. Following Proposition A.9, we know that when δ+(t)≥ctpoly⁡(κ1)d2\delta_{+}^{(t)}\geq\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}, it holds:

which gives the convergence error rate of δ+\delta_{+} after T1T_{1} iterations. The same holds for δ−(t)\delta_{-}^{(t)}. ∎

Finally, we have an estimate of how big each coordinate is for the neurons at the end of Stage 1.1, which can be given by the output layer weights {ai}i=1d\{a_{i}\}_{i=1}^{d}. We show the following claim, which will be used in the proof of Stage 2.

In the setting of Lemma A.2, at iteration T1T_{1} (recalling that T1=Θ(poly⁡(κ1)dlog⁡dη)T_{1}=\Theta(\frac{\operatorname{poly}(\kappa_{1})d\log d}{\eta})), for every v∈Sgv\in\mathcal{S}_{g} and every i∈[d]i\in[d], we have that

Let us first show the upper bound. For every v(0)∈Sgv^{(0)}\in\mathcal{S}_{g}. By the update rule, we have that

On the other hand, we have that by Eq (A.18), it holds:

A.2 Stage 1.2: Proof of Convergence for Higher Order Tensors

In this section, we prove Lemma A.3, which shows that by the end of Stage 1, a small fraction of neurons have won the lottery ticket by growing much larger than a typical neuron. This stage runs for approximately Θκ(d2ηlog⁡d)\Theta_{\kappa}(\frac{d^{2}}{\eta\log d}) many iterations (or T2−T1T_{2}-T_{1} more precisely). The proof of Lemma A.3 is organized as follows.

First, in Proposition A.10, we show that the dynamic is mainly determined by the 4th order gradients by bounding the gradients contributed by the 0th and 2nd order losses so that, as described in Section 3.1. Based on this result, we can relate the dynamic of this substage to tensor power method.

Second, we provide a lower bound on the norm of every neuron in Claim A.3. Based on this result, we prove Claim A.4 that shows the growth of good neurons. This leads to the proof of Lemma A.3 in Section A.2.1.

Finally, we prove the inductive hypothesis H1\mathcal{H}_{1} in Section A.2.2.

We describe the following proposition to bound the gradients of 4th or higher tensors.

In the setting of Lemma A.3, suppose that Proposition A.1 holds. Consider any iteration t∈[T1+1,T2]t\in[T_{1}+1,T_{2}] and any neuron v∈Sgv\in\mathcal{S}_{g}. Suppose that for every s≤ts\leq t, ∥v(s)∥∞≤poly⁡(log⁡d)d\|v^{(s)}\|_{\infty}\leq\frac{\operatorname{poly}(\log d)}{\sqrt{d}}. Then for every i∈[d]i\in[d], the gradient of vv at iteration tt satisfies

The result mainly follows from combining Proposition A.7 for the gradient coming from 4th and higher order losses with Proposition A.2 for the gradient of 0th and 2nd order losses. The only remaining term is

By the definition of Sg\mathcal{S}_{g}, we know that at T1T_{1} every v∈Sgv\in\mathcal{S}_{g} satisfies

We will maintain the following condition by induction.

Now suppose the following is true at some iteration t≥T1t\geq T_{1}, then we have that

Thus, for iteration t+1t+1, using Eq (A.28) we know that

Hence, we have that for every ii with ∣vi(t)∣2≤poly⁡(κ1)ctd|v^{(t)}_{i}|^{2}\leq\frac{\operatorname{poly}(\kappa_{1})c_{t}}{d}, it holds that

Hence for every t≤T2t\leq T_{2}, as long as ∣vi(T1)∣2≤poly⁡(κ1)ctd|v^{(T_{1})}_{i}|^{2}\leq\frac{\operatorname{poly}(\kappa_{1})c_{t}}{d}, we have:

This proves inequality (A.29) for t+1t+1. ∎

Next, we use the following claim to maintain a lower bound on the norm of each neuron.

In the setting of Lemma A.3, suppose Proposition A.1 holds. For every v∈Sgv\in\mathcal{S}_{g}, the norm of vv at any iteration t∈[T1+1,T2]t\in[T_{1}+1,T_{2}] satisfies ∥v(t)∥22≥Ω(1κ)\|v^{(t)}\|_{2}^{2}\geq\Omega\left(\frac{1}{\kappa}\right).

By the update rule, using Proposition A.7 we know that for every p∈[d]p\in[d]:

Combined with Proposition A.10 , we have that for every neuron vv, ∥v(t)∥22=Ω(1κ)\|v^{(t)}\|_{2}^{2}=\Omega\left(\frac{1}{\kappa}\right) for every t∈[T1,T2]t\in[T_{1},T_{2}]. ∎

Provided with the gradient bound and norm lower bound, we are now ready to prove the main result of Stage 1.2. Towards showing Lemma A.3, we prove the following claim, which shows that if a neuron has grown beyond poly⁡log⁡(d)d\frac{\operatorname{poly}\log(d)}{d} at a certain iteration T2′T_{2}^{\prime}, then this neuron will become basis-like at iteartion T2T_{2}.

In the setting of Lemma A.3, suppose that Proposition A.1 holds. For every v∈Sgv\in\mathcal{S}_{g}, suppose at iteration T2′T_{2}^{\prime} (recalling that T2′=T2−d2ηpoly⁡log⁡(d)T_{2}^{\prime}=T_{2}-\frac{d^{2}}{\eta\operatorname{poly}\log(d)}), only one coordinate i∈[d]i\in[d] satisfies ∣vi(T2′)∣≥log⁡10dd|v_{i}^{(T_{2}^{\prime})}|\geq\frac{\log^{10}d}{\sqrt{d}} and all the other coordinates satisfies ∣vj(T2′)∣≤(log⁡d)2d|v_{j}^{(T_{2}^{\prime})}|\leq\frac{(\log d)^{2}}{\sqrt{d}}, then at iteration T2T_{2}, we have that

In other words, the claim says that for neuron vv, its ii-th coordinate at iteration T2T_{2}, denoted by ∣vi(T2)∣|v_{i}^{(T_{2})}|, will be as large as poly⁡(d)\operatorname{poly}(d), which implies that this neuron has won the lottery. We describe the proof of Claim A.4.

We shall prove the claim by doing an induction. Consider the condition ∣vi(t)∣≥poly⁡(log⁡d)d|v_{i}^{(t)}|\geq\frac{\operatorname{poly}(\log d)}{\sqrt{d}} and all the other coordinates satisfies ∣vj(t)∣≤2(log⁡d)2d|v_{j}^{(t)}|\leq\frac{2(\log d)^{2}}{\sqrt{d}} for t∈[T2′,T]t\in[T_{2}^{\prime},T]. Suppose it is true up to iteration tt, consider iteration t+1t+1. When p=ip=i, we have that ai(vi(t))2j−2≥ar(vr(t))2j−2a_{i}(v_{i}^{(t)})^{2j-2}\geq a_{r}(v_{r}^{(t)})^{2j-2} for every r≠ir\not=i, hence this implies that (using the fact that B1,2j>B2,2jB_{1,2j}>B_{2,2j} and B1,4B_{1,4} is greater than B2,4B_{2,4} plus a fixed constant):

where Qi(t)Q_{i}^{(t)} is defined in the proof of Claim A.3. With equation (A.33), this implies that

which provides a direct the lower bound on (vi(t+1))2(v^{(t+1)}_{i})^{2}. Now, to show the upper bound of the other coordinates, recall that we have shown ∥v(t)∥22=Ω(1κ)\|v^{(t)}\|_{2}^{2}=\Omega\left(\frac{1}{\kappa}\right) for every t∈[T1,T2]t\in[T_{1},T_{2}],

Hence we prove all the other p≠ip\not=i satisfies ∣vp(t)∣≤2(log⁡d)2d|v_{p}^{(t)}|\leq\frac{2(\log d)^{2}}{\sqrt{d}} as long as T2−T2′≤d2ηlog⁡9(d)T_{2}-T_{2}^{\prime}\leq\frac{d^{2}}{\eta\log^{9}(d)}, which complete the induction. In the end, since ∣vi(t)∣≥poly⁡(log⁡d)d|v_{i}^{(t)}|\geq\frac{\operatorname{poly}(\log d)}{\sqrt{d}} and all the other coordinates satisfies ∣vj(t)∣≤2(log⁡d)2d|v_{j}^{(t)}|\leq\frac{2(\log d)^{2}}{\sqrt{d}} for every t∈[T2′,T]t\in[T_{2}^{\prime},T], we can further simplify Eq (A.36) as:

which directly gives us the bound ∣vi(T2)∣2=Ω(1λ0poly⁡(d))|v_{i}^{(T_{2})}|^{2}=\Omega\left(\frac{1}{\lambda_{0}\operatorname{poly}(d)}\right) at iteration T2T_{2}. ∎

Now we are ready to prove Lemma A.3. We define the union of good neurons as

where we recall that Γi\Gamma_{i} and ρ\rho have been defined before the statement of Lemma A.3. In the proof, we focus on the dynamic of a neuron vv until the point that ∥v∥∞≥poly⁡(log⁡d)d\|v\|_{\infty}\geq\frac{\operatorname{poly}(\log d)}{\sqrt{d}}. The key step is to track the dynamic via a tensor gradient update.

We focus on proving the following three statements.

For every v∉Spotv\notin\mathcal{S}_{pot}, ∥v(t)∥∞≥poly⁡(log⁡d)d\|v^{(t)}\|_{\infty}\geq\frac{\operatorname{poly}(\log d)}{\sqrt{d}} never happen for any t≤T2t\leq T_{2}.

For every v∈Sgoodv\in\mathcal{S}_{good}, ∥v(t)∥∞≥poly⁡(log⁡d)d\|v^{(t)}\|_{\infty}\geq\frac{\operatorname{poly}(\log d)}{\sqrt{d}} must happen for some t≤T2′t\leq T_{2}^{\prime} and when it happens, the condition in Claim A.4 meets for i=arg max⁡j∈[d]{vj(0)}i=\operatorname*{arg\,max}_{j\in[d]}\{v^{(0)}_{j}\}.

For every v∈Spot\Sbadv\in\mathcal{S}_{pot}\backslash\mathcal{S}_{bad}, ∥v(t)∥∞≥poly⁡(log⁡d)d\|v^{(t)}\|_{\infty}\geq\frac{\operatorname{poly}(\log d)}{\sqrt{d}} might happen for some t≤T2′t\leq T_{2}^{\prime}. If ∥v(t)∥∞≥poly⁡(log⁡d)d\|v^{(t)}\|_{\infty}\geq\frac{\operatorname{poly}(\log d)}{\sqrt{d}} happens for some t≤T2t\leq T_{2}, then the condition in Claim A.4 meets for i=arg max⁡j∈[d]{vj(0)}i=\operatorname*{arg\,max}_{j\in[d]}\{v^{(0)}_{j}\}.

The first and second statement of Lemma A.3 follow by combining the above three statements and Claim A.4. The third statement can be proved by standard anti-concentration inequalities for the Gaussian distribution. For the rest of the proof, we focus on proving the above three statements. We know by Proposition A.10 that when ∥v(t)∥∞≤poly⁡(log⁡d)d\|v^{(t)}\|_{\infty}\leq\frac{\operatorname{poly}(\log d)}{\sqrt{d}} the update of v(t)v^{(t)} at every iteration t∈[T1,T2]t\in[T_{1},T_{2}] is given by

For every ii, consider a process where p(T1),q(T1)=vi(T1)p^{(T_{1})},q^{(T_{1})}=v_{i}^{(T_{1})}, with

Along with Eq (A.38), we can see that for every tt where ∥v(t)∥∞≤poly⁡(log⁡d)d\|v^{(t)}\|_{\infty}\leq\frac{\operatorname{poly}(\log d)}{\sqrt{d}},

To analyze this process, we introduce the following differential equation

The solution is given as x2(t)=11τ2−2τ1tx^{2}(t)=\frac{1}{\frac{1}{\tau_{2}}-2\tau_{1}t}. Therefore, we can easily obtain that as long as ρ=Ω(ctpoly⁡(κ1)d2)\rho=\Omega\left(\frac{c_{t}\operatorname{poly}(\kappa_{1})}{d^{2}}\right), when τ1=B1,4ai\tau_{1}=B_{1,4}a_{i}, η2τ1T2′=1τ2\eta 2\tau_{1}T_{2}^{\prime}=\frac{1}{\tau_{2}} which implies that τ2=1η2τ1T2′=1η2(b4+b4′)aiT2′\tau_{2}=\frac{1}{\eta 2\tau_{1}T_{2}^{\prime}}=\frac{1}{\eta 2(b_{4}+b_{4}^{\prime})a_{i}T_{2}^{\prime}}, we have that

In the end, by Proposition A.1 and the definition of Sg\mathcal{S}_{g} (Eq (A.1)), we know that for every v∈Sgv\in\mathcal{S}_{g} and every i∈[d]i\in[d], we have that

Putting into the definition of τ2\tau_{2} we complete the proof. ∎

In addition, we state the following claim that will be used in Appendix B for the error analysis.

In the setting of Theorem 3.1, at the first iteration tt where ∥v(t)∥22>1λ0\|v^{(t)}\|_{2}^{2}>\frac{1}{\lambda_{0}}, i.e. the threshold where gradients are truncated, we have that

When ∥vˉ(t)∥∞2≤1poly⁡(κ1)\|\bar{v}^{(t)}\|_{\infty}^{2}\leq\frac{1}{\operatorname{poly}(\kappa_{1})}, we have that for p=arg max⁡r∈[d]{ar(vr(t))2}p=\operatorname*{arg\,max}_{r\in[d]}\{a_{r}(v_{r}^{(t)})^{2}\}, the following holds

The above implies that as long as ∥vˉ(t)∥∞2≤1poly⁡(κ1)\|\bar{v}^{(t)}\|_{\infty}^{2}\leq\frac{1}{\operatorname{poly}(\kappa_{1})}, we have:

After that, when ∥vˉ(t)∥∞2≥1poly⁡(κ1)\|\bar{v}^{(t)}\|_{\infty}^{2}\geq\frac{1}{\operatorname{poly}(\kappa_{1})}, we have that ∥vˉ(t)∥∞4≥1poly⁡(κ1)\|\bar{v}^{(t)}\|_{\infty}^{4}\geq\frac{1}{\operatorname{poly}(\kappa_{1})} as well, which implies

Hence, as long as ∥vˉ(t)∥∞2≥1poly⁡(κ1)\|\bar{v}^{(t)}\|_{\infty}^{2}\geq\frac{1}{\operatorname{poly}(\kappa_{1})}, Eq (A.33) implies that

On the other hand, we also have for every iteration, by Eq (A.34):

The above implies that (∗)(*) can only happen for poly⁡(κ1)dηlog⁡1λ0\frac{\operatorname{poly}(\kappa_{1})d}{\eta}\log\frac{1}{\lambda_{0}} iterations until the norm of vv is too large and gradient clipping happens. For these iterations when ∥vˉ(t)∥∞2≥1poly⁡(κ1)\|\bar{v}^{(t)}\|_{\infty}^{2}\geq\frac{1}{\operatorname{poly}(\kappa_{1})}, we can also easily see that

For all the other iterations when ∥vˉ(t)∥∞2≤1poly⁡(κ1)\|\bar{v}^{(t)}\|_{\infty}^{2}\leq\frac{1}{\operatorname{poly}(\kappa_{1})}, we have Eq (A.37) holds, which implies that as long as ∥v(t−1)∥22≤12λ0\|v^{(t-1)}\|_{2}^{2}\leq\frac{1}{2\lambda_{0}}:

A.2.2 Proof of the Inductive Hypothesis

Note that the first part of equation (A.3) has been shown in Claim A.1 — the second part can be shown via a similar proof of Claim A.1. For the rest of the proof, we focus on proving equation (A.2). The construction of the sequence {ct}t=1T2\left\{c_{t}\right\}_{t=1}^{T_{2}} will be shown below.

By inequality (A.27) in the proof of Claim A.1, we know that for every v(0)∈Sgv^{(0)}\in\mathcal{S}_{g} and t≤T1t\leq T_{1}, it holds that

which implies that for every t∈[T1]t\in[T_{1}], ct≤2κ1κc0c_{t}\leq 2\kappa_{1}\kappa c_{0}. Now, we focus on t∈[T1,T2]t\in[T_{1},T_{2}]. By Lemma A.2 ,we know that for every t≥T1t\geq T_{1} , we have that

By Proposition A.7, we have that for every v(0)∈Sv^{(0)}\in\mathcal{S}, ∣[∇≥4,v(t)]i∣≤C2ctκd2∣vi(t)∣\left|\left[\nabla_{\geq 4,v^{(t)}}\right]_{i}\right|\leq\frac{C}{2}\frac{c_{t}\kappa}{d^{2}}|v_{i}^{(t)}|. Hence,

Iterating the above equation over tt gives us the sequence {ct}t=1T2\left\{c_{t}\right\}_{t=1}^{T_{2}}. By maintaining that for every v∈Sv\in\mathcal{S}, the norm of vv at iteration tt satisfies ct≤poly⁡(κ1)c0c_{t}\leq\operatorname{poly}(\kappa_{1})c_{0} and the fact that T2≤d2ηc0poly⁡(κ1)T_{2}\leq\frac{d^{2}}{\eta c_{0}\operatorname{poly}(\kappa_{1})}, we have verified the running hypothesis H1\mathcal{H}_{1}. ∎

A.3 Stage 2.1: Obtaining a Warm Start Initialization

At the beginning of Stage 2, we reduce the gradient truncation parameter. This allows the basis-like neurons to continue to grow and we can obtain a warm start initialization at the end of Stage 2.1 in the sense described in Lemma A.5. The proof of Lemma A.5 consists of the following steps.

First, we analyze the 0th order loss in Claim A.6 and A.8.

Second, We analyze the 2nd order loss in Proposition A.11. Combined together, we prove Lemma A.5 in Section A.3.1.

Notations for gradients. To facilitate the analysis, we introduce several notations on the gradients of a neuron vv. We separate the gradient of vv into several components at the tt-th iteration as ∇v,2j=∇v,2j,sig+∇v,2j,¬pot+∇v,2j,bad+∇v,2j,pot\bad\nabla_{v,2j}=\nabla_{v,2j,sig}+\nabla_{v,2j,\neg pot}+\nabla_{v,2j,bad}+\nabla_{v,2j,pot\backslash bad}, where each term is given by

Recall that this substage runs for T3≤dlog⁡1.01dηT_{3}\leq\frac{d\log^{1.01}d}{\eta} iterations. We first focus on the update of the 0th order term Δ(t)\Delta^{(t)}. Let κ2\kappa_{2} denote epoly⁡(κ1)e^{\operatorname{poly}(\kappa_{1})}. We show the following claim.

In the setting of Lemma A.5, suppose that Proposition A.4 holds. Let δ\delta be any value in the range [poly⁡(κ2)ctd2,1κd][\frac{\operatorname{poly}(\kappa_{2})c_{t}}{d^{2}},\frac{1}{\kappa d}]. When Δ(t)≥δ\Delta^{(t)}\geq\delta, for any iteration t∈[T2+1,T3]t\in[T_{2}+1,T_{3}], we have that

Let us denote δ′=min⁡{δ,max⁡{C1,C2}10κd}\delta^{\prime}=\min\{\delta,\frac{\max\{C_{1},C_{2}\}}{10\kappa d}\}. We shall see that when Δ(t)≥δ\Delta^{(t)}\geq\delta, then for every ii with βi(t)+γi(t)≥ai−δ′4max⁡{C1,C2}\beta_{i}^{(t)}+\gamma_{i}^{(t)}\geq a_{i}-\frac{\delta^{\prime}}{4\max\{C_{1},C_{2}\}}, we have that

Therefore, using equation A.44, we have that

On the other hand, when γi(t)≥ai−δ′3max⁡{C1,C2}\gamma_{i}^{(t)}\geq a_{i}-\frac{\delta^{\prime}}{3\max\{C_{1},C_{2}\}}, we have that

In either case, we have that as long as βi(t)+γi(t)≥ai−δ′4max⁡{C1,C2}\beta_{i}^{(t)}+\gamma_{i}^{(t)}\geq a_{i}-\frac{\delta^{\prime}}{4\max\{C_{1},C_{2}\}}, it holds that

Using Δ(t)≥0\Delta^{(t)}\geq 0, we obtain that

Next, we focus on the other side when Δ(t)\Delta^{(t)} is negative. We first show the first lower bound on the neuron mass.

In the setting of Lemma A.5, suppose that Proposition A.4 holds. Then we have that for any t∈[T2+1,T3]t\in[T_{2}+1,T_{3}], the following holds:

Initially at t=0t=0, we have that L∞(P(0))=O(1d)L_{\infty}(\mathcal{P}^{(0)})=O\left(\frac{1}{d}\right)). Now, for every δ≤min⁡{C1,1}100κd\delta\leq\frac{\min\{C_{1},1\}}{100\kappa d}, when Δ(t)≤δ\Delta^{(t)}\leq\delta, we know that as long as βi(t)+γi(t)≤ai−2δC1\beta_{i}^{(t)}+\gamma_{i}^{(t)}\leq a_{i}-\frac{2\delta}{C_{1}} and βi(t)+γi(t)≥1poly⁡(d)\beta_{i}^{(t)}+\gamma_{i}^{(t)}\geq\frac{1}{\operatorname{poly}(d)}, we also have that

Thus, when βi(t)+γi(t)≤ai2\beta_{i}^{(t)}+\gamma_{i}^{(t)}\leq\frac{a_{i}}{2}, it can decrease at next iteration t+1t+1 only when δ=Ω(1κd)\delta=\Omega\left(\frac{1}{\kappa d}\right), in which case, the total decrement is bounded by exp⁡{−η∑t≤T∣Δ(t)∣1Δ(t)≥δ}\exp\{-\eta\sum_{t\leq T}|\Delta^{(t)}|1_{\Delta^{(t)}\geq\delta}\}. Therefore, taking δ=Θ(1κd)\delta=\Theta\left(\frac{1}{\kappa d}\right), with the fact that βi(0)≥1κd\beta_{i}^{(0)}\geq\frac{1}{\kappa d}, we obtain the result by combining equation A.42. ∎

Based on the above claim, we move on to the case when Δ(t)\Delta^{(t)} is negative. We show the following proposition.

In the setting of Lemma A.5, suppose that Proposition A.4 holds. Let δ\delta be any value in the range [poly⁡(κ2)ctd2,1κd][\frac{\operatorname{poly}(\kappa_{2})c_{t}}{d^{2}},\frac{1}{\kappa d}]. When Δ(t)≤−δ\Delta^{(t)}\leq-\delta, we have

We shall see that when Δ(t)≤−δ\Delta^{(t)}\leq-\delta, then for every ii with βi(t)+γi(t)≤ai+δ12max⁡{C1,C2}\beta_{i}^{(t)}+\gamma_{i}^{(t)}\leq a_{i}+\frac{\delta}{12\max\{C_{1},C_{2}\}}, we have that

On the other hand, γi(t)≤ai+δ′12max⁡{C1,C2}\gamma_{i}^{(t)}\leq a_{i}+\frac{\delta^{\prime}}{12\max\{C_{1},C_{2}\}} as well, this implies that

Notice that Δ(t)≤0\Delta^{(t)}\leq 0. This implies that

In the setting of Claim A.6 and A.8, for every T∈[T2+1,T3]T\in[T_{2}+1,T_{3}], the following holds:

To prove the above equation, we consider two scenarios. Using Claim A.6, for every δ∈[1d1.5,1κd]\delta\in\left[\frac{1}{d^{1.5}},\frac{1}{\kappa d}\right], we have:

Combined together, using the fact that T3≤dlog⁡1.01dηT_{3}\leq\frac{d\log^{1.01}d}{\eta}, we obtain equation (A.41). ∎

In the setting of Lemma A.5, suppose Proposition A.4 holds. There exists fixed constants C1,C2>0C_{1},C_{2}>0 such that for any t∈[T2+1,T3]t\in[T_{2}+1,T_{3}] and any i∈[d]i\in[d], the update of γi^(t),βi(t)\hat{\gamma_{i}}^{(t)},\beta_{i}^{(t)} satisfies that

Moreover, when γi(t)≥1poly⁡(d)\gamma_{i}^{(t)}\geq\frac{1}{\operatorname{poly}(d)}, we have that

The above claim implies that the update between the potential neurons and those not in the potential set differs by a multiplicative factor of ai−γi(t)a_{i}-\gamma_{i}^{(t)}. Intuitively, this gap allows us to show that the mass of potential neurons will converge and reduce the value of ai−γi(t)a_{i}-\gamma_{i}^{(t)}. On the other hand, the mass of bad neurons βi(t+1)\beta_{i}^{(t+1)} will remain polynomially small throughout the update, since its increment only scales with poly⁡(κ2)ct/d2\operatorname{poly}(\kappa_{2})c_{t}/d^{2} every iteration. We now describe the proof of the above proposition, which is based on a simple claim that bounds the gradient from irrelevant neurons in equation (A.44).

We first show the following claim. For every v∈Sgv\in\mathcal{S}_{g}, every i∈[d]i\in[d]:

To see that the above claim is true, for v∉Spotv\notin\mathcal{S}_{pot}, we can bound ∇v,2j,¬pot\nabla_{v,2j,\neg pot} as in Lemma A.7. For v∈Sbadv\in\mathcal{S}_{bad}, we can bound ∇v,2j,bad\nabla_{v,2j,bad} directly using equation (A.7). For v∈Si,potv\in\mathcal{S}_{i,pot}, we notice

On the other hand, when v∈Si,potv\in\mathcal{S}_{i,pot} and ∥v∥2≥d6\|v\|_{2}\geq d^{6}, we have that ∥vˉ−ei∥2≤1d4\|\bar{v}-e_{i}\|_{2}\leq\frac{1}{d^{4}} by Eq (A.6). This implies that for

By plugging in the claim in the beginning of the proof into the gradient update rule, we can prove the update rules for each set of neurons. For every v∉Spotv\notin\mathcal{S}_{pot} and every i∈[d]i\in[d], we have that

For every v∈Si,potv\in\mathcal{S}_{i,pot} with ∣vi∣≥d6|v_{i}|\geq d^{6}, we have that

By applying the above results on each set of neurons, we obtain the result of this claim. ∎

A.3.1 Proof of the Main Lemma

We are now ready to prove Lemma A.5. Based on the dynamic of 0th order tensor and the update of the 2nd order losses shown above, we prove the following proposition that shows βi(t)+γi(t)\beta_{i}^{(t)}+\gamma_{i}^{(t)} cannot be too far away from aia_{i} for too many iterations.

Moreover, for every δ≤1100κd\delta\leq\frac{1}{100\kappa d}, we have:

We consider an update step, then it holds that as long as βi(t)+γi(t)≤poly⁡(κ2)d\beta_{i}^{(t)}+\gamma_{i}^{(t)}\leq\frac{\operatorname{poly}(\kappa_{2})}{d}, using Claim A.10, the update of Φ(t)\Phi^{(t)} is given as:

Hence, consider the case that βi(t)+γi(t)=ai−ρ(t)\beta_{i}^{(t)}+\gamma_{i}^{(t)}=a_{i}-\rho^{(t)} for ρ(t)≥0\rho^{(t)}\geq 0, we have that γi(t)≤ai−ρ(t)\gamma_{i}^{(t)}\leq a_{i}-\rho^{(t)}. Hence in addition to Eq (A.48), we also have (using Claim A.7):

Note that originally Φ(0)=O(1d2)\Phi^{(0)}=O\left(\frac{1}{d^{2}}\right) using the fact that βi(0)≤2ai\beta^{(0)}_{i}\leq 2a_{i} and γi(0)≤1poly⁡(d)\gamma_{i}^{(0)}\leq\frac{1}{\operatorname{poly}(d)}, with Claim A.9, we have that for T≤dηlog⁡1.01dT\leq\frac{d}{\eta}\log^{1.01}d:

Similarly, we can see that when βi(t)+γi(t)=ai+ρ(t)\beta_{i}^{(t)}+\gamma_{i}^{(t)}=a_{i}+\rho^{(t)} for ρ(t)≥0\rho^{(t)}\geq 0, then either βi(t)≥ρ(t)/2\beta_{i}^{(t)}\geq\rho^{(t)}/2 or γi(t)≥ai−ρ(t)/2\gamma_{i}^{(t)}\geq a_{i}-\rho^{(t)}/2. In either case, we have that

Eventually, consider for every δ≤1100κd\delta\leq\frac{1}{100\kappa d}, when βi(t)+γi(t)≥ai+δ\beta_{i}^{(t)}+\gamma_{i}^{(t)}\geq a_{i}+\delta and ∣Δ(t)∣≤d2poly⁡(κ2)δ3|\Delta^{(t)}|\leq\frac{d^{2}}{\operatorname{poly}(\kappa_{2})}\delta^{3}, then we also have

Using equation A.43, we obtain that when T≤T3T\leq T_{3},

Based on the above proposition, we are ready to prove the main Lemma of Stage 2.1, which provides a warm start initialization at a certain iteration T3=Θ(dlog⁡d/η)T_{3}=\Theta(d\log d/\eta).

We first define T3T_{3} more precisely. We note that initially, for any i∈[d]i\in[d], γ^i(0)≤1/poly⁡(d)\hat{\gamma}_{i}^{(0)}\leq{1}/{\operatorname{poly}(d)} by construction. Using equation (A.41) and equation (A.46), by working on γ^i(t)\hat{\gamma}_{i}^{(t)} and noticing that γ^i(t)≤γi(t)\hat{\gamma}_{i}^{(t)}\leq\gamma_{i}^{(t)}, we have that there exists an iteration T(i)=O(dκlog⁡(1γ^i(0))/η)T^{(i)}=O({d\kappa\log(\frac{1}{\hat{\gamma}_{i}^{(0)}})}/{\eta}) such that at this iteration, γiT(i)≥110κd\gamma_{i}^{T^{(i)}}\geq\frac{1}{10\kappa d}. We shall fix T3T_{3} to be the maximum of T(i)T^{(i)} over i∈[d]i\in[d], which is on the order of Θ(dlog⁡d/η)\Theta(d\log d/\eta).

Next, similar to the proof of Proposition A.11, we consider the function

Let ii be the coordinate that achieves the maximum for the function above. We show that

Let μ=C1(ai−βi(t)−γi(t))+C2(ai−γi(t))\mu=C_{1}(a_{i}-\beta_{i}^{(t)}-\gamma_{i}^{(t)})+C_{2}(a_{i}-\gamma_{i}^{(t)}), ν=C1(ai−βi(t)−γi(t))\nu=C_{1}(a_{i}-\beta_{i}^{(t)}-\gamma_{i}^{(t)}), we have that

with βi(t)=μ−νC2−νC1\beta_{i}^{(t)}=\frac{\mu-\nu}{C_{2}}-\frac{\nu}{C_{1}}. So we have when γi(t)≥1poly⁡(κ3)d\gamma_{i}^{(t)}\geq\frac{1}{\operatorname{poly}(\kappa_{3})d},

When Φ(t)≥δ\Phi^{(t)}\geq\delta, we have that either μ2≥δ100\mu^{2}\geq\frac{\delta}{100}, or μ2≤δ100\mu^{2}\leq\frac{\delta}{100} and ν2≥δ2\nu^{2}\geq\frac{\delta}{2}. In the first case, we have that

Combining this equation with the bound in equation (A.41), we know that for δ=1dlog⁡0.01d\delta=\frac{1}{d\log^{0.01}d}, we have that Φ(t)≥δ\Phi^{(t)}\geq\delta can only happen for at most dlog⁡0.5dη\frac{d\log^{0.5}d}{\eta} many of the iterations within t∈[T2+1,T3]t\in[T_{2}+1,T_{3}]. Combining the above with equation (A.43), we obtain the desired result. ∎

A.4 Stage 2.2: The Final Substage

In this section, we present the proof of Lemma A.6 for the final substage. In the end, we prove the running inductive hypothesis H1\mathcal{H}_{1} in Proposition A.4.

Suppose the lemma holds at iteration tt, then using the condition at iteration tt, together with Φ(0)=O(1dlog⁡0.01d)\Phi^{(0)}=O\left(\frac{1}{d\log^{0.01}d}\right), we have that

Similar to the proof of Lemma A.8, we have that as long as Δ+(t)≥poly⁡(κ2)ctd2\Delta_{+}^{(t)}\geq\frac{\operatorname{poly}(\kappa_{2})c_{t}}{d^{2}} and

Then it must satisfy that Δ+(t+1)≤Δ+(t)(1−η1dpoly⁡(κ))\Delta_{+}^{(t+1)}\leq\Delta_{+}^{(t)}\left(1-\eta\frac{1}{d\operatorname{poly}(\kappa)}\right). Hence, if the maximizer of Φ\Phi is Δ+\Delta_{+}, Then it must be the case that

Now, consider another case when Δ−(t)≤(1−1poly⁡(κ))δ+(t)\Delta_{-}^{(t)}\leq\left(1-\frac{1}{\operatorname{poly}(\kappa)}\right)\delta_{+}^{(t)}, let ii be the argmax of {τj}j∈[d]\{\tau_{j}\}_{j\in[d]}, then we must have that

Hence as long as δ−(t)≥poly⁡(κ2)ctd2\delta_{-}^{(t)}\geq\frac{\operatorname{poly}(\kappa_{2})c_{t}}{d^{2}}, we have that

is a decreasing function of γi\gamma_{i} with slop at least 0.5γi0.5\gamma_{i} when γi≥ai2\gamma_{i}\geq\frac{a_{i}}{2}, which holds true using Eq (A.49). Combining Eq (A.51) and Eq (A.52), we have that if the maximizer of Φ\Phi is δ−\delta_{-}, the following is true

Consider another case when the maximizer is Δ−\Delta_{-}. Similar to the proof of Lemma A.8, as long as Δ−(t)≥poly⁡(κ2)ctd2\Delta_{-}^{(t)}\geq\frac{\operatorname{poly}(\kappa_{2})c_{t}}{d^{2}}, we have that

Hence, if the maximizer of Φ\Phi is Δ−\Delta_{-}, then it must be the case that

Moreover, using the fact that when Δ−(t)≤(1−1poly⁡(κ))δ+(t)\Delta_{-}^{(t)}\leq\left(1-\frac{1}{\operatorname{poly}(\kappa)}\right)\delta_{+}^{(t)}, let ii be the argmax of {−τj}j∈[d]\{-\tau_{j}\}_{j\in[d]}, then we must have that as long as δ+(t)≥poly⁡(κ2)ctd2\delta_{+}^{(t)}\geq\frac{\operatorname{poly}(\kappa_{2})c_{t}}{d^{2}}, we have that

The maximizer is δ+\delta_{+}. Then we must have that for every i∈[d]i\in[d], βi(t)≤Cδ+(t)\beta_{i}^{(t)}\leq C\delta_{+}^{(t)}, then we must have that ai−γ(t)≤Cδ+(t)a_{i}-\gamma^{(t)}\leq C\delta_{+}^{(t)} as well. Hence, it holds that

Hence if the maximizer of Φ\Phi is δ+\delta_{+}, it must be the case:

The maximizer is β+\beta_{+}. Then there is a j∈[d]j\in[d] such that βj(t)≥Cδ+(t)\beta_{j}^{(t)}\geq C\delta_{+}^{(t)} , βj(t)≥Cδ−(t)\beta_{j}^{(t)}\geq C\delta_{-}^{(t)} and βj(t)≥C∣Δ(t)∣\beta_{j}^{(t)}\geq C|\Delta^{(t)}|, we have that for this jj, it holds: let S=C1(βj(t)+γj(t)−aj)S=C_{1}\left(\beta_{j}^{(t)}+\gamma_{j}^{(t)}-a_{j}\right) and ρ=(aj−γj(t))\rho=(a_{j}-\gamma_{j}^{(t)}), we have: if ρ≤14βj(t)\rho\leq\frac{1}{4}\beta_{j}^{(t)}, then

On the other hand if ρ>14βj(t)\rho>\frac{1}{4}\beta_{j}^{(t)}, then using δ−(t)\delta_{-}^{(t)}, we have:

Hence if the maximizer of Φ\Phi is β+\beta_{+}, it must be the case:

To sum up, the result follows by combining Eq (A.56), (A.55), (A.54), (A.50) and (A.53). ∎

We first verify the inductive hypothesis for t≤T3t\leq T_{3}. The bound for v∉Spotv\notin\mathcal{S}_{pot} follows from Claim A.9 and Proposition A.11. We prove the bound for v∈Spotv\in\mathcal{S}_{pot} by tracking the gradient descent dynamic. Following Eq (A.31), for every neuron vv, and every p∈[d]p\in[d], define

Hence, using Eq (A.44), we have that for every i∈[d]i\in[d]

Hence, we have that for every i,j∈[d]i,j\in[d],

Now, if ∣vˉi(t)∣2,∣vˉj(t)∣2≤ctd|\bar{v}_{i}^{(t)}|^{2},|\bar{v}_{j}^{(t)}|^{2}\leq\frac{c_{t}}{d}, we also have that

Hence using Proposition A.11 we show that when ∣vˉi(0)∣2,∣vˉj(0)∣2≤c0d|\bar{v}_{i}^{(0)}|^{2},|\bar{v}_{j}^{(0)}|^{2}\leq\frac{c_{0}}{d}, then ∣vˉi(t)∣2,∣vˉj(t)∣2≤ctd|\bar{v}_{i}^{(t)}|^{2},|\bar{v}_{j}^{(t)}|^{2}\leq\frac{c_{t}}{d} as well for every t≤T3t\leq T_{3}. Now, we need to give an upper bound on the the coordinates of the neurons. For every v∉Spotv\notin\mathcal{S}_{pot}, we know that all coordinates j∈[d]j\in[d] satisfies that ∣vˉj(0)∣2≤c0d|\bar{v}_{j}^{(0)}|^{2}\leq\frac{c_{0}}{d}. Hence, by Eq (A.57), we have that

Hence using Proposition A.11 and Claim A.9, we have proved Eq (A.5) and Eq (A.6).

Next, we proceed to the norm of neurons v∈Si,goodv\in\mathcal{S}_{i,good}. For this neuron, using the fact that ∣vi(t)∣≥d6|v_{i}^{(t)}|\geq d^{6} and equation A.45, we have that

Hence, for every tt, using Eq (A.46) we obtain that:

Notice that for every neuron v∈Sgv\in\mathcal{S}_{g}, we have that ∣vˉi(0)∣2≤c0d|\bar{v}_{i}^{(0)}|^{2}\leq\frac{c_{0}}{d} can happen for at most O(log⁡0.01d)O(\log^{0.01}d) many i∈[d]i\in[d]. Denote this set as Qv\mathcal{Q}_{v}, we have that [vi(t+1)]2/[vi(t)]2\left[v_{i}^{(t+1)}\right]^{2}/\left[v_{i}^{(t)}\right]^{2} is at most

Hence, for every t≤T≤T3t\leq T\leq T_{3}, using Eq (A.41) and Eq (A.46), by working on γ^i\hat{\gamma}_{i} and notice that γ^i≤γi\hat{\gamma}_{i}\leq\gamma_{i}, we conclude that for every i∈[d]i\in[d].

Combining the above equation with Eq (A.46) we have that for every i∈[d]i\in[d]:

This proves that gradient truncation never happens during this substage. Now, apply Lemma A.3, which says that

We complete the proof of the first our statements. For the last statement on γ,β\gamma,\beta, Claim A.8 also proves the upper bound on γi(t)+βi(t)\gamma_{i}^{(t)}+\beta_{i}^{(t)} as in equations (A.8) and (A.9). Taking δ=1κd\delta=\frac{1}{\kappa d}, we can show that

Next verify the running inductive hypothesis H2\mathcal{H}_{2} for T3≤t≤T4T_{3}\leq t\leq T_{4}. Based on Lemma A.6, we have the following bounds on the update of each coordinate of each neuron. For every i∈[d]i\in[d], using Eq (A.44), we have that

Note that by the definition of Φ(t)\Phi^{(t)} at Lemma A.6, we have that

Hence, we obtain that for t∈[T3,T4]t\in[T_{3},T_{4}], with Lemma A.6:

Hence as long as T4≤d1+10QηT_{4}\leq\frac{d^{1+10Q}}{\eta}, we obtain the running hypothesis H2\mathcal{H}_{2} at this substage. ∎

Appendix B Proof of the Finite-Width Case

where Ξw(t)\Xi_{w}^{(t)} is an extra error term that arises from the sampling error of the empirical loss.

Our main result in this section is that provided with polynomially many neuron samples and training samples, the errors ξw(t)\xi_{w}^{(t)} and Ξw(t)\Xi_{w}^{(t)} in equation (B.2) remain polynomially small throughout Algorithm 1. We first state the result for Stage 1.

The proof of Lemma B.1 can be found in Section B.2.3. Next, we consider the error propagation of Stage 2.1. We show that the norm of ξw\xi_{w} is much smaller than that of ww.

The proof of Lemma B.2 involves carefully studying the error term and follows a similar argument to Lemma B.1. The details can be found in Section B.3. Finally, we consider the error terms in the final stage. We use a different error analysis. At iteration T3T_{3}, let us consider the set

Let Ssingleton=∪i=1dSi,singleton\mathcal{S}_{singleton}=\cup_{i=1}^{d}\mathcal{S}_{i,singleton}. Consider the set

where we recall the definition of Spot\mathcal{S}_{pot} in Proposition A.4. We state the error propagation of the final substage as follows.

The proof of Lemma B.3 can be found in Section B.4. Based on the analysis of error propagation, we are now ready to prove our main result. We prove Theorem 1.1 as follows.

where L^(W)\hat{L}(W) denotes the empirical loss.

These statements together give us the following

Finally, combined with Claim 2.1 and Theorem 3.1 we complete the proof of Theorem 1.1. ∎

where A=diag({ai}i∈[d])A=\text{diag}(\{a_{i}\}_{i\in[d]}). Using the update of equation (B.2) for the finite-width case, the gradient of vv for the 0th and 2nd order terms over the population loss is given by

The first order terms of the error term ξv\xi_{v} for the neuron vv is given by

The first inequality is obviously true. Now we consider the second inequality, we have that

Next we consider the first order tensor. The first order gradient in the finite-width case for the population loss is

The 1st order loss in the gradient is zero in the infinite-width case of Section A. The first-order expansion of the error is given by:

We have the following claim for the error in the first order gradients.

B.2 Stage 1.2: Analysis of Higher Order Tensor Decompositions

In this substage, we consider the error terms of the gradients of the higher order tensor decompositions. Towards showing the error propagation in Lemma B.1, our proof outline is as follows.

We decompose the error of the gradients into individual terms that we analyze one by one.

In Proposition B.4, we provide an upper bound on the average norm of the error. In Proposition B.7, we bound the error of the individual terms from the decomposition. Finally, we present the proof of Lemma B.1 in Section B.2.3.

We begin by writing down the gradient of higher order terms for the population loss..

As long as for every w∈Sw\in\mathcal{S} (recalling its definition in Def. A.2), ∥ξw∥2≤1poly⁡(d)\|\xi_{w}\|_{2}\leq\frac{1}{\operatorname{poly}(d)}, then we have

Next we show that the norm of the error in each individual neuron can also be bounded.

The proof of Proposition B.4 and Proposition B.5 is left to Section B.2.3.

In addition, we show that the second order terms in ξ\xi that contains ∥ξw∥2p\|\xi_{w}\|_{2}^{p} and ∥ξv∥2q\|\xi_{v}\|_{2}^{q} for p+q≥3p+q\geq 3 are of a lower order compared to the first order terms. Informally, we know that ∥ξw∥2\|\xi_{w}\|_{2} and ∥ξv∥2\|\xi_{v}\|_{2} are less than λ02\lambda_{0}^{2}. Meanwhile, ∥w∥2\|w\|_{2} and v∥2v\|_{2} are at least Ω(1d)\Omega\left(\frac{1}{d}\right), for every w,v∈Sgw,v\in\mathcal{S}_{g} by Lemma A.3. Combined together, we show the following result.

B.2.2 Individual Error Norm bound

Based on the decomposition above, we provide several helper claims for bounding the error of the gradient terms. First, for v∈Sv\in\mathcal{S}, we have the following claim.

In the setting of Proposition B.4, we have that

The second inequality in the Lemma follows from the fact that (⟨w,v⟩)w,v,(wˉ,vˉ⟩)w,v,(⟨ξw,ξv⟩)w,v(\langle w,v\rangle)_{w,v},(\bar{w},\bar{v}\rangle)_{w,v},(\langle\xi_{w},\xi_{v}\rangle)_{w,v} forms PSD matrices, and the Hadamard product of PSD matrices is PSD. ∎

We also have the following claim, which serves as an upper bound of

In the setting of Proposition B.4, we have that

As a corollary, combine the above inequality with Proposition A.1, we obtain

The proof is a direct calculation, using ⟨wˉ,vˉ⟩2≤ctd\langle\bar{w},\bar{v}\rangle^{2}\leq\frac{c_{t}}{d} for w∈Sw\in\mathcal{S}, we have that

Now, we can easily calculate that (using the Eq (A.2))

which completes the proof. For the other two inequalities, we can bound them in the exact same way. ∎

The final claim aims to bound the rest of the terms.

In the setting of Proposition B.4, we have that

As a corollary, combining the above inequality with Proposition A.1, we obtain

Below we also consider the error individually, we will mainly focus on the error term with ξv\xi_{v}.

In the setting of Proposition B.4, we have that

For w∉Sw\notin\mathcal{S}, we can naively bound ∣⟨w,vˉ⟩⟨wˉ,vˉ⟩2⟨ξv,vˉ⟩⟨w,ξv⟩∣≤∥w∥22∥ξv∥22|\langle w,\bar{v}\rangle\langle\bar{w},\bar{v}\rangle^{2}\langle\xi_{v},\bar{v}\rangle\langle w,\xi_{v}\rangle|\leq\|w\|_{2}^{2}\|\xi_{v}\|_{2}^{2}. Hence, using Eq (A.4), we have:

This claim together with Claim B.5 implies that

In the setting of Proposition B.4, for every v∈Sv\in\mathcal{S}, we have that

Now we move on to the harder terms, we have the following claim.

In the setting of Proposition B.4, for every v∈Sv\in\mathcal{S}, we have that for p=11,13p=11,13:

We first consider p=11p=11. Let Q2j,v,11′=−(b2j+b2j′)∑i(ai⟨ei,vˉ⟩2j−2eiei⊤).Q_{2j,v,11}^{\prime}=-\left(b_{2j}+b_{2j}^{\prime}\right)\sum_{i}\left(a_{i}\langle e_{i},\bar{v}\rangle^{2j-2}e_{i}e_{i}^{\top}\right). We have that

Let us assume that ∥vˉ∥∞=1−δ\|\bar{v}\|_{\infty}=1-\delta for some value δ≥0\delta\geq 0, then we have that ∥vˉ−er∥22=O(δ)\|\bar{v}-e_{r}\|_{2}^{2}=O(\delta).

Using the fact that b2j,b2j′=Θ(1j2)b_{2j},b_{2j}^{\prime}=\Theta(\frac{1}{j^{2}}), we know that

Note that ∑j≥21j(1−δ)j=(1−δ)log⁡1δ\sum_{j\geq 2}\frac{1}{j}(1-\delta)^{j}=(1-\delta)\log\frac{1}{\delta} we obtain:

which completes the proof. On the other hand, for p=13p=13, let Q2j,v,13′=b2j′(∑iai⟨ei,vˉ⟩2j−1vˉei⊤).Q_{2j,v,13}^{\prime}=b_{2j}^{\prime}\left(\sum_{i}a_{i}\langle e_{i},\bar{v}\rangle^{2j-1}\bar{v}e_{i}^{\top}\right).. We have that

We can bound the terms in a similar way. ∎

Using the aforementioned claims, we conclude the proof of the following proposition.

In the setting of Proposition B.4, for every vv, we have that

For p=4,6,8,9,10p=4,6,8,9,10, using Claim B.6, we have

When ∥w∥2,∥ξw∥2≤1λ0\|w\|_{2},\|\xi_{w}\|_{2}\leq\frac{1}{\lambda_{0}}, for p=2,3,5,7p=2,3,5,7, the following is true

For p=11,12,13,14p=11,12,13,14, we have that for p=11,13p=11,13, using Claim B.8, we get

B.2.3 Proof of Error Propagation

Based on the individual error norm bound and the average error norm bound, we are ready to prove the main result of stage 1. We first state the proof of the individual error norm bound.

On the other hand, by Eq (A.33), we have that for this vv, if the gradient clipping is not performed, then by the definition of Sg\mathcal{S}_{g}, we have that ∥vˉ(t)∥2≥1log⁡d\|\bar{v}^{(t)}\|_{2}\geq\frac{1}{\log d}. Therefore,

which implies that after t′=O(λ0dlog⁡3d∥ξv(t)∥2η)t^{\prime}=O\left(\frac{\sqrt{\lambda_{0}}d\log^{3}d\|\xi_{v}^{(t)}\|_{2}}{\eta}\right), many iterations, if gradient clipping is not performed, we should have that

Since each iteration shall introduce at most O(η1λ0)O\left(\eta\frac{1}{\lambda_{0}}\right) amount error, so we have:

This gives us the final error bound of the individual error when combined with Claim B.5. ∎

Next we state the proof of the average error norm bound.

Using Proposition B.7 (together with Eq (B.4)) and by the definition of Eq (B.3), we can obtain the desired result. ∎

Based on Proposition B.4, Proposition B.5, and Proposition B.7, we are ready to prove Lemma B.1.

and for every neuron vv, ∥ξv(t)∥22≤poly⁡(d)λ0Ξ\|\xi_{v}^{(t)}\|_{2}^{2}\leq\frac{\operatorname{poly}(d)}{\lambda_{0}}\Xi. Suppose this is true for all t≤T0t\leq T_{0}, then consider t=T0+1t=T_{0}+1. We apply Proposition B.4, which says that as long as for every w∈Sw\in\mathcal{S}, ∥ξw∥2≤1d3\|\xi_{w}\|_{2}\leq\frac{1}{d^{3}}, we have that

By m≥poly⁡(d)poly⁡(λ0)m\geq\frac{\operatorname{poly}(d)}{\operatorname{poly}(\lambda_{0})}, a simple Chernoff bound gives us:

Now, using the update rule of Eq (A.17) and in Proposition A.9, we have that

Note that at iteration 0, ε0=0\varepsilon_{0}=0. This implies that for t+1t+1: εt+1≤poly⁡(d)Ξ\varepsilon_{t+1}\leq\operatorname{poly}(d)\Xi as well. Combine this with Proposition B.7 on the individual norm bound we complete the proof. ∎

B.3 Stage 2.1: Analysis After Reducing the Gradient Truncation Parameter

For every individual neuron vv and every value α≥1\alpha\geq 1, the following holds

The proof of this claim is quite straightforward. We have that for p=1,2p=1,2, we use Claim B.3, which gives us:

For p=3p=3, we use that ∣⟨wˉ,vˉ⟩∣≤1|\langle\bar{w},\bar{v}\rangle|\leq 1 and

For p=8,9,10,11p=8,9,10,11, the result can be obtained similarly. For p=12,13p=12,13, we use that

Finally, the individual error bound comes from the following simple calculation.

Note that this substage has T3T_{3} many iterations, where T3T_{3} is upper bounded by dC(κ)log⁡dη\frac{dC(\kappa)\log d}{\eta} for some value C(κ)>0C(\kappa)>0 that only depends on κ\kappa. By by taking α=1\alpha=1 in Claim B.9, the rest of the proof is similar to the proof of Lemma B.1. We omit the details. ∎

B.4 Stage 2.2: The Final Substage

We provide the proof of Lemma B.3, which analyzes the error propagation in the final substage. Recall that Si,singleton\mathcal{S}_{i,singleton} and Signore\mathcal{S}_{ignore} have been defined in the beginning of this section. At the beginning of Stage 2.2 when t=T3+1t=T_{3}+1, we do a modification:

If vv in Si,singleton\mathcal{S}_{i,singleton} we will just set vˉ=ei\bar{v}=e_{i} and keep the norm not changed.

If vv in Signore\mathcal{S}_{ignore}, then we will just set v=0v=0.

Thus, we can see that v(t)=0v^{(t)}=0 for every v∈Signorev\in\mathcal{S}_{ignore} and for every t>T3t>T_{3}. We define a new update for the infinite neuron process at this substage for v∈Si,singletonv\in\mathcal{S}_{i,singleton}. We define v+,v−v_{+},v_{-} such that at every iteration t≥T3t\geq T_{3}:

We will replace vv in the infinite neuron process with two neurons v+,v−v_{+},v_{-}. For the simplicity of notation, we write v+v_{+} simply as vv. For the other neurons, the update does not change.

By the running hypothesis H1\mathcal{H}_{1} in Proposition A.4, we have that at iteration T3T_{3},

Moreover, throughout the entire process, by Lemma A.6, we will always have that

For every v∈Sgv\in\mathcal{S}_{g}, we have that:

For every v∉Spotv\notin\mathcal{S}_{pot} (cf. Lemma A.3 for the definition),

We consider v∈Si,singletonv\in\mathcal{S}_{i,singleton}. For these neurons, we have that

as the gradient of vv involving only a single neuron ww. For p=5,6,7,8,10,13,14,16p=5,6,7,8,10,13,14,16. Now, for p=3p=3 we have that

For the second term, we have that when j=ij=i, we have for every w∈Sj,singletonw\in\mathcal{S}_{j,singleton}: ⟨wˉ,ξv⟩=0\langle\bar{w},\xi_{v}\rangle=0. Otherwise, when j≠ij\not=i, we have that ⟨wˉ,vˉ⟩=0\langle\bar{w},\bar{v}\rangle=0. Therefore,

For p=9,15p=9,15, following the same calculation by dividing ww into three parts we can easily conclude that

Next, we consider the error of the neurons not in Spot\mathcal{S}_{pot}. We use a direct corollary of Claim B.9 , except that for every v∉Spotv\notin\mathcal{S}_{pot}, it holds that ⟨vˉ,wˉ⟩2≤ctd\langle\bar{v},\bar{w}\rangle^{2}\leq\frac{c_{t}}{d} instead of 11 for every vector ww. We state the result as follows.

Based on Claim B.10 and B.11, we prove Lemma B.3.

Let α=d\alpha=\sqrt{d} in Claim B.11, we show the following result the bound. For every v∉Spotv\notin\mathcal{S}_{pot},

Together with the individual error bound as in Claim B.9, we can obtain the desired result using a similar proof to Lemma B.1. The details are omitted. ∎

Appendix C Proof of Lower Bound

We follow the proof of Theorem 2 in Allen-Zhu and Li [2019a] for proving the lower bound. We first describe the construction of the hardness distribution W\mathcal{W}. We first show the following lemma.

For a positive integer rr, for every d≥r2d\geq r^{2} which is a multiple of rr, there exists at least H=dΩ(r)H=d^{\Omega(r)} many sets C(j)={C1(j)∈[d],⋯ ,Cd/r(j)∈[d]}\mathcal{C}^{(j)}=\{\mathcal{C}_{1}^{(j)}\in[d],\cdots,\mathcal{C}_{d/r}^{(j)}\in[d]\} for j=1,…,Qj=1,\dots,Q such that

For every 1≤i≤d/r1\leq i\leq d/r and 1≤j≤H1\leq j\leq H, Ci(j)\mathcal{C}_{i}^{(j)} is a subset of [d][d] of size rr.

For every 1≤i≠i′≤d/r1\leq i\not=i^{\prime}\leq d/r and 1≤j≤H1\leq j\leq H, Ci(j)∩Ci′(j)=∅\mathcal{C}_{i}^{(j)}\cap\mathcal{C}_{i^{\prime}}^{(j)}=\emptyset.

For every 1≤i,i′≤d/r1\leq i,i^{\prime}\leq d/r and 1≤j≠j′≤H1\leq j\not=j^{\prime}\leq H, Ci(j)≠Ci′(j′)\mathcal{C}_{i}^{(j)}\not=\mathcal{C}_{i^{\prime}}^{(j^{\prime})}.

We consider a uniformly at random distribution over the set C={C1,⋯ ,Cd/r}\mathcal{C}=\{\mathcal{C}_{1},\cdots,\mathcal{C}_{d/r}\}, where Ci\mathcal{C}_{i} is a subset of [d][d] of size rr and for every i≠i′i\not=i^{\prime}, we have that Ci∩Ci′=∅\mathcal{C}_{i}\cap\mathcal{C}_{i^{\prime}}=\emptyset. Let us sample QQ many sets {C(j)}j∈[Q]\{\mathcal{C}^{(j)}\}_{j\in[Q]} from it, then using union bound, we have that:

Hence when d≥r2d\geq r^{2}, for some H=dO(r)H=d^{O(r)}, the above probability is smaller than one. This proves the existence of these sets. ∎

Now, we define the distribution W\mathcal{W}. Recall that the Hadamard transform of dimension rr is a unitary matrix in dimension rr whose entries are all ∈{−1/r,1/r}\in\{-1/\sqrt{r},1/\sqrt{r}\}.

For every rr that is a power of 22, for every dd that is a multiple of rr bigger than r2r^{2}, we generate W\mathcal{W} as:

Pick C\mathcal{C} uniformly at random from the set {C(j)}j∈[H]\{\mathcal{C}^{(j)}\}_{j\in[H]} given by Lemma C.1.

where hq⋆h^{\star}_{q} is the i-th column of the Hadamard transform of dimension rr.

Sample b1,⋯ ,bdb_{1},\cdots,b_{d} independent from $$ uniformly at random. Define

The proof of the lower bound relies on the following Lemma.

To prove this Lemma, we use Lemma F.2F.2 and the proof of Corollary 7.17.1 in Allen-Zhu and Li [2019a], which says the following.

Using this Corollary, we can prove Lemma C.2.

Hence we have that by τi∈{−1,1}\tau_{i}\in\{-1,1\}:

Using the fact that ∣h(μ)∣≤V|h(\mu)|\leq V, we have that

Notice that with probability at least 1−er21-e^{r^{2}} over μ\mu, we have that λμ≤rO(r)\lambda_{\mu}\leq r^{O(r)}. Note that λμ≥0\lambda_{\mu}\geq 0 as well. Thus, using Markov’s inequality we complete the proof. ∎

Next we can derive the following corollary of Lemma C.2. For two vectors x,yx,y with the same dimension, we denote x∘yx\circ y as the entry-wise product of x,yx,y.

Let p1,⋯ ,prp_{1},\cdots,p_{r} be rr vectors in {−1/r,1/r}r\{-1/\sqrt{r},1/\sqrt{r}\}^{r}, let q1,⋯ ,qrq_{1},\cdots,q_{r} be i.i.d. random variable chosen uniformly at random from $,define, defineF_{\mu}(\tau)=\sum_{i\in[r]}q_{i}\left|\langle p_{i}\circ\mu,\tau\rangle\right|,wehavethatwithprobabilityatleast, we have that with probability at leastr^{-O(r)}overover\mu\sim\mathcal{N}(0,\operatorname{Id}_{d\times d})andandq$:

Finally, we can complete the proof of Theorem 1.2.

We prove by contradiction. Suppose on the contrary that equation (1.3) does not hold. Then, there exists ≥0.01\geq 0.01 fraction of {ai,w⋆}i∈[d]\{a_{i},w^{\star}\}_{i\in[d]} generated from W\mathcal{W} such that for some w(R)w^{(R)} we have R(x):=wR⊤ϕ(x)\mathcal{R}(x):=w_{R}^{\top}\phi(x), and it holds that

We consider x=xˉ∘τx=\bar{x}\circ\tau where xˉ∼N(0,Id⁡d×d)\bar{x}\sim\mathcal{N}(0,\operatorname{Id}_{d\times d}) and τ∼Uniform({−1,1}d)\tau\sim Uniform(\{-1,1\}^{d}). Clearly, x∼N(0,Id⁡d×d)x\sim\mathcal{N}(0,\operatorname{Id}_{d\times d}) as well. Thus,

Therefore, by Markov’s inequality we have that with probability at least 0.9990.999 over the choice of xˉ\bar{x}, we have that

where λB\lambda_{\mathcal{B}} if the Fourier coefficient of the subset BB. Now, define λB⋆\lambda_{\mathcal{B}}^{\star} to be the Fourier coefficients of f⋆(xˉ∘τ)f^{\star}(\bar{x}\circ\tau) and λBR\lambda_{\mathcal{B}}^{\mathcal{R}} to be the Fourier coefficient of R(xˉ∘τ)\mathcal{R}(\bar{x}\circ\tau), we can observe that if we sample C\mathcal{C} from W\mathcal{W} to generate w⋆w^{\star} according to Definition C.1, then it holds that for every B⊂[d]\mathcal{B}\subset[d] of size rr, we have:

Moreover, using Corollary C.4, we can conclude that w.p. at least 0.9990.999 over W\mathcal{W},

Let us consider the set Sgd\mathcal{S}_{gd} of {ai,w⋆}i∈[d]\{a_{i},w^{\star}\}_{i\in[d]} generated from W\mathcal{W}. We call {ai,wi⋆}i∈[d]∈Sgd\{a_{i},w_{i}^{\star}\}_{i\in[d]}\in\mathcal{S}_{gd} if and only if the function f⋆f^{\star} defined using {ai,w⋆}i∈[d]\{a_{i},w^{\star}\}_{i\in[d]} satisfies Eq (C.2) and there is a w(R)w^{(R)} such that for R(x):=wR⊤ϕ(x)\mathcal{R}(x):=w_{R}^{\top}\phi(x) with

We already know that there are at least 0.9990.999 fraction {ai,w⋆}i∈[d]\{a_{i},w^{\star}\}_{i\in[d]} generated from W\mathcal{W} that satisfies Eq (C.2). By our assumption, there are ≥0.01\geq 0.01 fraction of {ai,w⋆}i∈[d]\{a_{i},w^{\star}\}_{i\in[d]} generated from W\mathcal{W} satisfying that for some w(R)w^{(R)} such that R(x):=wR⊤ϕ(x)\mathcal{R}(x):=w_{R}^{\top}\phi(x), it holds that

Thus, we can conclude ∣Sgd∣≥0.005∣W∣|\mathcal{S}_{gd}|\geq 0.005|\mathcal{W}|. Together with Lemma C.1 which shows that ∣W∣≥dΩ(r)|\mathcal{W}|\geq d^{\Omega(r)}, we know that ∣Sgd∣≥dΩ(r)|\mathcal{S}_{gd}|\geq d^{\Omega(r)}.

Now, we consider a matrix MM, whose rows are indexed by each set of {ai,wi⋆}i∈[d]∈Sgd\{a_{i},w_{i}^{\star}\}_{i\in[d]}\in\mathcal{S}_{gd} and Eq (C.1), whose columns are indexed by λB⋆\lambda_{\mathcal{B}}^{\star} with ∣B∣=r|\mathcal{B}|=r.

We know that this matrix is of size dΩ(r)×dΩ(r)d^{\Omega(r)}\times d^{\Omega(r)}. Moreover, for any matrix M′M^{\prime} satisfies that

where MiM_{i} is the ii-th row of MM. It must holds that rank⁡(M′)=dΩ(r)\operatorname*{rank}(M^{\prime})=d^{\Omega(r)}. We immediately complete the proof by contradiction, following exactly the same argument in the lower bound proof in Allen-Zhu and Li [2019a] while taking rr to be a sufficiently large constant. ∎