Learning time-scales in two-layers neural networks

Raphaël Berthier, Andrea Montanari, Kangjie Zhou

Introduction

It is a recurring empirical observation that the training dynamics of neural networks exhibits a whole range of surprising behaviors:

Plateaus. Plotting the training and test error as a function of SGD steps, using either small stepsize or large batches to average out stochasticity, reveals striking patterns. These error curves display long plateaus where barely anything seems to be happening, which are followed by rapid drops (Saad and Solla, 1995; Yoshida and Okada, 2019; Power et al., 2022).

Time-scales separation. The time window for this rapid descent is much shorter than the time spent in the plateaus. Additionally, subsequent phases of learning take increasingly longer times (Ghorbani et al., 2020a; Barak et al., 2022).

Incremental learning. Models learnt in the first phases of learning appear to be simpler than in later phases. Among others, Arpit et al. (2017) demonstrated that easier examples in a dataset are learned earlier; Kalimeris et al. (2019) showed that models learnt in the first phase of training correlate well with linear models; Gissin et al. (2019) showed that, in many simplified models, the dynamics of gradient descent explores the solution space in an incremental order of complexity; Power et al. (2022) demonstrated that, in certain settings, a function that approximates well the target is only learnt past the point of overfitting.

Understanding these phenomena is not a matter of intellectual curiosity. In particular, incremental learning plays a key role in our understanding of generalization in deep learning. Indeed, in this scenario, stopping the learning at a certain time tt amounts to controlling the complexity of the model learnt. The notion of complexity corresponds to the order in which the space of models is explored.

While a number of groups have developed models to explain these phenomena, it is fair to say that a complete picture is still lacking. An exhaustive overview of these works is out of place here. We will outline three possible explanations that have been developed in the past, and provide more pointers in Section 3.

Several early works (Saad and Solla, 1995; Fukumizu and Amari, 2000; Wei et al., 2008) pointed out that the parametrization of multi-layer neural networks presents symmetries and degeneracies. For instance, the function represented by a multilayer perceptron is invariant under permutations of the neurons in the same layer. As a consequence, the population risk has multiple local minima connected through saddles or other singular sub-manifolds. Dynamics near these sub-manifolds naturally exhibits plateaus. Further, random or agnostic initializations typically place the network close to such submanifolds.

Theory #​2#2\#2: Linear networks.

Following the pioneering work of Baldi and Hornik (1989), a number of authors, most notably Saxe et al. (2013); Li et al. (2020), studied the behavior of deep neural networks with linear activations. While such networks can only represent linear functions, the training dynamics is highly non-linear. As demonstrated in Saxe et al. (2013), learning happens through stages that correspond to the singular value decomposition of the input-output covariance. Time scales are determined by the singular values.

Theory #​3#3\#3: Kernel regime.

Following an initial insight of Jacot et al. (2018), a number of groups proved that, for certain initializations, the training dynamics and model learnt by overparametrized neural networks is well approximated by certain linearly parametrized models. In the limit of very wide networks, the training dynamics of these models converges in turn to the training dynamics of kernel ridge(less) regression (KRR) with respect to a deterministic kernel (independent of the random initialization.) We refer to Bartlett et al. (2021) for an overview and pointers to this literature. Recently Ghosh et al. (2021) show that, in high dimension, the learning dynamics of KRR also exhibits plateaus and waterfalls, and learns functions of increasing complexity over a diverging sequence of timescales.

While each of these theories offers useful insights, it is important to realize that they do not agree on the basic mechanism that explains plateaus, time-scales separation, and incremental learning. In theory #1\#1, plateaus are associated to singular manifolds and high-dimensional saddles, while in theories #2\#2 and #3\#3 they are related to a hierarchy of singular values of a certain matrix. In #2\#2, the relevant singular values are the ones of the input-output covariance, and the fact that these singular values are well separated is postulated to be a property of the data distribution. In contrast, in #3\#3 the relevant singular values are the eigenvalues of the kernel operator, and hence completely independent of the output (the target function). In this case, eigenvalues which are very different are proved to exist under natural high-dimensional distributions.

Not only these theories propose different explanations, but they are also motivated by very different simplified models. Theory #1\#1 has been developed only for networks with a small number of hidden units. Theory #2\#2 only applies to networks with multiple output units, because otherwise the input-output covariance is a d×1d\times 1 matrix and hence has only one non-trivial singular value. Finally, theory #3\#3 applies under the conditions of the linear (a.k.a. lazy) regime, namely large overparametrization and suitable initialization (see, e.g., Bartlett et al. (2021)).

In order to better understand the origin of plateaus, time-scales separation, and incremental learning, we attempt a detailed analysis of gradient flow for two-layer neural networks. We consider a simple data-generation model, and propose a precise scenario for the behavior of learning dynamics. We do not assume any of the simplifying features of the theories described above: activations are non-linear; the number of hidden neurons is large; we place ourselves outside the linear (lazy) regime.

Our analysis is based on methods from dynamical systems theory: singular perturbation theory and matched asymptotic expansions. Unfortunately, we fall short of providing a general rigorous proof of the proposed scenario, but we can nevertheless prove it in several special cases and provide a heuristic argument supporting its generality.

The rest of the paper is organized as follows. Section 2 describes our data distribution, learning model, and the proposed scenario for the learning dynamics. We review further related work in Section 3. Section 4 describes the reduction of the gradient flow to a ‘mean field’ dynamics that will be the starting point of our analysis. Section 5 presents numerical evidence of the proposed learning scenario. Finally, Sections 6 to 7 present our analysis of the learning dynamics.

Notations.

In this paper, we use the classical asymptotic notations. The notations f(ε)=o(g(ε))f(\varepsilon)=o(g(\varepsilon)) or g(ε)=ω(f(ε))g(\varepsilon)=\omega(f(\varepsilon)) as ε→0\varepsilon\to 0 both denote that ∣f(ε)∣/∣g(ε)∣→0|f(\varepsilon)|/|g(\varepsilon)|\to 0 in the limit ε→0\varepsilon\to 0. The notations f(ε)=O(g(ε))f(\varepsilon)=O(g(\varepsilon)) or g(ε)=Ω(f(ε))g(\varepsilon)=\Omega(f(\varepsilon)) both denote that the ratio ∣f(ε)∣/∣g(ε)∣|f(\varepsilon)|/|g(\varepsilon)| remains upper bounded in the limit. The notation f(ε)=Θ(g(ε))f(\varepsilon)=\Theta(g(\varepsilon)) or f(ε)≍g(ε)f(\varepsilon)\asymp g(\varepsilon) denote that f(ε)=O(g(ε))f(\varepsilon)=O(g(\varepsilon)) and g(ε)=O(f(ε))g(\varepsilon)=O(f(\varepsilon)) both hold. Finally, f(ε)∼g(ε)f(\varepsilon)\sim g(\varepsilon) denotes that f(ε)/g(ε)→1f(\varepsilon)/g(\varepsilon)\to 1 in the limit.

Setting and standard learning scenario

where (a,u):=(a1,⋯ ,am,u1,⋯ ,um)(a,u):=(a_{1},\cdots,a_{m},u_{1},\cdots,u_{m}) collectively denotes all the model’s parameter. The factor 1/m1/m in the definition is relevant for the initialization and learning rate. We anticipate that we will initialize the aia_{i}’s to be of order one, which results in second layer coefficients ai/m=Θ(1/m)a_{i}/m=\Theta(1/m). This is often referred to as the ‘mean-field initialization’ and is known to drive learning process out of the linear or kernel regime, see e.g. (Mei et al., 2018b; Chizat and Bach, 2018; Ghorbani et al., 2020b; Yang and Hu, 2020; Abbe et al., 2022).

The bulk of our work will be devoted to the analysis of projected gradient flow in (ai,ui)1⩽i⩽m(a_{i},u_{i})_{1\leqslant i\leqslant m} on the population risk

In Section 7, we will bound the distance between stochastic gradient descent (SGD) and gradient flow in population risk. As a consequence, we will establish finite sample generalization guarantees for SGD learning.

Projected gradient flow with respect to the risk \mathscrsfsR(a,u)\mathscrsfs{R}(a,u) is defined by the following ordinary differential equations (ODEs):

It is useful to make a few remarks about the definition of gradient flow:

The overall scaling of time is arbitrary, and the matching to SGD steps will be carried out in Section 7. The factors mm on the right-hand side are introduced for mathematical convenience, since the partial derivatives are of order 1/m1/m.

The factor ε\varepsilon introduced in the flow of the aia_{i}’s reflects the fact that usually SGD is run with respect to the overall second-layer weights (ai/m)1≤i≤m(a_{i}/m)_{1\leq i\leq m}. This would correspond to taking ε=1/m\varepsilon=1/m. However, we will keep ε\varepsilon as a free parameter independent of mm, and study the evolution for small ε\varepsilon.

We assume the initialization to be random with i.i.d. components (ai,init,ui,init)(a_{i,\rm{init}},u_{i,\rm{init}}):

In order to describe the polynomial approximations learnt during the training more explicitly, we decompose φ\varphi and σ\sigma into normalized Hermite polynomials:

As we will see, the incremental learning behavior arises for small ε\varepsilon. By the law of large numbers (see below), the following almost sure limit exists (provided PA{\rm P}_{A} is square integrable)

We are now in position to describe the scenario that we will study in the rest of the paper.

We say that the standard learning scenario holds up to level LL for a certain target function φ\varphi, activation σ\sigma, and distribution PA{\rm P}_{A}, if the followings hold:

There exist constants c2,…,cL+1>0c_{2},\dots,c_{L+1}>0 such that the following asymptotic holds as ε→0\varepsilon\to 0, t→0t\to 0:

Figure 1 provides a cartoon illustration of the standard learning scenario.

A specific realization of our general setup is determined by the triple (σ,φ,PA)(\sigma,\varphi,{\rm P}_{A}), In the rest of the paper, we will provide evidence showing that the standard learning scenario holds in a number of cases. Nevertheless, we can also construct examples in which it does not hold:

If the first Hermite k+1k+1 coefficients of φ\varphi vanish, φ0=⋯=φk=0\varphi_{0}=\dots=\varphi_{k}=0, k≥1k\geq 1, then the standard scenario does not hold. (See Appendix D.2 for the proof.)

In fact, we expect the standard scenario might fail every time one or more of the coefficients φk\varphi_{k} vanish, for k≥1k\geq 1. Appendix D.3 provides some heuristic justification for this failure.

We can compare the standard learning scenario described here to the ones in earlier literature and described as theory #1\#1, #2\#2, #3\#3 in the introduction. There appears points of contact, but also important differences with both theory #1\#1 and #3\#3:

As in theory #1\#1, the plateaus and separation of time scales arise because the trajectory of gradient flow is approximated by a sequence of motions along submanifolds in the space of parameters (a,u)(a,u). Along the ll-th such submanifold f(x;a,u)f(x;a,u) is well-approximated by a degree-ll polynomial. Escaping each submanifold takes an increasingly longer time.

This is reminding of the motion between saddles investigated in earlier work (Saad and Solla, 1995; Fukumizu and Amari, 2000; Wei et al., 2008). However, unlike in earlier work, we will see that this applies to networks with a large (possibly diverging) number of hidden neurons. Also, we identify the subsequent phases of learning with the polynomial decomposition of Eq. (7).

As in theory #3\#3, subsequent phases of learning correspond to increasingly accurate polynomial approximations of the target function φ(⟨u∗,x⟩)\varphi(\langle u_{*},x\rangle). However, the underlying mechanism and time scales are completely different. In the linear regime, the different time scales emerge because of increasingly small eigenvalues of the neural tangent kernel. In that case, the time required to learn degree-ll polynomials is of order dld^{l} (Ghosh et al., 2021).

In contrast, in the standard learning scenario, polynomials of degree ll are learnt on a time scale of order one in dd (and only depending on the learning rate ε\varepsilon). This of course has important implications when approximating gradient flow by SGD. Within the linear regime, the sample size required to learn polynomial of order ll scales like dld^{l} (Ghosh et al., 2021), while in the standard scenario, it is only of order dd (see Section 7).

Further related work

As we mentioned in the introduction, plateaus and time scales in the learning dynamics of kernel models were analyzed by Ghosh et al. (2021). A sharp analysis for the related random features model was developed by Bodin and Macris (2021).

Our analysis builds upon the mean-field description of learning in two-layer neural networks, which was developed in a sequence of works, see, e.g., (Mei et al., 2018b; Rotskoff and Vanden-Eijnden, 2018; Chizat and Bach, 2018; Mei et al., 2019). In particular, we leverage the fact that, for the data distribution (1), the population risk function is invariant under rotations around the axis u∗u_{*}, and this allows for a dimensionality reduction in the mean field description. Similar symmetry argument were used by Mei et al. (2018b) and, more recently, by Abbe et al. (2022).

The single-index model can be learnt using simpler methods than large two-layer networks. Limiting ourselves to the case of gradient descent algorithms, Mei et al. (2018a) proved that gradient descent with respect to the non-convex empirical risk R^n(u):=n−1∑i=1n(yi−φ(u⊤xi))2\widehat{R}_{n}(u):=n^{-1}\sum_{i=1}^{n}(y_{i}-\varphi(u^{\top}x_{i}))^{2} converges to a near global optimum, provided φ\varphi is strictly increasing. Ben Arous et al. (2021) considered online SGD under more challenging learning scenarios and characterized the time (sample size) for ∣⟨u,u∗⟩∣|\langle u,u_{*}\rangle| to become significantly larger than for a random unit vector uu.

Learning in overparametrized two-layer networks under model (1) (or its variations) has been studied recently by several groups. In particular, Ba et al. (2022) considers a training procedure which runs a single step gradient descent followed by freezing the first layer and performing ridge regression with respect to the second layer. This scheme is amenable to a precise characterization of the generalization error. Bietti et al. (2022) consider a similar scheme in which a first phase of gradient descent is run to achieve positive correlation with the unknown direction u∗u_{*}. Damian et al. (2022) also consider a two-phases scheme, and prove consistency and excess risk bounds for a more general class of target functions whereby the first equation in (1) is replaced by

with k≪dk\ll d. In particular, near optimal error bounds are obtained under a non-degeneracy condition on ∇2φ\nabla^{2}\varphi.

A similar model was studied by Barak et al. (2022) that bounds the sample complexity by dO(k)d^{O(k)} for learning parities on kk bits using gradient descent with large batches (if k=O(1)k=O(1), Barak et al. (2022) require O(1)O(1) steps with batch size dO(k)d^{O(k)}).

Let us emphasize that our objective is quite different from these works. We do not allow ourselves deviations from standard SGD and try to derive a precise picture of the successive phases of learning (in particular, we do not consider two-stage schemes or layer-by-layer learning). On the other hand, we focus on a relatively simple model.

To clarify the difference, it is perhaps useful to rephrase our claims in terms of sample complexity. While previous works show that the target function can be learnt with O(d)O(d) samples, we claim that it is learnt by online SGD with test error rr from about C(r,ε)dC(r,\varepsilon)d samples and characterize the dependence of C(r,ε)C(r,\varepsilon) on rr for small ε\varepsilon. (Falling short of a proof in the general case.)

After posting an initial version of this paper, we became aware that Arnaboldi et al. (2023) independently derived equations similar to (14)-(18), (25), (119). There are technical differences, and hence we cannot apply their results directly. However, Section 4.1 and Appendix A.4 are analogous to their work.

The large-network, high-dimensional limit

The first step of our analysis is a reduction of the system of ODEs (4), (5), with dimension m(d+1)m(d+1) to a system of ODEs in 2m2m dimensions. We will achieve this reduction in two steps:

First we reduce to a system in m(m+3)/2m(m+3)/2 dimensions for the variables aia_{i}, ⟨ui,uj⟩\langle u_{i},u_{j}\rangle, ⟨ui,u∗⟩\langle u_{i},u_{*}\rangle. This reduction is exact and is quite standard.

We then show that the products ⟨ui,uj⟩\langle u_{i},u_{j}\rangle can be eliminated, with an error O(1/m)O(1/m). As further discussed below, the resulting dynamics could also be derived from the mean field theory of Mei et al. (2018b); Rotskoff and Vanden-Eijnden (2018); Chizat and Bach (2018); Mei et al. (2019) (with the required modifications for the constraints ∥ui∥=1\|u_{i}\|=1).

Note that the above identities follow from (O’Donnell, 2014, Proposition 11.31). Throughout this section, we will make the following assumptions.

The distribution of weights at initialization, PA{\rm P}_{A} is supported on [−M1,M1][-M_{1},M_{1}].

The activation function is bounded: ∥σ∥∞≤M2\left\|{\sigma}\right\|_{\infty}\leq M_{2}. Additionally, the functions VV and UU are bounded and of class C2C^{2}, with uniformly bounded first and second derivatives over s∈s\in. A sufficient condition for this is

Responses are bounded, i.e., ∥φ∥∞≤M3\|\varphi\|_{\infty}\leq M_{3}.

We hereby briefly explain the sufficiency of L2L^{2}-boundedness of derivatives of σ\sigma and φ\varphi as claimed in Assumption A2. Suppose for example that ∥σ′∥L2≤M2\left\|{\sigma^{\prime}}\right\|_{L^{2}}\leq M_{2} and ∥φ′∥L2≤M2\left\|{\varphi^{\prime}}\right\|_{L^{2}}\leq M_{2}, then we have

where (a)(a) follows from Gaussian integration by parts and (b)(b) follows from Cauchy-Schwarz inequality.

Our first statement establishes reduction (i)(i) mentioned above. The proof of this fact is presented in Appendix A.1.

Define si=⟨ui,u∗⟩s_{i}=\langle u_{i},u_{*}\rangle, rij=⟨ui,uj⟩r_{ij}=\langle u_{i},u_{j}\rangle for i,j=1,…,mi,j=1,\dots,m. Then, letting R=(rij)i,j≤mR=(r_{ij})_{i,j\leq m}, we have

If (a(t),u(t))(a(t),u(t)) solve the gradient flow ODEs (4)-(5) then (a(t),s(t),R(t))(a(t),s(t),R(t)) are the unique solution of the following set of ODEs (note that rii=1r_{ii}=1 identically)

This discussion immediately yields the following consequence.

Let (a(t),u(t))(a(t),u(t)) be the solution of the gradient flow ODEs (4), (5) with initialization (6), and let (a0(t),s0(t),R0(t))(a^{0}(t),s^{0}(t),R^{0}(t)) be the unique solution of Eqs. (15) to (18), with initialization ai0(0)=ai(0)a^{0}_{i}(0)=a_{i}(0), si0(0)=0s^{0}_{i}(0)=0, rij0(0)=0r^{0}_{ij}(0)=0 for i≠ji\neq j. Then, for any fixed TT (possibly dependent on mm but not on dd), the followings holds with probability at least 1−exp⁡(−C′m)1-\exp(-C^{\prime}m) over the i.i.d. initialization (ai(0),ui(0))i∈[m](a_{i}(0),u_{i}(0))_{i\in[m]}:

Here C,C′C,C^{\prime} are absolute constants and MM only depends on the MiM_{i}’s in Assumptions A1-A3.

In order to state the reduction (ii)(ii) outlined above, we define the mean field risk as

Further, we denote by {ai\mboxmf(t),si\mboxmf(t)}i=1m\{a^{\mbox{\tiny\rm mf}}_{i}(t),s^{\mbox{\tiny\rm mf}}_{i}(t)\}_{i=1}^{m} the solution to the following ODEs:

Note that (23) would be identical to (15)-(16) if we had rij=sisjr_{ij}=s_{i}s_{j}. A priori, this is not the case. However, the two systems of equations are close to each other for large mm as made precise by our next proposition, which formalizes reduction (ii)(ii).

Let (ai(t),si(t),rij(t))1≤i<j≤m(a_{i}(t),s_{i}(t),r_{ij}(t))_{1\leq i<j\leq m} be the unique solution of the ODEs (15)-(18) with initialization si(0)=0s_{i}(0)=0, rij(0)=0r_{ij}(0)=0 for all 1≤i≠j≤m1\leq i\neq j\leq m. Let (ai\mboxmf(t),si\mboxmf(t))i≤m(a^{\mbox{\tiny\rm mf}}_{i}(t),s^{\mbox{\tiny\rm mf}}_{i}(t))_{i\leq m} be the unique solution of the ODEs (23) with initialization si\mboxmf(0)=0s^{\mbox{\tiny\rm mf}}_{i}(0)=0, ai\mboxmf(0)=ai(0)a^{\mbox{\tiny\rm mf}}_{i}(0)=a_{i}(0) for all i≤mi\leq m.

If assumptions A1-A3 hold, then for any T<∞T<\infty there exists a constant

(with MM depending on the constants {Mi}1≤i≤3\{M_{i}\}_{1\leq i\leq 3} appearing in Assumptions A1-A3 only) such that:

The proof of this proposition is deferred to Appendix A.3. Now, combining the propositions and corollaries in this section, we deduce that with high probability over the i.i.d. initialization,

Consider the empirical distributions of the neurons:

with (ai(t),si(t))i≤m(a_{i}(t),s_{i}(t))_{i\leq m}, (ai\mboxmf(t),si\mboxmf(t))i≤m(a^{\mbox{\tiny\rm mf}}_{i}(t),s^{\mbox{\tiny\rm mf}}_{i}(t))_{i\leq m} as in the statement of Proposition 2, i.e., solving (respectively) Eqs. (15)-(18) and Eq. (23) with initial conditions as given there.

Then, it is immediate to show that ρt\rho_{t} solves (in weak sense) the following continuity partial differential equation (PDE) (we refer to Ambrosio et al. (2005); Santambrogio (2015) for the definition of weak solutions and basic properties, and Appendix A.4 for a short derivation.)

where Ψ=(Ψa,Ψs)\Psi=(\Psi_{a},\Psi_{s}) is given by

which is the obvious extension of \mathscrsfsR\mboxmf(a,s)\mathscrsfs{R}_{\mbox{\tiny\rm mf}}(a,s) of Eq. (22) to general probability distributions. Proposition 2 implies that for any T<∞T<\infty, and under the above initial conditions,

Starting with Mei et al. (2018b); Chizat and Bach (2018); Rotskoff and Vanden-Eijnden (2018), several authors used continuity PDEs of the form (28) to study the learning dynamics of two-layer neural networks. Following the physics tradition, this is referred to as the ‘mean-field theory’ of two-layer neural networks. Appendix A.5 sketches an alternative approach to prove bounds of the form (25), (34) using the results of Mei et al. (2018b, 2019). The present derivation has the advantages of yielding a sharper bound and of being self-contained.

2 A general formulation

As mentioned above, the system of ODEs in Eq. (23) is a special case of the Wasserstein gradient flow of Eq. (28) whereby we set ρ0=m−1∑i=1mδ(ai\mboxmf(0),si\mboxmf(0))\rho_{0}=m^{-1}\sum_{i=1}^{m}\delta_{(a_{i}^{\mbox{\tiny\rm mf}}(0),s_{i}^{\mbox{\tiny\rm mf}}(0))}. In order to study the solutions of Eq. (28) (hence Eq. (23)) we adopt the following framework. Let (Ω,ρ)(\Omega,\rho) denote a probability space. Let a=a(ω,t)a=a(\omega,t) and s=s(ω,t)s=s(\omega,t) (ω∈Ω\omega\in\Omega, t⩾0t\geqslant 0) be two measurable functions satisfying (dropping dependencies in tt below)

Numerical solution

In Figure 2, we present the result of an Euler discretization of Eqs. (23) where φ\varphi is a degree-22 polynomial and σ\sigma is the ReLU activation: σ(s)=max⁡(s,0)\sigma(s)=\max(s,0),

These plots clearly display two of the features emphasized in the introduction: (i)(i) plateaus separated by periods of rapid improvement of the risk; (ii)(ii) increasingly long timescales (notice the logarithmic time axis in the second and third row).

In order to examine the incremental learning structure, we rewrite the risk \mathscrsfsR\mboxmf\mathscrsfs{R}_{\mbox{\tiny\rm mf}} of Eq. (22) by decomposing φ\varphi and σ\sigma in the basis of Hermite polynomials

We observe that, for small ε\varepsilon, the Hermite coefficients of φ\varphi are learned sequentially, in the order of their degree. When ε\varepsilon is sufficiently small (right plots), this incremental learning happens in well separated phases. The plateaus and waterfalls in the plots of \mathscrsfsR\mboxmf\mathscrsfs{R}_{\mbox{\tiny\rm mf}} correspond to the network learning increasingly higher degree polynomials.

In Figure 3 we plot the evolution of the values of the aia_{i} and sis_{i}, for i∈{1,…,m}i\in\{1,\dots,m\}. We observe that the order of magnitude of the aia_{i}’s and the sis_{i}’s increases when passing through the different phases of the incremental learning process.

Altogether, the results of Figures 2 and 3 are consistent with the standard learning scenario up to level L=2L=2 as per Definition 1. While we conjecture that incremental learning also occurs for higher-order polynomials, we found this hard to observe in numerical simulations.

First, as predicted in Definition 1, the times at which the components are learned are closer on a logarithmic scale as the degree increases. It is therefore increasingly difficult to observe time scales corresponding to higher degrees.

Second, we expect there to be a choice of the initialization (ai,init,ui,init)i∈[m](a_{i,\rm{init}},u_{i,\rm{init}})_{i\in[m]}, activation and target function, for which not all the components of φ\varphi are actually learnt. We observed empirically that this happens easily for small mm.

Timescales hierarchy in the gradient flow dynamics

We are interested in the behavior of the solution of the ODEs (35), initialized from s(ω,0)=0s(\omega,0)=0 for all ω\omega (as per Proposition 2). The standard learning scenario of Definition 1 concerns the behavior of solutions for ε→0\varepsilon\to 0. This type of questions can be addressed within the theory of dynamical systems using singular perturbation theory (Holmes, 2013) (‘singular’ refers to the fact that ε\varepsilon multiplies one of the highest-order derivatives).

As a side remark, we note that the system (35) can be seen as a slow-fast dynamical system, where the a(ω)a(\omega)’s are the fast variables and the s(ω)s(\omega)’s are the slow variables (Berglund, 2001). Formally, the time derivative of the a(ω)a(\omega)’s is multiplied by a factor (1/ε)(1/\varepsilon). From a dynamical systems perspective, the present case is made complicated because of a bifurcation when the s(ω)s(\omega)’s become non-zero.

The standard learning scenario provides a detailed description of this bifurcation. We will motivate this scenario using a classical, but non-rigorous, technique of singular perturbation theory, called the matched asymptotic expansion (Holmes, 2013, Chapter 2). This technique decomposes the approximation of the solution in several time scales on which a regular approximation holds. These time scales are traditionally called layers in the literature; however, we avoid this terminology due to the potential confusion with the layers of the neural network.

We will work mainly using the Hermite representation of the dynamical ODEs (35), which we write down for the reader’s convenience:

Sections 6.1-6.3 respectively describe the first three time scales of the matched asymptotic expansion of (39). This gives, for each time scale, an approximation of the a(ω)a(\omega), s(ω)s(\omega). In Appendix B.2, we detail how these sections induce an evolution of the risk alternating plateaus and rapid decreases, and support the standing learning scenario of Definition 1. Finally, in Section 6.4, we conjecture the behavior on longer time scales.

We denote ainit(ω)=a(ω,0)a_{\text{init}}(\omega)=a(\omega,0) and thus a⊥,inita_{\perp,\text{\rm{init}}} is the orthogonal projection of ainita_{\text{init}} on \mathds1⊥\mathds{1}^{\perp}.

1 First time scale: constant component

We define a “fast” time variable t1=t/εt_{1}=t/\varepsilon and replace it in Eq. (39). We expand the solutions a(ω)a(\omega) and s(ω)s(\omega) in powers of ε\varepsilon:

where a(0)(ω),a(1)(ω),a(2)(ω),…,s(0)(ω),s(1)(ω),s(2)(ω),…a^{(0)}(\omega),a^{(1)}(\omega),a^{(2)}(\omega),\dots,s^{(0)}(\omega),s^{(1)}(\omega),s^{(2)}(\omega),\dots are implicitly functions of t1t_{1}. They are initialized at

to be consistent with the initial condition a(ω,t1=0)=a(ω,t=0)=ainit(ω)a(\omega,t_{1}=0)=a(\omega,t=0)=a_{\rm{init}}(\omega) and s(ω,t1=0)=s(ω,t=0)=0s(\omega,t_{1}=0)=s(\omega,t=0)=0.

The basic assumption of matched asymptotic expansions is that terms of the same order in ε\varepsilon can be identified (with some limitations that we develop below). For now, let us identify terms of order 1=ε01=\varepsilon^{0}:

From (51) and (43), we have s(0)(ω)=0s^{(0)}(\omega)=0: time t1=O(1)⇔t=O(ε)t_{1}=O(1)\Leftrightarrow t=O(\varepsilon) is too short for the s(ω)s(\omega) to be of order 11.

Substituting s(0)(ω)=0s^{(0)}(\omega)=0 in (50), we obtain

which gives after integration (using (42)):

At this point, we have determined a(0)(ω)a^{(0)}(\omega) and s(0)(ω)s^{(0)}(\omega), and thus a(ω)=a(0)(ω)+O(ε)a(\omega)=a^{(0)}(\omega)+O(\varepsilon) and s(ω)=s(0)(ω)+O(ε)s(\omega)=s^{(0)}(\omega)+O(\varepsilon) up to a O(ε)O(\varepsilon) precision, which is sufficient to obtain a o(1)o(1)-approximation of the risk \mathscrsfsR\mboxmf,∗\mathscrsfs{R}_{\mbox{\tiny\rm mf},*} (see Section B.2). However, note that we could obtain more precise estimates by identifying higher-order terms in (44)-(49). For instance, identifying the O(ε)O(\varepsilon) terms in (47)-(49), we obtain ∂t1s(1)(ω)=a(0)(ω)σ1φ1\partial_{t_{1}}s^{(1)}(\omega)=a^{(0)}(\omega)\sigma_{1}\varphi_{1}. This shows that the s(ω)s(\omega) become non-zero, though only of order ε\varepsilon on the time scale t1≍1t_{1}\asymp 1; the inner-layer weights develop an infinitesimal correlation with the true direction u∗u_{*} thanks to the linear component of σ\sigma and φ\varphi.

The approximation constructed above should be considered as valid on the time scale t1≍1⇔t≍εt_{1}\asymp 1\Leftrightarrow t\asymp\varepsilon. The approximation breaks down when we reach a new time scale, at which the s(ω)s(\omega) are large enough for the a(ω)a(\omega) to be affected (at leading order) by the linear part of the functions. We detail the new time scale and its resolution in the next section.

2 Second time scale: linear component I

In this section, we seek a second, slower time scale, for which the behavior of the asymptotic expansion is different.

Consider t2=tεγt_{2}=\frac{t}{\varepsilon^{\gamma}}, where γ<1\gamma<1 is to be determined. We rewrite the system (39) using t2t_{2}, and expand the solutions a(ω)a(\omega) and s(ω)s(\omega):

(Since within the previous time scale we obtained s(ω)=O(ε)s(\omega)=O(\varepsilon), it is natural to assume s(0)(ω)=0s^{(0)}(\omega)=0.)

Similarly to what has been done in the previous time scale, we will substitute the expansions (54)-(55) in the equations (39) in order to compute the different terms in the expansion. However, this step also allows us to compute the exponents γ\gamma and δ\delta, that give respectively the new time scale and the size of the s(ω)s(\omega)’s.

Note that we should have proceeded similarly for the first time scale, by introducing a first time variable t1=tεγ′t_{1}=\frac{t}{\varepsilon^{\gamma^{\prime}}}, expanding a(ω),s(ω)a(\omega),s(\omega) in powers 1,εδ′,ε2δ′,…1,\varepsilon^{\delta^{\prime}},\varepsilon^{2\delta^{\prime}},\dots, and determining γ′\gamma^{\prime} and δ′\delta^{\prime} a posteriori. This would have led, indeed, to γ′=1\gamma^{\prime}=1 and δ′=1\delta^{\prime}=1. However, for simplicity, we preferred to fix these values that are natural a priori.

Finally, note that the expansions (40)-(41) and (54)-(55) are different, because they are valid on different time scales. In fact, the only coherence conditions that we require below is that the expansions match in a joint asymptotic where t1=tε→∞t_{1}=\frac{t}{\varepsilon}\to\infty and t2=tεγ→0t_{2}=\frac{t}{\varepsilon^{\gamma}}\to 0. We thus build different approximations for each one of the time scales, with some matching conditions; this justifies the name of matched asymptotic expansion.

We now return to our computations and substitute (54)-(55) in (39):

For the first time scale, we chose γ=δ=1\gamma=\delta=1, so that the terms of order εδ\varepsilon^{\delta} were negligible compared to ε1−γ∂t2a(0)(ω)\varepsilon^{1-\gamma}\partial_{t_{2}}a^{(0)}(\omega) in (56). This means that the linear components σ1,φ1\sigma_{1},\varphi_{1} of the functions had no effect on the a(ω)a(\omega) at leading order. We are now interested in a new time scale where ε1−γ∂t2a(0)(ω)\varepsilon^{1-\gamma}\partial_{t_{2}}a^{(0)}(\omega) and εδσ1φ1s(1)(ω)\varepsilon^{\delta}\sigma_{1}\varphi_{1}s^{(1)}(\omega) are of the same order, i.e., 1−γ=δ1-\gamma=\delta; then the linear components play a role in the dynamics.

Further, for s(1)(ω)s^{(1)}(\omega) to be non-zero, we need both sides of (58) to be of the same order, thus δ=γ\delta=\gamma. Putting together, this gives γ=δ=1/2\gamma=\delta=1/2.

Derivation of the ODEs for this time scale.

Let us summarize equations. For t2=tε\nicefrac12t_{2}=\frac{t}{\varepsilon^{\nicefrac{{1}}{{2}}}} and

First, we identify the terms of order 1=ε01=\varepsilon^{0}:

Second, we identify the terms of order ε\nicefrac12\varepsilon^{\nicefrac{{1}}{{2}}} in (59)-(61):

In (63), the first term of the right hand side depends on the unknown higher-order terms a(1)(ν)a^{(1)}(\nu); in fact, this is best interpreted as the Lagrange multiplier associated to the constraint (62). To eliminate this Lagrange multiplier, we use again the compact notations:

Matching.

The initialization of the ODEs (65)-(66) for the second time scale is determined by a classical procedure that matches with the previous time scale. In this paragraph, we denote a‾,s‾\underline{a},\underline{s} the approximation obtained in the first time scale (Section 6.1), and a‾,s‾\overline{a},\overline{s} the approximation in the second time scale, described above.

Consider an intermediate time scale t~=tεα\widetilde{t}=\frac{t}{\varepsilon^{\alpha}}, \nicefrac12<α<1\nicefrac{{1}}{{2}}<\alpha<1, and assume t~≍1\widetilde{t}\asymp 1 so that

In this intermediate regime, we want the approximations provided on the first and the second time scales to match: a‾(t~)\underline{a}(\widetilde{t}) and a‾(t~)\overline{a}(\widetilde{t}) (resp. s‾(t~)\underline{s}(\widetilde{t}) and s‾(t~)\overline{s}(\widetilde{t})) should match to leading order.

From the second time scale approximation,

By matching, Equations (73) and (75) should be coherent. Thus the ODE for the second time scale should be initialized from a‾(0)(0)=φ0σ0\mathds1+a⊥,init\overline{a}^{(0)}(0)=\frac{\varphi_{0}}{\sigma_{0}}\mathds{1}+a_{\perp,\rm{init}}.

Similarly, the matching procedure gives that the ODE for the second time scale should be initialized from s‾(1)=0\overline{s}^{(1)}=0.

Solution.

As we are done with the matching procedure, we now consider the solution in the second time scale only, that we denote again by aa, ss as in (65), (66). The matching procedure motivates us to consider the solution of (67)-(68) initialized at a⊥(0)(0)=a⊥,inita_{\perp}^{(0)}(0)=a_{\perp,\rm{init}}, s⊥(1)=0s_{\perp}^{(1)}=0. This gives

To conclude, we note that ⟨a(0),\mathds1⟩L2(ρ)=φ0σ0\langle a^{(0)},\mathds{1}\rangle_{L^{2}(\rho)}=\frac{\varphi_{0}}{\sigma_{0}} is constrained by (62). Further, from (64),

thus ⟨s(1),\mathds1⟩L2(ρ)=σ1φ1φ0σ0t2\langle s^{(1)},\mathds{1}\rangle_{L^{2}(\rho)}=\sigma_{1}\varphi_{1}\frac{\varphi_{0}}{\sigma_{0}}t_{2}.

We observe that a(0)a^{(0)} and s(1)s^{(1)} diverge as t2→∞t_{2}\to\infty. This implies that our approximation on the second time scale must break down at a certain point. Indeed, we analyzed this time scale under the assumption that both a(0)a^{(0)} and s(1)s^{(1)} are of order 11. However, since a(0)a^{(0)} and s(1)s^{(1)} diverge exponentially as t2→∞t_{2}\to\infty, as per Eq. (76), this assumption breaks down when t2≍log⁡(1/ε)t_{2}\asymp\log(1/\varepsilon).

More precisely, in (59) (resp. (61)), the O(ε)O(\varepsilon) term includes a term of the form

When a(0)a^{(0)} and s(1)s^{(1)} become of order ε−\nicefrac14\varepsilon^{-\nicefrac{{1}}{{4}}}, this term becomes of order ε\nicefrac14\varepsilon^{\nicefrac{{1}}{{4}}}, which is then of the same order as the term ε\nicefrac12σ1φ1s(1)(ω)\varepsilon^{\nicefrac{{1}}{{2}}}\sigma_{1}\varphi_{1}s^{(1)}(\omega) in (59) (resp. the term ε\nicefrac12σ1φ1a(0)(ω)\varepsilon^{\nicefrac{{1}}{{2}}}\sigma_{1}\varphi_{1}a^{(0)}(\omega) in (61)). At this point, these terms can not be neglected anymore. From (76), we have

Therefore, a(0)a^{(0)} and s(1)s^{(1)} become of order ε−\nicefrac14\varepsilon^{-\nicefrac{{1}}{{4}}} at the time t2∼14∣σ1φ1∣log⁡1εt_{2}\sim\frac{1}{4|\sigma_{1}\varphi_{1}|}\log\frac{1}{\varepsilon}, at which the approximation on the second time scale breaks down. We thus introduce a new time scale centered at this critical point.

3 Third time scale: linear component II

We now introduce the time t3=t2−14∣φ1σ1∣log⁡1εt_{3}=t_{2}-\frac{1}{4|\varphi_{1}\sigma_{1}|}\log\frac{1}{\varepsilon}. As t3t_{3} is only a translation from t2t_{2}, the ODEs in terms of t3t_{3} are the same as the ones in term of t2t_{2}. However, in this time scale, aa and ε\nicefrac12s\varepsilon^{\nicefrac{{1}}{{2}}}s have diverged. In coherence with the discussion above, we seek expansions of the form

Similarly to the second time scale, we substitute (77)-(78) in (39) and obtain

First, we identify the terms of order ε−\nicefrac14\varepsilon^{-\nicefrac{{1}}{{4}}}:

This means that aa has no component diverging in ε\varepsilon in the direction of \mathds1\mathds{1}.

Second, we identify the terms of order 1=ε01=\varepsilon^{0}:

Put together with (79), this equation ensures that the constant component of φ\varphi remains learned on this third time scale.

Third, we identify the terms of order ε\nicefrac14\varepsilon^{\nicefrac{{1}}{{4}}}:

where in the last equality we use (79). Thus we can rewrite (81) as

In Appendix B.1, we solve this system of ODEs and determine the initial condition by matching with the previous layer. The result is that

where λ=λ(t3)\lambda=\lambda(t_{3}) is the function

This solution finishes to describe how the linear part of the function φ\varphi is learned.

4 Conjectured behavior for larger time scales

The analysis of the previous sections naturally suggests the existence of a sequence of cutoffs. At each time scale, a new polynomial component of φ\varphi is learned within a window that is much shorter than the time elapsed before that phase started. Along this sequence, we expect ss and aa to grow to increasingly larger scales in ε\varepsilon (but ss remains o(1)o(1) while aa diverges).

More precisely, we assume that during the ll-th phase, the network learns the degree-ll component φl\varphi_{l}, and various quantities satisfy the following scaling behavior:

where ωl>0\omega_{l}>0 is an increasing sequence and βl,μl>0\beta_{l},\mu_{l}>0 are decreasing sequences. Further, while learning of this component takes place when t=O(εμl)t=O(\varepsilon^{\mu_{l}}), the actual evolution of the risk (and of the neural network) take place on much shorter scales, namely:

where νl\nu_{l} is also decreasing, with νl>μl\nu_{l}>\mu_{l}. The goal of this section is to provide heuristic arguments to conjecture the values of ωl\omega_{l}, βl\beta_{l}, μl\mu_{l} and νl\nu_{l}. We will base this conjecture on a rigorous analysis of a simplified model.

We capture the effect of learning dynamics on the previous time scales by the overall magnitude of the a(ω)a(\omega)’s and s(ω)s(\omega)’s at initialization. Namely, we choose the scale of initialization of the simplified model to be given by the end of the (l−1)(l-1)-th time scale, i.e., a(ω)≍ε−ωl−1a(\omega)\asymp\varepsilon^{-\omega_{l-1}} and s(ω)≍εβl−1s(\omega)\asymp\varepsilon^{\beta_{l-1}}. Further, in order for the (l−1)(l-1)-th component to be learned, namely

Based on this consideration, we introduce the rescaled variables

Rewriting Eq. (88) in terms of a~(ω)\widetilde{a}(\omega)’s and s~(ω)\widetilde{s}(\omega)’s, and using ωl=lβl\omega_{l}=l\beta_{l}, we get that

In order for the a~(ω)\widetilde{a}(\omega)’s and s~(ω)\widetilde{s}(\omega)’s to be learned simultaneously, we need 1−2lβl=2βl1-2l\beta_{l}=2\beta_{l}, which implies βl=1/2(l+1)\beta_{l}=1/2(l+1). Making a further change of the time variable t=ενlτt=\varepsilon^{\nu_{l}}\tau, where νl=2βl=1/(l+1)\nu_{l}=2\beta_{l}=1/(l+1), it follows that

Moreover, rewriting the risk in terms of the rescaled variables a~,s~\widetilde{a},\widetilde{s}, \mathscrsfsRl(τ)=\mathscrsfsRl(a~(τ),s~(τ))\mathscrsfs{R}_{l}(\tau)=\mathscrsfs{R}_{l}(\widetilde{a}(\tau),\widetilde{s}(\tau)) satisfies the ODE:

Note that with our choice of βl\beta_{l} and ωl\omega_{l}, we have ωl−ωl−1=βl−1−βl=1/2l(l+1)\omega_{l}-\omega_{l-1}=\beta_{l-1}-\beta_{l}=1/2l(l+1). This means that the a~(ω)\widetilde{a}(\omega)’s and s~(ω)\widetilde{s}(\omega)’s are initialized at the same scale, namely

The theorem below describes quantitatively the dynamics of the simplified model for small ε\varepsilon, and determines the value of μl\mu_{l} (recall that νl=1/(l+1)\nu_{l}=1/(l+1)):

Assume l≥2l\geq 2 and let (a~(ω,τ),s~(ω,τ))τ≥0(\widetilde{a}(\omega,\tau),\widetilde{s}(\omega,\tau))_{\tau\geq 0} be the unique solution of the ODE system (91), initialized as per Eq. (93) (note in particular that σlφla~(ω,0)s~(ω,0)l≍ε1/2l\sigma_{l}\varphi_{l}\widetilde{a}(\omega,0)\widetilde{s}(\omega,0)^{l}\asymp\varepsilon^{1/2l}). Then the followings hold:

and assume ρ(A)>0\rho(A)>0. For Δ∈(0,φl2/2)\Delta\in(0,\varphi_{l}^{2}/2), define

Then, for any fixed Δ\Delta we have τ(Δ)=Θ(ε−(l−1)/2l(l+1))\tau(\Delta)=\Theta(\varepsilon^{-(l-1)/2l(l+1)}) as ε→0\varepsilon\to 0. Further, if ρ\rho is a discrete probability measure, then there exists τ∗(ε)=Θ(ε−(l−1)/2l(l+1))\tau_{*}(\varepsilon)=\Theta(\varepsilon^{-(l-1)/2l(l+1)}) and, for any Δ>0\Delta>0 a constant c∗(Δ)>0c_{*}(\Delta)>0 independent of ε\varepsilon such that

namely the ll-th component is learnt in an O(1)O(1) time window around τ∗(ε)=Θ(ε−(l−1)/2l(l+1))\tau_{*}(\varepsilon)=\Theta(\varepsilon^{-(l-1)/2l(l+1)}).

If ρ(B)>0\rho(B)>0, then the same claims as in (a)(a) hold.

If neither of the conditions at points (a)(a), (b)(b) holds, and

for almost every ω∈Ω\omega\in\Omega. Then, for such ω∈Ω\omega\in\Omega and each Δ>0\Delta>0, there exists a constant C∗(ω,Δ)>0C_{*}(\omega,\Delta)>0 such that

meaning that s~(ω,τ)\widetilde{s}(\omega,\tau) converges to eventually.

We further note that τ=Θ(ε−(l−1)/2l(l+1))⟺t=Θ(εμl)\tau=\Theta(\varepsilon^{-(l-1)/2l(l+1)})\Longleftrightarrow t=\Theta(\varepsilon^{\mu_{l}}) with μl=1/2l\mu_{l}=1/2l, and τ=O(1)⟺t=O(ενl)\tau=O(1)\Longleftrightarrow t=O(\varepsilon^{\nu_{l}}) with νl=1/(l+1)\nu_{l}=1/(l+1).

The proof of Theorem 1 is deferred to Appendix B.3.

Under the conditions of cases (a)(a) and (b)(b), we see that the degree-ll component of the target function is learnt within an O(ε1/(l+1))O(\varepsilon^{1/(l+1)}) time window around t∗(l,ε)≍ε1/2lt_{*}(l,\varepsilon)\asymp\varepsilon^{1/2l}, which is consistent with the timescales conjectured in Definition 1.

Case (c)(c) corresponds to s(ω)/s(ω,0)s(\omega)/s(\omega,0) becoming close to in time t=O(εμl)t=O(\varepsilon^{\mu_{l}}), and staying at . In other words, the neurons become orthogonal to the target direction and play no role in learning higher-degree components any longer.

Informally, case (c)(c) couples the learning of different polynomial components. It can happen that the learning phase l−1l-1 induces an effective initialization (a~(ω,0), s~(ω,0))(\widetilde{a}(\omega,0),\ \widetilde{s}(\omega,0)) within the domain of case (c)(c).

We expect this not to be the case for suitable choices of initialization (or equivalently PA{\rm P}_{A}), φ\varphi, and σ\sigma. Establishing this would amount to establishing that the standard learning scenario holds.

Stochastic gradient descent and finite sample size

So far we focused on analyzing the projected gradient flow (GF) dynamics with respect to the population risk, as defined in Eqs. (4)-(5). In this section, we extract the implications of our analysis of GF on online projected stochastic gradient descent, which is a projected version of the SGD dynamics (151).

The projected SGD dynamics is specified as follows:

We prove that, for small η\eta, the projected SGD of Eq. (101) is close to the gradient flow of Eqs. (4)-(5). Throughout this section, we make the following assumptions similar to those assumed in Section 4:

We then require the functions VV and UU to be bounded and differentiable, with uniformly bounded and Lipschitz continuous gradients for all ∥u∥2,∥u′∥2≤2\left\|{u}\right\|_{2},\left\|{u^{\prime}}\right\|_{2}\leq 2:

Similar to Remark 4.1, we can show that a sufficient condition for Eq.s (104) and (105) is

where the constant M2′M_{2}^{\prime} depends uniquely on M2M_{2}.

Assume (x,y)∼\mathdsP(x,y)\sim\mathds{P}, then we require that y∈[−M3,M3]y\in[-M_{3},M_{3}] almost surely. Moreover, we assume that for all ∥u∥2≤2\left\|{u}\right\|_{2}\leq 2, both σ(⟨u,x⟩)\sigma(\langle u,x\rangle) and σ′(⟨u,x⟩)(x−⟨u,x⟩u)\sigma^{\prime}(\langle u,x\rangle)(x-\langle u,x\rangle u) are M3M_{3}-sub-Gaussian.

The following theorem upper bounds the distance between gradient flow and projected stochastic gradient descent dynamics.

Let θi(t)=(ai(t),ui(t))\theta_{i}(t)=(a_{i}(t),u_{i}(t)) be the solution of the GF ordinary differential equations (4)-(5). There exists a constant MM that only depends on the MiM_{i}’s from Assumptions A1-A3, such that for any T,z≥0T,z\geq 0 and

the following holds with probability at least 1−exp⁡(−z2)1-\exp(-z^{2}):

The proof is presented in Appendix C and follows the same scheme as in that of Theorem 1 part (B) in (Mei et al., 2019). The main difference with respect to that theorem is here we are interested in projected SGD (and GF) instead of plain SGD (and GF), hence an additional step of approximation is required, and the aia_{i}’s and uiu_{i}’s need to be treated separately. We next draw implications of the last result on learning by online SGD within the standard learning scenario.

Fix any δ>0\delta>0. Assume φ,σ\varphi,\sigma and the initialization PA{\rm P}_{A} be such that the standard learning scenario of Definition 1 holds up to level LL for some L≥2L\geq 2, and that

Then, there exist constants ε∗=ε∗(δ)\varepsilon_{*}=\varepsilon_{*}(\delta), T0=T0(δ)T_{0}=T_{0}(\delta), T=T(ε,δ)=T0(δ)ε1/(2L)T=T(\varepsilon,\delta)=T_{0}(\delta)\varepsilon^{1/(2L)} and M=M(ε,δ)M=M(\varepsilon,\delta) that depend on ε,δ\varepsilon,\delta (together with φ,σ\varphi,\sigma and PA{\rm P}_{A}) such that the following happens. Assume ε≤ε∗(δ)\varepsilon\leq\varepsilon_{*}(\delta) and m,d,zm,d,z are such that d≥Md\geq M, m≥max⁡(M,z)m\geq\max(M,z), and the step size η\eta and number of samples (equivalently, number of steps) nn satisfy

Then, with probability at least 1−e−z1-e^{-z}, the projected gradient descent algorithm of Eq. (101) achieves population risk smaller than δ\delta:

The proof of Theorem 3 is deferred to Appendix C.4.

In contrast, Theorem 3 shows that, within the standard learning scenario, O(d)O(d) samples and O(1)O(1) neurons are sufficient. Further as per Theorem 2, the learning dynamics is accurately described by the GF analyzed in the previous sections.

Acknowledgments

This work was supported by the NSF through award DMS-2031883, the Simons Foundation through Award 814639 for the Collaboration on the Theoretical Foundations of Deep Learning, the NSF grant CCF-2006489 and the ONR grant N00014-18-1-2729, and a grant from Eric and Wendy Schmidt at the Institute for Advanced Studies. Part of this work was carried out while Andrea Montanari was on partial leave from Stanford and a Chief Scientist at Ndata Inc dba Project N. The present research is unrelated to AM’s activity while on leave.

References

Appendix A Appendix to Section 4

This proves (14). Equation (15) follows directly:

To obtain equations (16)-(18), we now take gradients in (113):

This gives (16). Finally, we perform a similar computation to compute ∂trij=⟨∂tui,uj⟩+⟨ui,∂tuj⟩\partial_{t}r_{ij}=\langle\partial_{t}u_{i},u_{j}\rangle+\langle u_{i},\partial_{t}u_{j}\rangle. We compute only the first term, as the second term can be obtained by inverting ii and jj:

Adding the symmetric term ⟨ui,∂tuj⟩\langle u_{i},\partial_{t}u_{j}\rangle, we obtain (17)-(18).

A.2 Proof of Corollary 1

First, note that in the proof of Lemma 1, we obtain the following a priori estimate on the magnitude of the ai0a_{i}^{0}’s:

where MM only depends on the MiM_{i}’s in Assumptions A1-A3. Using a similar argument as that in the proof of Proposition 2, we obtain that for any t∈[0,T]t\in[0,T] and i∈[m]i\in[m],

then we know that G′(t)≤(M(1+t)2/ε2)G(t)G^{\prime}(t)\leq(M(1+t)^{2}/\varepsilon^{2})G(t). Applying Grönwall’s inequality yields

with probability at least 1−exp⁡(C′m)1-\exp(C^{\prime}m), where CC and C′C^{\prime} are both absolute constants. Therefore,

Next we upper bound the risk difference, by direct calculation,

with probability at least 1−exp⁡(−C′m)1-\exp(-C^{\prime}m), where the constant MM only depends on the MiM_{i}’s from Assumptions A1-A3. The conclusion now follows from taking the supremum over all t∈[0,T]t\in[0,T]. This completes the proof of Corollary 1.

A.3 Proof of Proposition 2

We consider rij⊥=rij−sisj=⟨ui,uj⟩−⟨ui,u∗⟩⟨u∗,uj⟩r_{ij}^{\perp}=r_{ij}-s_{i}s_{j}=\langle u_{i},u_{j}\rangle-\langle u_{i},u_{*}\rangle\langle u_{*},u_{j}\rangle, the dot product between uiu_{i} and uju_{j} that is out of the relevant subspace spanned by u∗u_{*}. We show that these variables satisfy the ODEs

By definition of rij⊥r_{ij}^{\perp}, we readily see that

If Assumptions A1-A3 hold, then we have for any fixed T>0T>0:

To begin with, using Eq. (119), we obtain that

Using the ODEs for the aia_{i}’s, we obtain that

where (i)(i) follows from our assumptions and the fact that \mathscrsfsR(a(t),u(t))≤\mathscrsfsR(a(0),u(0))\mathscrsfs{R}(a(t),u(t))\leq\mathscrsfs{R}(a(0),u(0)), since ∂t\mathscrsfsR(a,u)≤0\partial_{t}\mathscrsfs{R}(a,u)\leq 0 by gradient flow equations. Moreover, the constant MM only depends on the MiM_{i}’s. Since ∣ai(0)∣≤M1|a_{i}(0)|\leq M_{1} for all i∈[m]i\in[m], we know that ∣ai(t)∣≤M(1+t/ε)|a_{i}(t)|\leq M(1+t/\varepsilon) for all t≥0t\geq 0, thus leading to the following estimate:

where the constant MM only depends on the MiM_{i}’s in our assumptions. At initialization, we know that ∑i,j=1mrij⊥(0)2=m\sum_{i,j=1}^{m}r_{ij}^{\perp}(0)^{2}=m. Applying Grönwall’s inequality yields that

To this end, we define S(t)=∑i=1m((ai(t)−ai\mboxmf(t))2+(si(t)−si\mboxmf(t))2)S(t)=\sum_{i=1}^{m}\left(\left(a_{i}(t)-a_{i}^{\mbox{\tiny\rm mf}}(t)\right)^{2}+\left(s_{i}(t)-s_{i}^{\mbox{\tiny\rm mf}}(t)\right)^{2}\right). By our assumption, S(0)=0S(0)=0. Moreover, using the same technique as in the proof of Lemma 1, we know that ∣ai\mboxmf(t)∣≤M(1+t/ε)|a_{i}^{\mbox{\tiny\rm mf}}(t)|\leq M(1+t/\varepsilon) for all i∈[m]i\in[m]. According to Eq.s (15)-(18) and Eq. (23), we deduce that

where in (i)(i) we use the Cauchy-Schwarz inequality and the inequality of arithmetic and geometric means, and (ii)(ii) follows from the conclusion of Lemma 1. Similarly, we obtain that

Combining the above estimates, we finally deduce that

Applying Grönwall’s inequality immediately implies

which further leads to Eq. (120) and concludes the proof of Proposition 2. The “consequently” part can be shown via direct calculation, but we include it here for the sake of completeness. By definition, for any t∈[0,T]t\in[0,T] we have

A.4 Derivation of the mean field dynamics (28)

where (i)(i) follows from the ODE satisfied by the (ai\mboxmf(t),si\mboxmf(t))(a_{i}^{\mbox{\tiny\rm mf}}(t),s_{i}^{\mbox{\tiny\rm mf}}(t))’s, and in (ii)(ii) we use integration by parts. We thus obtain that

A.5 Details of the alternative mean field approach

where Ψ‾=(Ψ‾a,Ψ‾u)\overline{\Psi}=(\overline{\Psi}_{a},\overline{\Psi}_{u}) is given by

A remarkable property of the equation (124) is that it preserves invariance to rotations orthogonal to u∗u_{*}. Indeed, assume that ρ‾\overline{\rho} is invariant to rotations orthogonal to u∗u_{*}. In this case, we show that Ψ‾a(a,u;ρ‾)\overline{\Psi}_{a}\left(a,u;\overline{\rho}\right) and ⟨u∗,Ψ‾u(a,u;ρ‾)⟩\langle u_{*},\overline{\Psi}_{u}\left(a,u;\overline{\rho}\right)\rangle depend only on s:=⟨u,u∗⟩s:=\langle u,u_{*}\rangle and s1:=⟨u1,u∗⟩s_{1}:=\langle u_{1},u_{*}\rangle. Let u⊥u^{\perp} (resp. u1⊥u_{1}^{\perp}) denote the component of uu (resp. u1u_{1}) orthogonal to u∗u_{*}. Let RR denote a random uniform rotation orthogonal to u∗u_{*}. By the rotation invariance of ρ‾\overline{\rho},

The random variable B(d)=⟨u⊥∥u⊥∥,Ru1⊥∥u1⊥∥⟩B^{(d)}=\left\langle\frac{u^{\perp}}{\|u^{\perp}\|},R\frac{u_{1}^{\perp}}{\|u_{1}^{\perp}\|}\right\rangle is a one dimensional projection of a random variable uniform on the unit sphere of the hyperplane orthogonal to u∗u_{*}; thus it has the density pB(d)(b)∝(1−b2)d/2−2p_{B^{(d)}}(b)\propto(1-b^{2})^{d/2-2} (see, e.g., [Frye and Efthimiou, 2012, Lemma 4.17]). Denote

In the equation above, we have ⟨u∗,(Id−uu⊤)s1u∗⟩=s1(1−s2)\langle u_{*},(I_{d}-uu^{\top})s_{1}u_{*}\rangle=s_{1}(1-s^{2}) and as ⟨u∗,Ru1⊥⟩=0\langle u_{*},Ru_{1}^{\perp}\rangle=0 a.s., we have

Of course, a discrete measure of the form (123) can not be invariant to rotations orthogonal to u∗u_{*}. However, if the uiu_{i} are initialized uniformly on the unit sphere, then the measure ρ‾0\overline{\rho}_{0} converges to a measure with the rotation invariance as m→∞m\to\infty. One can then apply the results of Mei et al. to control the deviations from this limit. Let us thus assume that ρ‾0\overline{\rho}_{0} satisfies the rotation invariance. Define the map φ(a,u)=(a,⟨u,u∗⟩)\varphi(a,u)=(a,\langle u,u_{*}\rangle). Then, from (125), (126), the push-forward ρt\rho_{t} of ρ‾t\overline{\rho}_{t} through the map φ\varphi satisfies the continuity equation

where Ψ(d)=(Ψa(d),Ψs(d))\Psi^{(d)}=(\Psi^{(d)}_{a},\Psi^{(d)}_{s}) is given by

Appendix B Calculations for the analysis of mean-field gradient flow

In order to solve the system (83), we start from an associated one-dimensional ODE.

The solution λ=λ(t3)\lambda=\lambda(t_{3}) of the ODE

For simplicity, denote α=∣σ1∣\alpha=|\sigma_{1}|, β=∣φ1∣\beta=|\varphi_{1}| and γ=∣σ1∣∥a⊥,init∥L2(ρ)2\gamma={|\sigma_{1}|}\left\|a_{\perp,\rm{init}}\right\|_{L^{2}(\rho)}^{2}. Then

This is Bernoulli differential equation (see, e.g., Encyclopedia of Mathematics ). In this situation, the classical trick is to reduce the problem to a linear inhomogeneous first-order equation by considering

Solving this linear inhomogeneous first-order equation gives

Let λ=λ(t3)\lambda=\lambda(t_{3}) be a solution of (127) and consider

Then a(−1),s(1)a^{(-1)},s^{(1)} are solutions of the constrained ODE system (79), (82). Indeed,

thus the constraint (79) is satisfied. Further

A similar computation shows that the differential equation for s(1)s^{(1)} is also satisfied. This concludes that (129) is a valid candidate to solve the third time scale.

To determine the value of the initialization λ(0)\lambda(0) we perform a matching procedure with the previous time scale. In this paragraph, we denote a‾,s‾\underline{a},\underline{s} the approximation obtained in the second time scale (Section 6.2), and a‾,s‾\overline{a},\overline{s} the approximation in the third time scale (Section 6.3 and above).

Consider an intermediate time scale t~=t2−clog⁡1ε\widetilde{t}=t_{2}-c\log\frac{1}{\varepsilon} with 0<c<14∣σ1φ1∣0<c<\frac{1}{4|\sigma_{1}\varphi_{1}|}. Assume t~≍1\widetilde{t}\asymp 1. Then

From the approximation (76) on the second time scale,

From the approximation on the third time scale,

Note that as t3→−∞t_{3}\to-\infty, from (128),

By matching, Equations (130) and (131) should be coherent. This gives

One could check similarly that s(1)s^{(1)} also satisfies the matching conditions under the same constraint, and thus that (129) are indeed the solutions of the third time scale.

B.2 Induced approximation of the risk

In this section, we show that the behavior of aa and ss derived in Sections 6.1–6.3 leads to an evolution of the risk alternating plateaus and rapid decreases, in agreement with the standard scenario of Definition 1. For the convenience of the reader, we recall the expression (36) of the risk

This describes, in a more detailed form, the first transition in Definition 1.

This second time scale does not induce any transition of the risk \mathscrsfsR\mboxmf,∗\mathscrsfs{R}_{\mbox{\tiny\rm mf},*} (but was necessary to understand the divergence of aa and ε−\nicefrac12s\varepsilon^{-\nicefrac{{1}}{{2}}}s).

where in (a)(a) we used (84) and in (b)(b) (85). Thus as ε→0\varepsilon\to 0,

This describes, in a more detailed form, the second transition in Definition 1.

B.3 Proof of Theorem 1

Throughout the proof, we will use the shorthand \mathscrsfsRl(τ)\mathscrsfs{R}_{l}(\tau) to represent \mathscrsfsRl(a~(τ),s~(τ))\mathscrsfs{R}_{l}(\widetilde{a}(\tau),\widetilde{s}(\tau)). First, note that according to the ODE satisfied by \mathscrsfsRl\mathscrsfs{R}_{l} (Eq. (92)), we know that \mathscrsfsRl\mathscrsfs{R}_{l} must be non-increasing, thus for small enough ε>0\varepsilon>0,

According to the comparison theorem for system of ODEs, we know that ∣a~(ω,τ)∣≤a^(ω,τ)|\widetilde{a}(\omega,\tau)|\leq\widehat{a}(\omega,\tau), ∣s~(ω,τ)∣≤s^(ω,τ)|\widetilde{s}(\omega,\tau)|\leq\widehat{s}(\omega,\tau) for all τ≥0\tau\geq 0 where

The above system of ODEs can be solved analytically via integration. First, we note that

which implies that (further note s^(ω,0)2=la^(ω,0)2\widehat{s}(\omega,0)^{2}=l\widehat{a}(\omega,0)^{2})

The ODE system then reduces to ∂τa^(ω)=2ll/2∣σl∣∣φl∣a^(ω)l\partial_{\tau}\widehat{a}(\omega)=2l^{l/2}|\sigma_{l}||\varphi_{l}|\widehat{a}(\omega)^{l}, which admits the solution

Since a^(ω,0)=Θ(ε1/2l(l+1))\widehat{a}(\omega,0)=\Theta(\varepsilon^{1/2l(l+1)}), we know that a^(ω,τ),s^(ω,τ)=o(1)\widehat{a}(\omega,\tau),\widehat{s}(\omega,\tau)=o(1) until τ=Θ(ε−(l−1)/2l(l+1))−O(1)=Θ(ε−(l−1)/2l(l+1))\tau=\Theta(\varepsilon^{-(l-1)/2l(l+1)})-O(1)=\Theta(\varepsilon^{-(l-1)/2l(l+1)}), which means that a~(ω,τ),s~(ω,τ)=o(1)\widetilde{a}(\omega,\tau),\widetilde{s}(\omega,\tau)=o(1) until τ=Ω(ε−(l−1)/2l(l+1))\tau=\Omega(\varepsilon^{-(l-1)/2l(l+1)}). As a consequence,

until τ=Ω(ε−(l−1)/2l(l+1))\tau=\Omega(\varepsilon^{-(l-1)/2l(l+1)}). This means that the learning of the ll-th component will not begin until τ=Ω(ε−(l−1)/2l(l+1))\tau=\Omega(\varepsilon^{-(l-1)/2l(l+1)}), namely τ(Δ)=Ω(ε−(l−1)/2l(l+1))\tau(\Delta)=\Omega(\varepsilon^{-(l-1)/2l(l+1)}) for any fixed Δ>0\Delta>0. Note that the above argument applies to all of the settings in the theorem statement.

Next, we show that for any fixed Δ>0\Delta>0, τ(Δ)=O(ε−(l−1)/2l(l+1))\tau(\Delta)=O(\varepsilon^{-(l-1)/2l(l+1)}), which means that the ll-th component can be learnt in O(ε−(l−1)/2l(l+1))O(\varepsilon^{-(l-1)/2l(l+1)}) time. To prove our claim by contradiction, assume that there exists Δ>0\Delta>0 and a sequence εk↓0\varepsilon_{k}\downarrow 0, such that

By definition of τ(Δ)\tau(\Delta), we know that ∀τ≤τ(Δ)\forall\tau\leq\tau(\Delta),

Now, assume the condition of setting (a) holds and denote

Then by definition and our assumption that a~(ω,0)\widetilde{a}(\omega,0) is of the same order as s~(ω,0)\widetilde{s}(\omega,0), we know that A=∪ε0>0,η>0Aε0,ηA=\cup_{\varepsilon_{0}>0,\eta>0}A_{\varepsilon_{0},\eta}. Since ρ(A)>0\rho(A)>0, there exists ε0,η>0\varepsilon_{0},\eta>0 such that ρ(Aε0,η)>0\rho(A_{\varepsilon_{0},\eta})>0. Note that here we can choose ε0\varepsilon_{0} and η\eta to be arbitrarily small since the set Aε0,ηA_{\varepsilon_{0},\eta} is non-increasing in ε0\varepsilon_{0} and η\eta. For ω∈Aε0,η\omega\in A_{\varepsilon_{0},\eta} and τ≤τ(Δ)\tau\leq\tau(\Delta), we have

Moreover, we know that at initialization, ∣a~(ω,0)∣,∣s~(ω,0)∣>ηε1/2l(l+1)|\widetilde{a}(\omega,0)|,|\widetilde{s}(\omega,0)|>\eta\varepsilon^{1/2l(l+1)}. Using the ODE comparison theorem and a similar argument as that in proving τ(Δ)=Ω(ε−(l−1)/2l(l+1))\tau(\Delta)=\Omega(\varepsilon^{-(l-1)/2l(l+1)}), we deduce that for sufficiently large kk such that ε=εk<ε0\varepsilon=\varepsilon_{k}<\varepsilon_{0}, there exist constants C,C′>0C,C^{\prime}>0 that does not depend on ε\varepsilon satisfying the following: For all ω∈Aε0,η\omega\in A_{\varepsilon_{0},\eta} and τ≥Cε−(l−1)/2l(l+1)\tau\geq C\varepsilon^{-(l-1)/2l(l+1)},

This further implies that at time τ\tau,

According to Eq. (92), we know that \mathscrsfsRl\mathscrsfs{R}_{l} will decrease to exponentially fast in an O(1)O(1) time window after τ=Cε−(l−1)/2l(l+1)\tau=C\varepsilon^{-(l-1)/2l(l+1)}, which contradicts our assumption (136). This proves that τ(Δ)=O(ε−(l−1)/2l(l+1))\tau(\Delta)=O(\varepsilon^{-(l-1)/2l(l+1)}) under setting (a). Next, we show that setting (b) can be reduced to setting (a). Under setting (b), let us denote

Then similar to the previous argument, there exists ε0,η>0\varepsilon_{0},\eta>0 such that ρ(Bε0,η)>0\rho(B_{\varepsilon_{0},\eta})>0, and further we can choose ε0\varepsilon_{0} and η\eta to be arbitrarily small. For ω∈Bε0,η\omega\in B_{\varepsilon_{0},\eta}, we have

Hence, both a~(ω)2\widetilde{a}(\omega)^{2} and s~(ω)2\widetilde{s}(\omega)^{2} will decrease at initialization. Moreover, Eq. (91) implies that

Integrating both sides of the above equation, we obtain that

which is close to (s~(ω,0)2−s~(ω,τ)2)/l\left(\widetilde{s}(\omega,0)^{2}-\widetilde{s}(\omega,\tau)^{2}\right)/l as long as s~(ω,τ)=O(1)\widetilde{s}(\omega,\tau)=O(1). To be accurate, let us define

then we know that s~(ω,τa,ω)=Ω(ε1/2l(l+1))\widetilde{s}(\omega,\tau_{a,\omega})=\Omega(\varepsilon^{1/2l(l+1)}) and τa,ω=O(ε−(l−1)/2l(l+1))\tau_{a,\omega}=O(\varepsilon^{-(l-1)/2l(l+1)}) under the assumption (136), where the latter claim can be proved through making the change of variable a~′(ω)=ε−1/2l(l+1)a~(ω)\widetilde{a}^{\prime}(\omega)=\varepsilon^{-1/2l(l+1)}\widetilde{a}(\omega) and s~′(ω)=ε−1/2l(l+1)s~(ω)\widetilde{s}^{\prime}(\omega)=\varepsilon^{-1/2l(l+1)}\widetilde{s}(\omega). Note that after the time point τa,ω\tau_{a,\omega}, the sign of a~(ω)\widetilde{a}(\omega) changes. Hence, φlσla~(ω)s~(ω)l>0\varphi_{l}\sigma_{l}\widetilde{a}(\omega)\widetilde{s}(\omega)^{l}>0, and a~(ω,τ)2\widetilde{a}(\omega,\tau)^{2} and s~(ω,τ)2\widetilde{s}(\omega,\tau)^{2} will begin to increase for τ≥τa,ω\tau\geq\tau_{a,\omega}. Similarly, we can show that in O(ε−(l−1)/2l(l+1))O(\varepsilon^{-(l-1)/2l(l+1)}) time after τa,ω\tau_{a,\omega}, both a~(ω)\widetilde{a}(\omega) and s~(ω)\widetilde{s}(\omega) become of order ε1/2l(l+1)\varepsilon^{1/2l(l+1)}, and we still have φlσla~(ω)s~(ω)l>0\varphi_{l}\sigma_{l}\widetilde{a}(\omega)\widetilde{s}(\omega)^{l}>0. This reduces our case (b)(b) to case (a)(a).

We have proven that under settings (a) and (b), τ(Δ)=Θ(ε−(l−1)/2l(l+1))\tau(\Delta)=\Theta(\varepsilon^{-(l-1)/2l(l+1)}) for any fixed Δ∈(0,φl2/2)\Delta\in(0,\varphi_{l}^{2}/2). This means that some of the neurons (a~(ω),s~(ω))(\widetilde{a}(\omega),\widetilde{s}(\omega)) become of order Ω(1)\Omega(1) and the ll-th component of the target function is learnt at a timescale of order ε−(l−1)/2l(l+1)\varepsilon^{-(l-1)/2l(l+1)}. Next, we show that if the probability measure ρ\rho is discrete, then the evolution of \mathscrsfsRl\mathscrsfs{R}_{l} actually happens in an O(1)O(1) time window. It suffices to prove that, for any Δ>0\Delta>0 a small constant (Δ<φl2/4\Delta<\varphi_{l}^{2}/4),

as ε→0\varepsilon\to 0. Note that by continuity and monotonicity of \mathscrsfsRl\mathscrsfs{R}_{l}, we have

By definition of \mathscrsfsRl\mathscrsfs{R}_{l}, we know that ∀τ≥τ(φl2/2−Δ)\forall\tau\geq\tau(\varphi_{l}^{2}/2-\Delta),

Denote by {(a~i,s~i)}i∈[m]\{(\widetilde{a}_{i},\widetilde{s}_{i})\}_{i\in[m]} the realizations of {(a~(ω),s~(ω))}ω∈Ω\{(\widetilde{a}(\omega),\widetilde{s}(\omega))\}_{\omega\in\Omega} under the discrete measure ρ\rho, and by {pi}i∈[m]\{p_{i}\}_{i\in[m]} the point masses of ρ\rho. Then, we know that

which implies that ∃j∈[m]\exists j\in[m], s.t. ∣a~j(τ)s~j(τ)l∣≥rl(Δ)\left|\widetilde{a}_{j}(\tau)\widetilde{s}_{j}(\tau)^{l}\right|\geq r_{l}(\Delta). Applying Lemma 3 yields

It then follows from Eq. (92) that \mathscrsfsRl\mathscrsfs{R}_{l} will decrease to exponentially fast, and Eq. (140) holds consequently. This completes the proof for settings (a) and (b).

We then focus on the case (c). By our assumption, for almost every ω\omega there exists η>0\eta>0 (may depend on ω\omega) such that

for sufficiently small ε\varepsilon. Therefore, s~(ω,τ)2\widetilde{s}(\omega,\tau)^{2} and a~(ω,τ)2\widetilde{a}(\omega,\tau)^{2} will keep decreasing until one of them reaches , which means that

According to Eq. (139) and the inequality s~(ω,0)2<(l−η)a~(ω,0)2\widetilde{s}(\omega,0)^{2}<(l-\eta)\widetilde{a}(\omega,0)^{2}, a~(ω,τ)2\widetilde{a}(\omega,\tau)^{2} will not reach until s~(ω,τ)2\widetilde{s}(\omega,\tau)^{2} reaches . Furthermore, for any τ≥0\tau\geq 0,

Using again the comparison theorem for ODE, we get that

Since s~(ω,0)≍ε1/2l(l+1)\widetilde{s}(\omega,0)\asymp\varepsilon^{1/2l(l+1)}, it follows immediately that for any Δ>0\Delta>0, there exists a constant C∗(ω,Δ)>0C_{*}(\omega,\Delta)>0 such that

This completes the discussion for case (c), thus concluding the proof of Theorem 1.

Let r>0r>0 be a constant that does not depend on ε\varepsilon. Then there exists a constant c=c(l,r)>0c=c(l,r)>0 that only depends on ll and rr such that the following holds: For any a>0a>0, s>0s>0 satisfying asl≥ras^{l}\geq r and ε2βls2≤1\varepsilon^{2\beta_{l}}s^{2}\leq 1, we have

Otherwise, 1−ε2βls2≥1/21-\varepsilon^{2\beta_{l}}s^{2}\geq 1/2, and consequently

where the last line follows from the AM-GM inequality. This completes the proof. ∎

Appendix C Proofs of Theorem 2 and 3: learning with projected SGD

We will prove Theorem 2 which bounds the distance between GF and projected SGD in sub-Sections C.1 through C.3, with sub-Section C.4 devoted to the proof of Theorem 3. Throughout this section, we use MM to refer to any constant that only depends on the MiM_{i}’s from Assumptions A1-A3, whereas the value of MM can change from line to line. We start with an elementary lemma that establishes the Lipschitz continuity of the gradient flow trajectory:

First, notice that along the trajectory of gradient flow, the risk must be non-increasing. In fact, we have

where the last line follows from our assumption. Since ∣ai(0)∣≤M|a_{i}(0)|\leq M, we know that ∣ai(t)∣≤M(1+t/ε)|a_{i}(t)|\leq M(1+t/\varepsilon), and ∣ai(t)−ai(s)∣≤ε−1M(t−s)|a_{i}(t)-a_{i}(s)|\leq\varepsilon^{-1}M(t-s). Moreover, according to Eq. (5), we have

In what follows we define two discretized versions of Eq.s (4) and (5), namely the gradient descent (GD) and stochastic gradient descent (SGD) dynamics. They will serve as important intermediate objects for our proof.

where we recall from Eq.s (102) and (103):

By convention, we have V(s)=V(s;1,1)V(s)=V(s;1,1) and U(s)=U(s;1,1)U(s)=U(s;1,1) for s∈s\in.

The iteration equations for one-pass SGD read:

Note that Eq. (151) can also be written as:

For notational simplicity, we denote θi(t)=(ai(t),ui(t))\theta_{i}(t)=(a_{i}(t),u_{i}(t)) for i∈[m]i\in[m] and t≥0t\geq 0, and

and Hε(θ,ρ)=(ε−1F(θ,ρ),G(θ,ρ))H_{\varepsilon}(\theta,\rho)=(\varepsilon^{-1}F(\theta,\rho),G(\theta,\rho)). Then, Eq.s (4) and (5) and Eq. (150) can be rewritten as

respectively. The lemma below will be used several times in the proof.

Denoting ρ(m)=(1/m)∑i=1mδθi\rho^{(m)}=(1/m)\sum_{i=1}^{m}\delta_{\theta_{i}} and ρ′(m)=(1/m)∑i=1mδθi′\rho^{\prime(m)}=(1/m)\sum_{i=1}^{m}\delta_{\theta^{\prime}_{i}}. If ∥ui∥2≤C\left\|{u_{i}}\right\|_{2}\leq C and ∥ui′∥2≤C\left\|{u^{\prime}_{i}}\right\|_{2}\leq C for all i∈[m]i\in[m] (CC is any fixed absolute constant, for example, here we can take C=2C=2), then we have

where the constant MM only depends on the MiM_{i}’s. As a consequence, we obtain that

Second, using again triangle inequality, we deduce that

This completes the proof of Lemma 5, since the “as a consequence” part follows naturally from the upper bounds obtained earlier. ∎

Following the notation and assumption of Lemma 5, we have

By definition of the risk function and triangle inequality, we deduce that

For any s∈[0,t]s\in[0,t], by Lemma 4 and 5 we have (denote [s]=η⌊s/η⌋[s]=\eta\lfloor s/\eta\rfloor, and notice that we can take C=2C=2 since t≤TΔt\leq T_{\Delta})

Using again Lemma 4 and 5, we obtain that

For s≤t≤TΔs\leq t\leq T_{\Delta}, we have Δ(s)2≤Δ(s)\Delta(s)^{2}\leq\Delta(s). Hence,

Therefore, for all T≥0T\geq 0 and η≤1/(Mexp⁡((ε−1+1)MT(1+ε−1T)2))\eta\leq 1/(M\exp((\varepsilon^{-1}+1)MT(1+\varepsilon^{-1}T)^{2})), we have

This proves T≤TΔT\leq T_{\Delta}, and consequently

Finally, with the aid of Lemma 6, we get the following upper bound on the difference between the risk of gradient flow and gradient descent:

There exists a constant MM that only depends on the MiM_{i}’s, such that for any T≥0T\geq 0 and

C.2 Difference between GD and SGD

The proof for this section is almost identical to Appendix C.5 in [Mei et al., 2019]. The only difference is that, here we need to verify that (Id−uu⊤)σ′(⟨u,x⟩)x(I_{d}-uu^{\top})\sigma^{\prime}(\langle u,x\rangle)x is an M3M_{3}-sub-Gaussian random vector. This follows from the identity (Id−uu⊤)x=x−⟨u,x⟩u(I_{d}-uu^{\top})x=x-\langle u,x\rangle u and Assumption A3. We thus obtain the following interpolation bound between GD and SGD:

There exists a constant MM that only depends on the MiM_{i}’s, such that for any T,z≥0T,z\geq 0 and

the following happens with probability at least 1−exp⁡(−z2)1-\exp(-z^{2}): For all t∈[0,T]t\in[0,T], we have

C.3 Difference between SGD and projected SGD

The aim of this section is to prove a coupling bound between the trajectory of SGD and that of projected SGD, thus finally leading to an upper bound on the difference between the risk of projected gradient flow and projected SGD. To begin with, let us fix T,z≥0T,z\geq 0 and choose

as in Theorem 2, where MM is a large enough constant (to be determined later). Define

then for k≤min⁡(T,Tθ)/ηk\leq\min(T,T_{\theta})/\eta and i∈[m]i\in[m], we have (note that here t=kηt=k\eta)

Denoting Fk=σ(θˉ(0),z1,⋯ ,zk)\mathcal{F}_{k}=\sigma(\bar{\theta}(0),z_{1},\cdots,z_{k}), we know from Assumption A3 that, conditioning on Fk\mathcal{F}_{k}, σ′(⟨uˉi(k),xk+1⟩)xk+1\sigma^{\prime}(\langle\bar{u}_{i}(k),x_{k+1}\rangle)x_{k+1} is an M3M_{3}-sub-Gaussian random vector. By well-known results on Euclidean norm of sub-Gaussian random vectors (see, e.g., Jin et al. ), we know that there exists a constant MM satisfying

Choosing δ=ηexp⁡(−z2)/(mT)\delta=\eta\exp(-z^{2})/(mT) and applying a union bound gives

Therefore, with probability at least 1−exp⁡(−z2)1-\exp(-z^{2}), for all k≤min⁡(T,Tθ)/ηk\leq\min(T,T_{\theta})/\eta and i∈[m]i\in[m], we have

Hence, we deduce from the definition of Δi(k)\Delta_{i}(k) that

thus leading to (using the same argument as in the proof of Lemma 5)

Moreover, by (conditional) sub-Gaussianity of the G^i\widehat{G}_{i}’s, we know that

Combining the above estimates, it then follows that

Using the same proof technique as in Appendix C.5 of Mei et al. , we conclude that

Similarly as in the proof of Theorem 4, we define

Then, for l≤min⁡(T,Tθ,TΔ)/ηl\leq\min(T,T_{\theta},T_{\Delta})/\eta, we have

Proceeding with the same argument, it follows that

Applying Grönwall’s inequality (discrete version) yields that

Combining the above estimates gives the following:

There exists a constant MM that only depends on the MiM_{i}’s, such that for any T,z≥0T,z\geq 0 and

the following happens with probability at least 1−exp⁡(−z2)1-\exp(-z^{2}): For all t∈[0,T]t\in[0,T], we have

Theorem 2 then follows as a result of combining Theorem 4, Theorem 5, and Theorem 6.

C.4 Proof of Theorem 3

By our assumption, we know that the standard learning scenario holds up to level LL, and that

Then, according to Definition 1, there exists ε∗=ε∗(δ)\varepsilon_{*}=\varepsilon_{*}(\delta), T0=T0(δ)T_{0}=T_{0}(\delta) such that for all ε≤ε∗\varepsilon\leq\varepsilon_{*} and T=T(ε,δ)=T0(δ)ε1/2LT=T(\varepsilon,\delta)=T_{0}(\delta)\varepsilon^{1/2L}, one has

Moreover, from Section 4 we know that with probability at least 1−e−C′m1-e^{-C^{\prime}m} over the i.i.d. initialization,

where M′M^{\prime} only depends on (σ,φ,PA)(\sigma,\varphi,{\rm P}_{A}). Now we choose ε≤ε∗\varepsilon\leq\varepsilon_{*} and T=T(ε,δ)=T0(δ)ε1/2LT=T(\varepsilon,\delta)=T_{0}(\delta)\varepsilon^{1/2L}. It then follows that

According to Theorem 2, we know that with probability at least 1−exp⁡(−z)1-\exp(-z),

with n=T/η=T(ε,δ)/ηn=T/\eta=T(\varepsilon,\delta)/\eta. We now take

Then, by our choice of mm and dd, we know that \mathscrsfsR(a(T),u(T))≤2δ/3\mathscrsfs{R}(a(T),u(T))\leq 2\delta/3. Further, taking

The above happens with probability 1−exp⁡(−C′m)−exp⁡(−z)1-\exp(-C^{\prime}m)-\exp(-z). Hence, our conclusion follows naturally from the assumption m≥zm\geq z.

Appendix D Counterexamples to the standard learning scenario

For any fixed (a,u)=(ai,ui)1≤i≤m(a,u)=(a_{i},u_{i})_{1\leq i\leq m}, we have

Moreover, the risk is always lower bounded by

We consider the reduced mean-field equations (23):

Note that if φ0=φ1=0\varphi_{0}=\varphi_{1}=0, then V′(s)=s⋅v(s)V^{\prime}(s)=s\cdot v(s) for some continuous function vv. Denoting a=(a1,⋯ ,am)a=(a_{1},\cdots,a_{m}) and s=(s1,⋯ ,sm)⊤s=(s_{1},\cdots,s_{m})^{\top}, the above equation regarding the evolution of the sis_{i}’s can be written as

where A(a,s)A(a,s) is a matrix-valued function satisfying

Using the similar a priori estimate as in the proof of Lemma 1, we can show that

for any finite time TT, which immediately implies that s(t)≡0s(t)\equiv 0 for t∈[0,T]t\in[0,T]. Therefore, we won’t be able to learn any component of φ\varphi with degree ≥1\geq 1.

We may assume σk≠0\sigma_{k}\neq 0, and analyze the simplified ODE system (91), which reduces to

Therefore, most of the neurons cannot evolve to the magnitude of Ω(1)\Omega(1) in the process of learning the kk-th component, and therefore fails to provide an effective initialization for learning the next component φk+1\varphi_{k+1}.