Does Preprocessing Help Training Over-parameterized Neural Networks?

Zhao Song, Shuo Yang, Ruizhe Zhang

Introduction

Over the last decade, deep learning has achieved dominating performance over many areas, e.g., computer vision [LBBH98, KSH12, SLJ+15, HZRS16], natural language processing [CWB+11, DCLT18], game playing [SHM+16, SSS+17] and beyond. The computational resource requirement for deep neural network training grows very quickly. Designing a fast and provable training method for neural networks is, therefore, a fundamental and demanding challenge.

Almost all deep learning models are optimized by gradient descent (or its variants). The total training time can be split into two components, the first one is the number of iterations and the second one is the cost per spent per iteration. Nearly all the iterative algorithms for acceleration can be viewed as two separate lines of research correspondingly, the first line is aiming for an algorithm that has as small as possible number of iterations, the second line is focusing on designing as efficient as possible data structures to improve the cost spent per iteration of the algorithm [Vai89, CLS19, LSZ19, JLSW20, JKL+20, JSWZ21]. In this paper, our major focus is on the second line.

Is it possible to improve the cost per iteration of training neural network algorithm? E.g., is o(mnd)o(mnd) possible?

We provide a new theoretical framework for speeding up neural network training by: 1) adopting the shifted neural tangent kernel; 2) showing that only a small fraction (o(m)o(m)) of neurons are activated for each input data in each training iteration; 3) identifying the sparsely activated neurons via geometric search; 4) proving that the algorithm can minimize the training loss to zero in a linear convergence rate.

We provide two theoretical results 1) our first result (Theorem 6.1) builds a dynamic half-space report data structure for the weights of a neural network, to train neural networks in sublinear cost per iteration; 2) our second result (Theorem 6.2) builds a static half-space report data-structure for the input data points of the training data set for training a neural network in sublinear time.

The goal of our paper is to theoretically characterize the acceleration brought by the high-dimensional geometric data structure. Specifically, our algorithm and analysis are built upon the HSR data structures [AEM92] which can find all the points that have large inner products and support efficient data update. Note that HSR comes with a stronger recovery guarantee than LSH, in the sense that HSR, whereas LSH is guaranteed to find some of those points.

Convergence via over-parameterization.

Over the last few years, there has been a tremendous work studying the convergence result of deep neural network explicilty or implicitly based on neural tangent kernel (NTK) [JGH18], e.g. [LL18, DZPS19, AZLS19a, AZLS19b, DLL+19, ADH+19a, ADH+19b, SY19, CGH+19, ZMG19, CG19, ZG19, OS20, LSS+20, JT20, ZPD+20, HLSY21, BPSW21]. It has been shown that (S)GD can train a sufficiently wide NN with random initialization will converge to a small training error in polynomial steps.

Challenges and Techniques

Empirical works combine high-dimensional search data structures (e.g., LSH) with neural network training, however, they do not work theoretically due to the following reasons:

Without shifting, the number of activated (and therefore updated) neurons is Θ(m)\Theta(m). There is no hope to theoretically prove o(m)o(m) complexity (See Challenge 1).

Approximate high-dimensional search data structures might miss some important neurons, which can potentially prevent the training from converging (see Challenge 2).

We propose a shifted ReLU activation that is guaranteed to have o(m)o(m) number of activated neurons. Along with the shifted ReLU, we also propose a shifted NTK to rigorously provide a convergence guarantee (see Solution 1).

We adopt an exact high-dimensional search data structure that better couples with the shifted NTK. It takes o(m)o(m) time to identify the activated neurons and fits well with the convergence analysis as it avoids missing important neurons (see Solution 2).

Solution 1

The problem actually comes from the activation function. In practice, people use a shifted ReLU function σb(x)=max⁡{⟨wr,x⟩,br}\sigma_{b}(x)=\max\{\langle w_{r},x\rangle,b_{r}\} to train neural networks. The main observation of our work is that threshold implies sparsity. We consider the setting where all neurons have a unified threshold parameter bb. Then, by the concentration of Gaussian distribution, there will be O(exp⁡(−b2)⋅m)O(\exp(-b^{2})\cdot m) activated neurons after the initialization.

The next step is to show that the number of activated neurons will not blow up too much in the following training iterations. In [DZPS19, SY19], they showed that the weights vectors are changing slowly during the training process. In our work, we open the black box of their proof and show a similar phenomenon for the shifted ReLU function. More specifically, a key component is to prove that for each training data, a large fraction of neurons will not change their status (from non-activated to activated and vice versa) in the next iteration with high probability. To achieve this, they showed that this is equivalent to the event that a standard Gaussian random variable in a small centered interval [−R,R][-R,R], and applied the anti-concentration inequality to upper-bound the probability. In our setting, we need to upper-bound the probability of z∼N(0,1)z\sim{\cal N}(0,1) in a shifted interval [b−R,b+R][b-R,b+R]. On the one hand, we can still apply the anti-concentration inequality by showing that the probability is at most Pr⁡[z∈[−R,R]]\Pr[z\in[-R,R]]. On the other hand, this probability is also upper-bounded by Pr⁡[z>b−R]\Pr[z>b-R], and for small RR, we can apply the concentration inequality for a more accurate estimation. In the end, by some finer analysis of the probability, we can show that with high probability, the number of activated neurons in each iteration is also O(exp⁡(−b2)⋅m)O(\exp(-b^{2})\cdot m) for each training data. If we take b=Θ(log⁡m)b=\Theta(\sqrt{\log m}), we only need to deal with truly sublinear in mm of activated neurons in the forward evaluation.

Challenge 2: How to find the small subset of activated neurons?

A linear scan of the neurons will lead to a time complexity linear in mm, which we hope to avoid. Randomly sampling or using LSH for searching can potentially miss important neurons which are important for a rigorous convergence analysis.

Solution 2

Preliminaries

This section is organized as follows. Section 3.1 introduces the neural network and present problem formulation. Section 3.2 presents the half-space report data-structure, Section 3.3 proposes our new sparsity-based Characterizations.

1 Problem Formulation

We say function ff is 2NN(m,b)\mathsf{2NN}(m,b) for simplicity.

For each rr, we sample wr(0)∼N(0,Id)w_{r}(0)\sim\mathcal{N}(0,I_{d})

For each rr, we sample ara_{r} from {−1,+1}\{-1,+1\} uniformly at random

To update the weights from iteration kk to iteration k+1k+1, we follow the standard update rule of the GD algorithm,

The ODE of the gradient flow is defined as

2 Data Structure for Half-Space Reporting

The half-space range reporting problem is an important problem in computational geometry, which is formally defined as following:

Let Tinit{\cal T}_{\mathsf{init}} denote the pre-processing time to build the data structure, Tquery{\cal T}_{\mathsf{query}} denote the time per query and Tupdate{\cal T}_{\mathsf{update}} time per update.

We use the data-structure proposed in [AEM92] to solve the half-space range reporting problem, which admits the interface summarized in Algorithm 1. Intuitively, the data-structure recursively partitions the set SS and organizes the points in a tree data-structure. Then for a given query (a,b)(a,b), all kk points of SS with sgn⁡(⟨a,x⟩−b)≥0\operatorname{sgn}(\langle a,x\rangle-b)\geq 0 are reported quickly. Note that the query (a,b)(a,b) here defines the half-space HH in Definition 3.5.

Adapted from [AEM92], the algorithm comes with the following complexity:

Part 1. Tquery(n,d,k)=Od(n1−1/⌊d/2⌋+k){\cal T}_{\mathsf{query}}(n,d,k)=O_{d}(n^{1-1/\lfloor d/2\rfloor}+k), amortized Tupdate=Od(log⁡2(n)){\cal T}_{\mathsf{update}}=O_{d}(\log^{2}(n)).

Part 2. Tquery(n,d,k)=Od(log⁡(n)+k){\cal T}_{\mathsf{query}}(n,d,k)=O_{d}(\log(n)+k), amortized Tupdate=Od(n⌊d/2⌋−1){\cal T}_{\mathsf{update}}=O_{d}(n^{\lfloor d/2\rfloor-1}).

We remark that Part 1 will be used in Theorem 6.1 and Part 2 will be used in Theorem 6.2.

3 Sparsity-based Characterizations

In this section, we consider the ReLU function with a nonzero threshold: σb(x)=max⁡{0,x−b}\sigma_{b}(x)=\max\{0,x-b\}, which is commonly seen in practise, and also has been considered in theoretical work [ZPD+20].

We first define the set of neurons that are firing at time tt.

We propose a new “sparsity” lemma in this work. It shows that σb\sigma_{b} gives the desired sparsity.

Let b>0b>0 be a tunable parameter. If we use the σb\sigma_{b} as the activation function, then after the initialization, with probability at least 1−n⋅exp⁡(−Ω(m⋅exp⁡(−b2/2)))1-n\cdot\exp(-\Omega(m\cdot\exp(-b^{2}/2))), it holds that for each input data xix_{i}, the number of activated neurons ki,0k_{i,0} is at most O(m⋅exp⁡(−b2/2))O(m\cdot\exp(-b^{2}/2)), where mm is the total number of neurons.

By the concentration of Gaussian distribution, the initial fire probability of a single neuron is

Hence, for the indicator variable 1r∈Si,fire⁡(0)\mathbf{1}_{r\in{\cal S}_{i,\operatorname{fire}}(0)}, we have

By standard concentration inequality (Lemma B.1),

where k0:=m⋅exp⁡(−b2/2)k_{0}:=m\cdot\exp(-b^{2}/2). If we choose t=k0t=k_{0}, then we have:

Then, by union bound over all i∈[n]i\in[n], we have that with high probability

the number of initial fire neurons for the sample xix_{i} is bounded by ki,0≤2m⋅exp⁡(−b2/2)k_{i,0}\leq 2m\cdot\exp(-b^{2}/2). ∎

The following remark gives an example of setting the threshold bb, and will be useful for showing the sublinear complexity in the next section.

If we choose b=0.4log⁡mb=\sqrt{0.4\log m} then k0=m4/5k_{0}=m^{4/5}. For t=m4/5t=m^{4/5}, Eq. (5) implies that

Training Neural Network with Half-Space Reporting Data Structure

In this section, we present two sublinear time algorithms for training over-parameterized neural networks. The first algorithm (Section 4.1) relies on building a high-dimensional search data-structure for the weights of neural network. The second algorithm (Section 4.2) is based on building a data structure for the input data points of the training set. Both of the algorithms use the HSR to quickly identify the fired neurons to avoid unnecessary calculation. The time complexity and the sketch of the proof are provided after each of the algorithms.

We first introduce the algorithm that preprocesses the weights wrw_{r} for r∈[m]r\in[m], which is commonly used in practice [CLP+21, CMF+20, KKL20]. Recall 2NN(m,b)\mathsf{2NN}(m,b) is f(W,x,a):=1m∑r=1marσb(⟨wr,x⟩)f(W,x,a):=\frac{1}{\sqrt{m}}\sum_{r=1}^{m}a_{r}\sigma_{b}(\langle w_{r},x\rangle). By constructing a HSR data-structure for wrw_{r}’s, we can quickly find the set of active neurons Si,fire⁡S_{i,\operatorname{fire}} for each of the training sample xix_{i}. See pseudo-code in Algorithm 2.

In the remaining part of this section, we focus on the time complexity analysis of Algorithm 2. The convergence proof will be given in Section 5.

Given nn data points in dd-dimensional space. Running gradient descent algorithm (Algorithm 2) on 2NN(m,b=0.4log⁡m)\mathsf{2NN}(m,b=\sqrt{0.4\log m}) (Definition 3.1) the expected cost per-iteration of the gradient descent algorithm is

The first term ∑i=1nT\textscQuery(m,d,ki,t)\sum_{i=1}^{n}{\cal T}_{\textsc{Query}}(m,d,k_{i,t}) corresponds to the running time of querying the active neuron set Si,fire⁡(t)S_{i,\operatorname{fire}}(t) for all training samples i∈[n]i\in[n]. With the first result in Corollary 3.6, the complexity is bounded by O~(m1−Θ(1/d)nd)\widetilde{O}(m^{1-\Theta(1/d)}nd).

The second term (T\textscDelete+T\textscInsert)⋅∣∪i∈[n]Si,fire⁡(t)∣({\cal T}_{\textsc{Delete}}+{\cal T}_{\textsc{Insert}})\cdot|\cup_{i\in[n]}S_{i,\operatorname{fire}}(t)| corresponds to updating wrw_{r} in the high-dimensional search data-structure (Lines 9 and 10). Again with the first result in Corollary 3.6, we have T\textscDelete+T\textscInsert=O(log⁡2m){\cal T}_{\textsc{Delete}}+{\cal T}_{\textsc{Insert}}=O(\log^{2}m). Combining with the fact that ∣∪i∈[n]Si,fire⁡(t)∣≤∣∪i∈[n]Si,fire⁡(0)∣≤O(nm4/5)|\cup_{i\in[n]}S_{i,\operatorname{fire}}(t)|\leq|\cup_{i\in[n]}S_{i,\operatorname{fire}}(0)|\leq O(nm^{4/5}), the second term is bounded by O(nm4/5log⁡2m)O(nm^{4/5}\log^{2}m).

The third term is the time complexity of gradient calculation restricted to the set Si,fire⁡(t){\cal S}_{i,\operatorname{fire}}(t). With the bound on ∑i∈[n]ki,t\sum_{i\in[n]}k_{i,t} (Lemma C.10), we have d∑i∈[n]ki,t≤O(m4/5nd)d\sum_{i\in[n]}k_{i,t}\leq O(m^{4/5}nd).

Putting them together completes the proof. ∎

2 Data Preprocessing

While the weights preprcessing algorithm is inspired by the common practise, the dual relationship between the input xix_{i} and model weights wrw_{r} inspires us to preprocess the dataset before training (i.e., building HSR data-structure for xix_{i}). This largely improves the per-iteration complexity and avoids the frequent updates of the data structure since the training data is fixed. More importantly, once the training dataset is preprocessed, it can be reused for different models or tasks, thus one does not need to perform the expensive preprocessing for each training.

The corresponding pseudocode is presented in Algorithm 3. With xix_{i} preprocessed, we can query HSR with weights wrw_{r} and the result S~r,fire⁡\widetilde{S}_{r,\operatorname{fire}} is the set of training samples xix_{i} for which wrw_{r} fires for. Given S~r,fire⁡\widetilde{S}_{r,\operatorname{fire}} for r∈[m]r\in[m], we can easily reconstruct the set Si,fire⁡S_{i,\operatorname{fire}}, which is the set of neurons fired for sample xix_{i}. The forward and backward pass can then proceed similar to Algorithm 2.

At the end of each iteration, we will update S~r,fire⁡\widetilde{S}_{r,\operatorname{fire}} based on the new wrw_{r} estimation and update Si,fire⁡S_{i,\operatorname{fire}} accordingly. For Algorithm 3, the HSR data-structure is static for the entire training process. This is the main difference from Algorithm 2, where the HSR needs to be updated every time step to account for the changing weights wrw_{r}.

We defer the convergence analysis to Section 5 and focus on the time complexity analysis of Algorithm 2 in the rest of this section. We consider dd being a constant for the rest of this subsection.

Given nn data points in dd-dimensional space. Running gradient descent algorithm (Algorithm 2) on 2NN(m,b=0.4log⁡m)\mathsf{2NN}(m,b=\sqrt{0.4\log m}) (Definition 3.1), the expected per-iteration running time of initializing S~r,fire⁡,Si,fire⁡\widetilde{S}_{r,\operatorname{fire}},S_{i,\operatorname{fire}} for r∈[m],i∈[n]r\in[m],i\in[n] is O(mlog⁡n+m4/5n).O(m\log n+m^{4/5}n). The cost per iteration of the training algorithm is O(m4/5nlog⁡n).O(m^{4/5}n\log n).

We analyze the initialization and training parts separately.

In Lines 4 and 5, the sets S~r,fire⁡,Si,fire⁡\widetilde{S}_{r,\operatorname{fire}},S_{i,\operatorname{fire}} for r∈[m],i∈[n]r\in[m],i\in[n] are initialized. For each r∈[m]r\in[m], we need to query the data structure the set of data points xx’s such that σb(wr(0)⊤x)>0\sigma_{b}(w_{r}(0)^{\top}x)>0. Hence, the running time of this step is

where the second step follows from ∑r=1mk~r,0=∑i=1nki,0\sum_{r=1}^{m}\widetilde{k}_{r,0}=\sum_{i=1}^{n}k_{i,0}.

Training

Consider training the neural network for TT steps. For each step, first notice that the forward and backward computation parts (Line 7 - 9) are the same as previous algorithm. The time complexity is O(m4/5nlog⁡n)O(m^{4/5}n\log n).

We next show that maintaining S~r,fire⁡,r∈[m]\widetilde{S}_{r,\operatorname{fire}},r\in[m] and Si,fire⁡,i∈[n]S_{i,\operatorname{fire}},i\in[n] (Line 10 - 14) takes O(m4/5nlog⁡n)O(m^{4/5}n\log n) time. For each fired neuron r∈[m]r\in[m], we first remove the indices of data in the sets Si,fireS_{i,\mathsf{fire}}, which takes time

Then, we find the new set of xx’s such that σb(⟨wr(t+1),x⟩)>0\sigma_{b}(\langle w_{r}(t+1),x\rangle)>0 by querying the half-space reporting data structure. The total running time for all fired neurons is

Then, we update the index sets Si,fire⁡S_{i,\operatorname{fire}} in time O(m4/5n)O(m^{4/5}n). Therefore, each training step takes O(m4/5nlog⁡n)O(m^{4/5}n\log n) time, which completes the proof. ∎

Convergence of Our Algorithm

We state the result of our training neural network algorithms (Lemma 5.2) can converge in certain steps. An important component in our proof is to find out a lower bound on minimum eigenvalue of the continuous Hessian matrix λmin⁡(Hcts⁡)\lambda_{\min}(H^{\operatorname{cts}}). It turns out to be an anti-concentration problem of the Gaussian random matrix. In [OS20], they gave a lower bound on λmin⁡(Hcts⁡)\lambda_{\min}(H^{\operatorname{cts}}) for ReLU function with b=0b=0, assuming the input data are separable. One of our major technical contribution is generalizing it to arbitrary b≥0b\geq 0.

With proposition 5.1, we are ready to show the convergence rate of training an over-parameterized neural network with shifted ReLU function.

Suppose input data-points are δ\delta-separable, i.e., δ:=min⁡i≠j{∥xi−xj∥2,∥xi+xj∥2}\delta:=\min_{i\neq j}\{\|x_{i}-x_{j}\|_{2},\|x_{i}+x_{j}\|_{2}\}. Let m=poly⁡(n,1/δ,log⁡(n/ρ))m=\operatorname{poly}(n,1/\delta,\log(n/\rho)) and η=O(λ/n2)\eta=O(\lambda/n^{2}). Let b=Θ(log⁡m)b=\Theta(\sqrt{\log m}). Then

Note that the randomness is over initialization. Eventually, we choose T=λ−2n2log⁡(n/ϵ)T=\lambda^{-2}n^{2}\log(n/\epsilon) where ϵ\epsilon is the final accuracy.

This result shows that despite the shifted ReLU and sparsely activated neurons, we can still retain the linear convergence. Combined with the results on per-step complexity in the previous section, it gives our main theoretical results of training deep learning models with sublinear time complexity (Theorem 6.1 and Theorem 6.2).

Main Classical Results

We present two theorems (under classical computation model) of our work, showing the sublinear running time and linear convergence rate of our two algorithms. We leave the quantum application into Appendix G. The first algorithm is relying on building a high-dimensional geometric search data-structure for the weights of a neural network.

Given nn data points in dd-dimensional space. We preprocess the initialization weights of the neural network. Running gradient descent algorithm (Algorithm 2) on a two-layer, mm-width, over-parameterized ReLU neural network will minimize the training loss to zero, and the expected running time of gradient descent algorithm (per iteration) is

The second algorithm is based on building a data structure for the input data points of the training set. Our second algorithm can further reduce the cost per iteration from m1−1/dm^{1-1/d} to truly sublinear in mm, e.g. m4/5m^{4/5}.

Given nn data points in dd-dimensional space. We preprocess all the data points. Running gradient descent algorithm (Algorithm 3) on a two-layer, mm-width, over-parameterized ReLU neural network will minimize the training loss to zero, and the expected running time of gradient descent algorithm (per iteration) is

Discussion and Limitations

In this paper, we propose two sublinear algorithms to train neural networks. By preprocessing the weights of the neuron networks or preprocessing the training data, we rigorously prove that it is possible to train a neuron network with sublinear complexity, which overcomes the Ω(mnd)\Omega(mnd) barrier in classical training methods. Our results also offer theoretical insights for many previously established fast training methods.

Our algorithm is intuitively related to the lottery tickets hypothesis [FC18]. However, our theoretical results can not be applied to explain lottery tickets immediately for two reasons: 1) the lottery ticket hypothesis focuses on pruning weights; while our results identify the important neurons. 2) the lottery ticket hypothesis identifies the weights that need to be pruned after training (by examining their magnitude), while our algorithms accelerate the training via preprocessing. It would be interesting to see how our theory can be extended to the lottery ticket hypothesis.

One limitation of our work is that the current analysis framework does not provide a convergence guarantee for combining LSH with gradient descent, which is commonly seen in many empirical works. Our proof breaks as LSH might miss important neurons which potentially ruins the convergence analysis. Instead, we refer to the HSR data structure, which provides a stronger theoretical guarantee of successfully finding all fired neurons.

References

Roadmap.

In Section A, we present our main algorithms. In Section B, we provide some preliminaries. In Section C, we provide sparsity analysis. We show convergence analysis in Section D. In Section E, we show how to combine the sparsity, convergence, running time all together. In Section F, we show correlation between sparsity and spectral gap of Hessian in neural tangent kernel. In Section G, we discuss how to generalize our result to quantum setting.

Appendix A Complete Algorithms

In this section, we present three algorithms (Alg. 4, Alg. 5 and Alg. 6) which are the complete version of Alg. 1, Alg. 2 and Alg. 3.

Appendix B Preliminaries

B.1 Probabilities

Let Z∼N(0,σ2)Z\sim{\mathcal{N}}(0,\sigma^{2}). Then, for t>0t>0,

B.2 Half-space reporting data structures

The time complexity of HSR data structure is:

Let dd be a fixed constant. Let tt be a parameter between nn and n⌊d/2⌋n^{\lfloor d/2\rfloor}. There is a dynamic data structure for half-space reporting that uses Od,ϵ(t1+ϵ)O_{d,\epsilon}(t^{1+\epsilon}) space and pre-processing time, Od,ϵ(nt1/⌊d/2⌋log⁡n+k)O_{d,\epsilon}(\frac{n}{t^{1/\lfloor d/2\rfloor}}\log n+k) time per query where kk is the output size and ϵ>0\epsilon>0 is any fixed constant, and Od,ϵ(t1+ϵ/n)O_{d,\epsilon}(t^{1+\epsilon}/n) amortized update time.

Part 1.Tinit(n,d)=Od(nlog⁡n){\cal T}_{\mathsf{init}}(n,d)=O_{d}(n\log n), Tquery(n,d,k)=Od,ϵ(n1−1/⌊d/2⌋+ϵ+k){\cal T}_{\mathsf{query}}(n,d,k)=O_{d,\epsilon}(n^{1-1/\lfloor d/2\rfloor+\epsilon}+k), amortized Tupdate=Od,ϵ(log⁡2(n)){\cal T}_{\mathsf{update}}=O_{d,\epsilon}(\log^{2}(n)).

Part 2.Tinit(n,d)=Od,ϵ(n⌊d/2⌋+ϵ){\cal T}_{\mathsf{init}}(n,d)=O_{d,\epsilon}(n^{\lfloor d/2\rfloor+\epsilon}), Tquery(n,d,k)=Od,ϵ(log⁡(n)+k){\cal T}_{\mathsf{query}}(n,d,k)=O_{d,\epsilon}(\log(n)+k), amortized Tupdate=Od,ϵ(n⌊d/2⌋−1+ϵ){\cal T}_{\mathsf{update}}=O_{d,\epsilon}(n^{\lfloor d/2\rfloor-1+\epsilon}).

B.3 Basic algebras

Appendix C Sparsity Analysis

In [DZPS19, SY19], they proved the following lemma for b=0b=0. Here, we provide a more general statement for any b≥0b\geq 0.

For any shift parameter b≥0b\geq 0, we define continuous version of shifted NTK Hcts⁡H^{\operatorname{cts}} and discrete version of shifted NTK Hdis⁡H^{\operatorname{dis}} as:

We define λ:=λmin⁡(Hcts⁡)\lambda:=\lambda_{\min}(H^{\operatorname{cts}}).

Let m=Ω(λ−1nlog⁡(n/ρ))m=\Omega(\lambda^{-1}n\log(n/\rho)) be number of samples of Hdis⁡H^{\operatorname{dis}}, then

We will use the matrix Chernoff bound (Theorem B.4) to provide a lower bound on the least eigenvalue of discrete version of shifted NTK Hdis⁡H^{\operatorname{dis}}.

Hence, Hr⪰0H_{r}\succeq 0. We need to upper-bound ∥Hr∥\|H_{r}\|. Naively, we have

since for each entry at (i,j)∈[n]×[n](i,j)\in[n]\times[n],

Hence, by matrix Chernoff bound (Theorem B.4) and choosing choose m=Ω(λ−1n⋅log⁡(n/ρ))m=\Omega(\lambda^{-1}n\cdot\log(n/\rho)), we can show

C.2 Handling Hessian if perturbing weight

We present a tool which is inspired by a list of previous work [DZPS19, SY19].

Part 1, ∥H(W~)−H(W)∥F≤n⋅min⁡{c⋅exp⁡(−b2/2),3R}\|H(\widetilde{W})-H(W)\|_{F}\leq n\cdot\min\{c\cdot\exp(-b^{2}/2),3R\} holds with probability at least 1−n2⋅exp⁡(−m⋅min⁡{c′⋅exp⁡(−b2/2),R/10})1-n^{2}\cdot\exp(-m\cdot\min\{c^{\prime}\cdot\exp(-b^{2}/2),R/10\}).

Part 2, λmin⁡(H(W))≥34λ−n⋅min⁡{c⋅exp⁡(−b2/2),3R}\lambda_{\min}(H(W))\geq\frac{3}{4}\lambda-n\cdot\min\{c\cdot\exp(-b^{2}/2),3R\} holds with probability at least 1−n2⋅exp⁡(−m⋅min⁡{c′⋅exp⁡(−b2/2),R/10})−ρ1-n^{2}\cdot\exp(-m\cdot\min\{c^{\prime}\cdot\exp(-b^{2}/2),R/10\})-\rho.

where the first step follows from definition of Frobenius norm, the last third step follows from by defining

For simplicity, we use srs_{r} to sr,i,j,bs_{r,i,j,b} (note that we fixed (i,j)(i,j) and bb).

Note that event Ai,rA_{i,r} happens iff ∣wr⊤xi−b∣≤R|w_{r}^{\top}x_{i}-b|\leq R happens.

Prior work [DZPS19, SY19] only one way to bound Pr⁡[Ai,r]\Pr[A_{i,r}]. We present two ways of arguing the upper bound on Pr⁡[Ai,r]\Pr[A_{i,r}]. One is anti-concentration, and the other is concentration.

where the last step follows from R<1/bR<1/b and c1≥exp⁡(1−R2/2)c_{1}\geq\exp(1-R^{2}/2) is a constant.

If the event ¬Ai,r\neg A_{i,r} happens and the event ¬Aj,r\neg A_{j,r} happens, then we have

If the event Ai,rA_{i,r} happens or the event Aj,rA_{j,r} happens, then we obtain

Define s‾=1m∑r=1msr\overline{s}=\frac{1}{m}\sum_{r=1}^{m}s_{r}. Thus, we are able to use Lemma B.1,

Define s‾=1m∑r=1msr\overline{s}=\frac{1}{m}\sum_{r=1}^{m}s_{r}. Thus, it gives

where c2:=2c1,c3:=38c1c_{2}:=2c_{1},c_{3}:=\frac{3}{8}c_{1} are some constants.

Define s‾=1m∑r=1msr\overline{s}=\frac{1}{m}\sum_{r=1}^{m}s_{r}. By Lemma B.1,

For the second part, by Lemma C.2, Pr⁡[λmin⁡(H(W~))≥0.75⋅λ]≥1−ρ\Pr[\lambda_{\min}(H(\widetilde{W}))\geq 0.75\cdot\lambda]\geq 1-\rho. Hence,

which happens with probability 1−n2⋅exp⁡(−m⋅min⁡{c3⋅exp⁡(−b2/2),R/10})−ρ1-n^{2}\cdot\exp(-m\cdot\min\{c_{3}\cdot\exp(-b^{2}/2),R/10\})-\rho by the union bound. ∎

C.3 Total movement of weights

For t≥0t\geq 0, let H(t)H(t) be an n×nn\times n matrix with (i,j)(i,j)-th entry:

We follow the standard notation Dcts⁡D_{\operatorname{cts}} in Lemma 3.5 in [SY19].

We state a tool from previous work [DZPS19, SY19] (more specifically, Lemma 3.4 in [DZPS19], Lemma 3.6 in [SY19]). Since adding the shift parameter bb to NTK doesn’t affect the proof of the following lemma, thus we don’t provide a proof and refer the readers to prior work.

∥W(t)−W(0)∥∞,2≤Dcts⁡\|W(t)-W(0)\|_{\infty,2}\leq D_{\operatorname{cts}},

C.4 Bounded gradient

The proof of Lemma 3.6 in [SY19] implicitly implies the following basic property of gradient.

C.5 Upper bound on the movement of weights per iteration

The following Claim is quite standard in the literature, we omitt the details.

C.6 Bounding the number of fired neuron per iteration

We define the set of neurons that are flipping at time tt:

Over all the iterations of training algorithm, there are some neurons that never flip states. We provide a mathematical formulation of that set,

For each i∈[n]i\in[n], let Si⊂[m]S_{i}\subset[m] denote the set of neurons that are never flipped during the entire training process,

In Lemma 3.8, we already show that ki,0=O(m⋅exp⁡(−b2/2))k_{i,0}=O(m\cdot\exp(-b^{2}/2)) for all i∈[n]i\in[n] with high probability. We can show that it also holds for t>0t>0.

Let b≥0b\geq 0 be a parameter, and let σb(x)=max⁡{x,b}\sigma_{b}(x)=\max\{x,b\} be the activation function. For each i∈[n],t∈[T]i\in[n],t\in[T], ki,tk_{i,t} is the number of activated neurons at the tt-th iteration. For 0<t≤T0<t\leq T, with probability at least 1−n⋅exp⁡(−Ω(m)⋅min⁡{R,exp⁡(−b2/2)})1-n\cdot\exp\left(-\Omega(m)\cdot\min\{R,\exp(-b^{2}/2)\}\right), ki,tk_{i,t} is at most O(mexp⁡(−b2/2))O(m\exp(-b^{2}/2)) for all i∈[n]i\in[n].

The base case of t=0t=0 is shown by Lemma 3.8 that ki,0=O(m⋅exp⁡(−b2/2))k_{i,0}=O(m\cdot\exp(-b^{2}/2)) for all i∈[n]i\in[n] with probability at least 1−nexp⁡(−Ω(m⋅exp⁡(−b2/2)))1-n\exp(-\Omega(m\cdot\exp(-b^{2}/2))).

Assume that the statement holds for 0,…,t−10,\dots,t-1. By Claim C.7, we know ∀k<t\forall k<t,

If we take t:=mexp⁡(−b2/2)t:=m\exp(-b^{2}/2), we have that

By a union bound for i∈[n]i\in[n], we obtain with probability

the number of activated neurons for xix_{i} at the tt-th iteration of the algorithm is

where the last step follows from ki,0=O(mexp⁡(−b2/2))k_{i,0}=O(m\exp(-b^{2}/2)) by Lemma 3.8.

The Lemma is then proved for all t=0,…,Tt=0,\dots,T. ∎

Let R≤1/bR\leq 1/b. For i∈[n]i\in[n], let SiS_{i} be the set defined by Eq. (6). Part 1. For r∈[m]r\in[m], r∉Sir\notin S_{i} if and only if ∣⟨wr(0),xi⟩−b∣<R|\langle w_{r}(0),x_{i}\rangle-b|<R. Part 2. If wr(0)∼N(0,Id)w_{r}(0)\sim\mathcal{N}(0,I_{d}), then

Part 1. We first note that r∉Si⊂[m]r\notin S_{i}\subset[m] is equivalent to the event that

Assume that ∥w−wr(0)∥2=R\|w-w_{r}(0)\|_{2}=R. Then, we can write w=wr(0)+R⋅vw=w_{r}(0)+R\cdot v with ∥v∥2=1\|v\|_{2}=1 and ⟨w,xi⟩=⟨wr(0),xi⟩+R⋅⟨v,xi⟩\langle w,x_{i}\rangle=\langle w_{r}(0),x_{i}\rangle+R\cdot\langle v,x_{i}\rangle.

Now, suppose there exists a ww such that 1⟨wr(0),xi⟩≥b≠1⟨w,xi⟩≥b\mathbf{1}_{\langle w_{r}(0),x_{i}\rangle\geq b}\neq\mathbf{1}_{\langle w,x_{i}\rangle\geq b}.

Since ∥xi∥2=1\|x_{i}\|_{2}=1 and ⟨v,xi⟩∈\langle v,x_{i}\rangle\in, we can see that the above conditions hold if and only if

In other words, r∉Sir\notin S_{i} if and only if ∣⟨wr(0),xi⟩−b∣<R|\langle w_{r}(0),x_{i}\rangle-b|<R.

where the last step follows from R<1/bR<1/b. ∎

Appendix D Convergence Analysis

The following Claim provides an upper bound for initialization. Prior work only shows it for b=0b=0, we generalize it to b≥0b\geq 0. The modification to the proof of previous Claim 3.10 in [SY19] is quite straightforward, thus we omit the details here.

Let b≥0b\geq 0 denote the NTK shifted parameter. Let parameter ρ∈(0,1)\rho\in(0,1) denote the failure probability. Then

D.2 Bounding progress per iteration

In previous work, [SY19] define HH and H⊥H^{\bot} only for b=0b=0. In this section, we generalize it to b≥0b\geq 0. Let us define two shifted matrices HH and H⊥H^{\bot}

Following the same proof as Claim 3.9 [SY19], we can show that the following Claim. The major difference between our claim and Claim 3.9 in [SY19] is, they only proved it for the case b=0b=0. We generalize it to b≥0b\geq 0. The proof is several basic algebra computations, we omit the details here.

The nontrivial parts in our analysis is how to bound B1,B2,B3B_{1},B_{2},B_{3} and B4B_{4} for the shifted cases (We will provide a proof later). Once we can bound all these terms, we can show the following result for one iteration of the algorithm:

D.3 Upper bound on the norm of dual Hessian

The proof of the following fact is similar to Fact C.1 in [SY19]. We generalize the b=0b=0 to b≥0b\geq 0. The same bound will hold as Fact C.1 in [SY19] if we replace 1wr(k)⊤xi≥0{\bf 1}_{w_{r}(k)^{\top}x_{i}\geq 0} by 1wr(k)⊤xi≥b{\bf 1}_{w_{r}(k)^{\top}x_{i}\geq b}. Thus, we omit the details here.

Let b≥0b\geq 0. Let shifted matrix H(k)⊥H(k)^{\bot} be defined as Eq. (8). For all k≥0k\geq 0, we have

D.4 Bounding the gradient improvement term

By Lemma C.2, there exists constants c,c′>0c,c^{\prime}>0 such that

If we have R≤λ12nR\leq\frac{\lambda}{12n} or b≥2⋅log⁡(4cn/λ)b\geq\sqrt{2\cdot\log(4cn/\lambda)}, then

D.5 Bounding the blowup by the dual Hessian term

Using Fact D.4, we have ∥H(k)⊥∥F≤nm2∑i=1n∣S‾i∣2\|H(k)^{\bot}\|_{F}\leq\frac{n}{m^{2}}\sum_{i=1}^{n}|\overline{S}_{i}|^{2}.

By Lemma C.10, ∀i∈{1,2,⋯ ,n}\forall i\in\{1,2,\cdots,n\}, it has

Hence, with probability at least 1−ρ01-\rho_{0}

D.6 Bounding the blowup by the flip-neurons term

D.7 Bounding the blowup by the prediction movement term

The proof of the following Claim is quite standard and simple in literature, see Claim 3.14 in [SY19]. We omit the details here.

D.8 Putting it all together

The goal of this section to combine all the convergence analysis together.

Let η=λ/(4n2)\eta={\lambda}/({4n^{2}}), R=λ/(12n)R=\lambda/(12n), let b∈[0,n]b\in[0,n], and

We know with probability ≥\geq 1−2n2⋅exp⁡(−Ω(m)⋅min⁡{R,exp⁡(−b2/2)})−ρ1-2n^{2}\cdot\exp(-\Omega(m)\cdot\min\{R,\exp(-b^{2}/2)\})-\rho,

Claim C.7 requires the following relationship between DD and RR,

By Claim D.1, we can upper bound the prediction error at the initialization,

Claim D.5 (where 0<c<e0<c<e is a constant) requires an upper bound on RR,Due to the relationship between bb and λ\lambda, we are not allowed to choose bb in an arbitrary function of λ\lambda. Thus, we should only expect to use RR to fix the problem.

Combing the lower bound and upper bound of RR, it implies the lower bound on mm in our Lemma statement.

which is dominated by the lower bound on mm in our lemma statement, thus we can ignore it.

However, by Theorem F.1, it will always hold for any b>0b>0.

where it follows from taking η:=λ/(4n2)\eta:={\lambda}/{(4n^{2})} and R=λ/(12n)R=\lambda/(12n).

Therefore, we can take the choice of the parameters m,b,Rm,b,R and Eqs. (11), (12) imply

Appendix E Combine

Let nn denote the number of points. Let dd denote the dimension of points. Let ρ∈(0,1/10)\rho\in(0,1/10) denote the failure probability. Let δ\delta be the separability of data points. For any parameter α∈(0,1]\alpha\in(0,1], we choose b=0.5(1−α)log⁡mb=\sqrt{0.5(1-\alpha)\log m}, if

If we preprocess the initial weights of the neural network, then we choose α=1−1/Θ(d)\alpha=1-1/\Theta(d) to get the desired running time.

If we preprocess the training data points, then we choose α\alpha to be an arbitrarily small constant to get the desired running time.

Since we know the upper bound of λ−1\lambda^{-1}, thus we need to choose

Let us choose b=0.5(1−α)log⁡mb=\sqrt{0.5(1-\alpha)\log m}, for any α∈(0,1]\alpha\in(0,1].

Given nn data points in dd-dimensional space. Running gradient descent algorithm on a two-layer ReLU (over-parameterized) neural network with mm neurons in the hidden layers is able to minimize the training loss to zero, let Tinit{\cal T}_{\mathsf{init}} denote the preprocessing time and Citer{\cal C}_{\mathsf{iter}} denote the cost per iteration of gradient descent algorithm.

If we preprocess the initial weights of the neural network (Algorithm 2), then

If we preprocess the training data points (Algorithm 3), then

Appendix F Bounds for the Spectral Gap with Data Separation

Let b≥0b\geq 0. Recall the continuous Hessian matrix Hcts⁡H^{\operatorname{cts}} is defined by

Let λ:=λmin⁡(Hcts⁡)\lambda:=\lambda_{\min}(H^{\operatorname{cts}}). Then, we have

Then, Hcts⁡H^{\operatorname{cts}} can be written as

where A∘BA\circ B denotes the Hadamard product between AA and BB.

By Claim B.7, and since ∥xi∥2=1\|x_{i}\|_{2}=1 for all i∈[n]i\in[n], we only need to show:

where the first step follows from Jensen’s inequality, the second step follows from Markov’s inequality, the last step follows from ∥a∥2=1\|a\|_{2}=1.

Since this is true for all aa, we find Eq. (17) with c12c2=1100c_{1}^{2}c_{2}=\frac{1}{100} by choosing c1=1/2,c2=1/25c_{1}=1/2,c_{2}=1/25 as described later.

For 0≤γ≤1/20\leq\gamma\leq 1/2, Gaussian small ball guarantees

Then, by Theorem 3.1 in [LS01] (Claim B.2), we have

Next, we argue that zi:=⟨Q‾g‾,xi⟩z_{i}:=\langle\overline{Q}\overline{g},x_{i}\rangle is small for all i≠1i\neq 1. For a fixed i≥2i\geq 2, observe that

Let τi,1:=⟨xi,x1⟩\tau_{i,1}:=\langle x_{i},x_{1}\rangle.

Then, from Gaussian anti-concentration bound (Lemma B.3) and variance bound on ziz_{i}, we have

Define E{\cal E} to be the following event:

where the last step follows from choosing γ:=δ4n∈[0,1/2]\gamma:=\frac{\delta}{4n}\in[0,1/2].

where the third step follows from τi,1=xi⊤x1\tau_{i,1}=x_{i}^{\top}x_{1}.

On the event E{\cal E}, by Claim F.2, we have that 1τi,1⋅g1+zi>b=1zi>(1−τi,1)b\mathbf{1}_{\tau_{i,1}\cdot g_{1}+z_{i}>b}=\mathbf{1}_{z_{i}>(1-\tau_{i,1})b}.

Furthermore, conditioned on E{\cal E}, g1,g‾g_{1},\overline{g} are independent as ziz_{i}’s are function of g‾\overline{g} alone. Hence, E{\cal E} can be split into two equally likely events that are symmetric with respect to g1g_{1} i.e. g1≥bg_{1}\geq b and g1<bg_{1}<b.

Now, using max⁡{∣a∣,∣b∣}≥∣a−b∣/2\max\{|a|,|b|\}\geq|a-b|/2, we find

where e1:=[10⋯0]⊤e_{1}:=\begin{bmatrix}1&0&\cdots&0\end{bmatrix}^{\top}, and the sixth step follows from ∥x1∥2=1\|x_{1}\|_{2}=1, the last step follows from the concentration of Gaussian distribution. In Line 5 and 6 of the above proof, ww is sampled from N(0,Id){\cal N}(0,I_{d}). ∎

If τi,1>0\tau_{i,1}>0, then ∣zi−(1−τi,1)b∣>+τi,1γ|z_{i}-(1-\tau_{i,1})b|>+\tau_{i,1}\gamma implies that 1τi,1⋅g1+zi>b=1zi>(1−τi,1)b\mathbf{1}_{\tau_{i,1}\cdot g_{1}+z_{i}>b}=\mathbf{1}_{z_{i}>(1-\tau_{i,1})b}.

If τi,1<0\tau_{i,1}<0, then ∣zi−(1−τi,1)b∣>−τi,1γ|z_{i}-(1-\tau_{i,1})b|>-\tau_{i,1}\gamma implies that 1τi,1⋅g1+zi>b=1zi>(1−τi,1)b\mathbf{1}_{\tau_{i,1}\cdot g_{1}+z_{i}>b}=\mathbf{1}_{z_{i}>(1-\tau_{i,1})b}.

That is, if ∣zi−(1−τi,1)b∣>∣τi,1∣γ|z_{i}-(1-\tau_{i,1})b|>|\tau_{i,1}|\gamma, then we have 1τi,1⋅g1+zi>b=1zi>(1−τi,1)b\mathbf{1}_{\tau_{i,1}\cdot g_{1}+z_{i}>b}=\mathbf{1}_{z_{i}>(1-\tau_{i,1})b}.

Case 1. We can assume τi,1>0\tau_{i,1}>0. By assumption, we know that g1∈[b−γ,b+γ]g_{1}\in[b-\gamma,b+\gamma].

According to the range of ziz_{i}, it implies zi>(1−τi,1)bz_{i}>(1-\tau_{i,1})b.

If zi>(1−τi,1)bz_{i}>(1-\tau_{i,1})b, then by the range of ziz_{i}, we have zi>(1−τi,1)b+τi,1γz_{i}>(1-\tau_{i,1})b+\tau_{i,1}\gamma.

Case 2. The τi,1<0\tau_{i,1}<0 case can be proved in a similar way. ∎

Appendix G Quantum Algorithm for Training Neural Network

In this section, we provide a quantum-classical hybrid approach to train neural networks with truly sub-quadratic time per iteration. The main observation is that the classical HSR data structure can be replaced with the Grover’s search algorithm in quantum.

We first state our main result in below, showing the running time of our quantum training algorithm:

Given nn data points in dd-dimensional space. Running gradient descent algorithm on a two-layer, mm-with, over-parameterized, and ReLU neural network will minimize the training loss to zero, let Citer{\cal C}_{\mathsf{iter}} denote the cost per iteration of gradient descent algorithm. Then, we have

by applying Grover’s search algorithm for the neurons (Algorithm 7) or the input data points (Algorithm 8).

We remark that previous works ([KLP19, AHKZ20]) on training classical neural networks use the quantum linear algebra approach, which achieves quantum speedup in the linear algebra operations in the training process. For example, [KLP19] used the block encoding technique to speedup the matrix multiplication in training convolutional neural network (CNN). [AHKZ20] used the quantum inner-product estimation to reduce each neuron’s computational cost. One drawback of this approach is that the quantum linear algebra computation incurs some non-negligible errors. Hence, extra efforts of error analysis are needed to guarantee that the intermediate errors will not affect the convergence of their algorithms.

Compared with the previous works, the only quantum component of our algorithm is Grover’s search. So, we do not need to worry about the quantum algorithm’s error in the training process. And we are able to use our fast training framework to exploit a sparse structure, which makes the Grover’s search algorithm run very fast, and further leads to a truly sub-quadratic training algorithm.

We also remark the difference between two algorithms in this quantum section the first algorithm runs Grover’s search for each data point to find the activated neurons, while the second one runs Grover’s search for each neuron to find the data points that make it activated. The advantage of Algorithm 8 is it uses less quantum resources, since its search space is of size O(n)O(n) and the first algorithm’s search space is of size O(m)O(m).

We first state a famous result about the quadratic quantum speedup for the unstructured search problem using Grover’s search algorithm.

Given access to the evaluation oracle for an unknown function f:[n]→{0,1}f:[n]\rightarrow\{0,1\} such that ∣f−1(1)∣=k|f^{-1}(1)|=k for some unknown number k≤nk\leq n, we can find all ii’s in f−1(1)f^{-1}(1) in O~(nk)\widetilde{O}(\sqrt{nk})-time quantumly.

For t=0,1,…,Tt=0,1,\dots,T, the time complexity of the tt-th iteration in Algorithm 7 is

Since ki,t≤mk_{i,t}\leq m for all i∈[n]i\in[n], the running time per iteration of Algorithm 7 is O~(ndm⋅max⁡i∈[n]ki,t)\widetilde{O}(nd\sqrt{m}\cdot\max_{i\in[n]}\sqrt{k_{i,t}}), which completes the proof of the lemma. ∎

The following lemma proves the running time of Algorithm 8.

For t=0,1,…,Tt=0,1,\dots,T, the time complexity of the tt-th iteration in Algorithm 8 is

The classical part is quite similar to Algorithm 6, which takes O(nd⋅max⁡i∈[n]ki,t)O(nd\cdot\max_{i\in[n]}k_{i,t})-time per iteration.

Therefore, the cost per iteration is O~(∑r∈[m](nk~r,t)1/2)\widetilde{O}(\sum_{r\in[m]}(n\widetilde{k}_{r,t})^{1/2}), and the lemma is then proved. ∎

Combining Lemma G.5 and Lemma G.6 proves the main result of this section:

In Section E, we prove that ki,t=m4/5k_{i,t}=m^{4/5} with high probability for all i∈[n]i\in[n] if we take b=0.4log⁡mb=\sqrt{0.4\log m}. Hence, by Lemma G.5, each iteration in Algorithm 7 takes

time in quantum. On the other hand, by Lemma G.6, each iteration in Algorithm 8 takes quantum time

where the second step is by ∑r∈[m]k~r,t=∑i∈[n]ki,t\sum_{r\in[m]}\widetilde{k}_{r,t}=\sum_{i\in[n]}k_{i,t}, which completes the proof of the corollary. ∎