Phase diagram of Stochastic Gradient Descent in high-dimensional two-layer neural networks
Rodrigo Veiga, Ludovic Stephan, Bruno Loureiro, Florent Krzakala, Lenka Zdeborová
Introduction
Descent-based algorithms such as stochastic gradient descent (SGD) and its variants are the workhorse of modern machine learning. They are simple to implement, efficient to run and most importantly: they work well in practice. A detailed understanding of the performance of SGD is a major topic in machine learning. Quite recently, significant progress was achieved in the context of learning in shallow neural networks. In a series of works, it was shown that the optimisation of wide two-layer neural networks can be mapped to a convex problem in the space of probability distributions over the weights . This remarkable result implies global convergence of two-layer networks towards perfect learning provided that the number of hidden neurons is large, the learning rate is sufficiently small and enough data is at disposition. This line of work is commonly referred to as the mean-field or the hydrodynamic limit of neural networks. Mathematically, these works showed that one could describe the entire dynamics using a partial differential equation (PDE) in dimensions.
In a different, and older, line of work one-pass SGD for two-layer neural networks with a finite number of hidden units, synthetic Gaussian input data and teacher-generated labels has been widely studied starting with the seminal work of . These works consider the limit of high-dimensional data and show, in particular, that the stochastic process driven by gradient updates converge to a set of deterministic ordinary differential equations (ODEs) as the input dimension and the learning rate is proportional to . The validity of these ODEs in this limit was proven by . However, the picture drawn from the analysis of these ODEs is slightly different from the mean-field/hydrodynamic picture: in this case SGD can get stuck for long time in minima associated to no specialization of the hidden units to the teacher hidden units, and even when it converges to specializing minima, it fails to perfectly learn (i.e. to achieve zero population risk). In fact, in this analysis, the interplay between the limit of the learning rate going to zero and appeared to be fundamental.
One should naturally wonder about the link between these two sets of works with, on the one hand a -dimensional PDE (with large ), and on the other a -dimensional ODE (with large ). In this work we aim to build a bridge between these two approaches for studying one-pass SGD.
Our starting point is the framework from , which we build upon and expand to a much broader range of choices of learning rate, time scales, and hidden layer width. This allows us to provide a sharp characterisation of the performance of SGD for two-layer neural networks in high-dimensions. We show it depends on the precise way in which the limit is taken, and in particular on how the quantity of data, the hidden layer width, and the learning rate scale as . For different choices of scaling, we can observe scenarios such as perfect learning, imperfect learning with an unavoidable error, or even no learning at all.
As a consequence of our analysis, we provide a phase diagram (see Figure 1(a)) describing the possible scenarios arising in the high-dimensional setting. Our main contributions are as follow:
We rigorously show that the dynamics of SGD can be captured by a set of deterministic ODEs, considerably extending the proof of to accommodate for general time scalings defined by an arbitrary learning rate, and a general range of hidden layer width. We provide much finer non-asymptotic guarantees which are crucial for our subsequent analysis.
From the analysis of the ODEs, we derive a phase diagram of SGD for two-layer neural networks in the high-dimensional input layer limit . In particular, scaling both the learning rate and hidden layer width with the input dimension as
we identify four different learning regimes which are summarized in Figure 1(a):
Perfect learning (green region, ): we show that perfect learning (zero population risk) can be asymptotically achieved with samples even for tasks with additive noise.
Plateau (blue line ): learning reaches a plateau related to the noise strength. The point goes back to the classical work of .
Bad learning (orange region ): here the noise dominates the learning process.
No ODEs (red region ): the stochastic process associated to SGD is not guaranteed to converge to a set of deterministic ODEs. This region is thus outside the scope of our analysis.
To better illustrate this phase diagram we present in Figure 1(b) a solution of the ODEs in all three regimes.
Deterministic dynamical descriptions of one-pass stochastic gradient descent in high-dimensions have a long tradition in the statistical physics community, starting with single- and two-layer neural networks with few hidden units . The seminal work by overcame previous limitations by constructing a set of deterministic ODEs for two-layer networks with any finite number of hidden units, paving the way for a series of important contributions . This line of work corresponds to the case of Figure 1(a). One of our goal is to generalize this picture beyond fixed hidden layer size and learning rate.
Reproducibility
A code is provided at https://github.com/rodsveiga/phdiag_sgd.
Setting
Typically, one minimizes the empirical risk over the full data set. Instead, learning with one-pass gradient descent minimizes directly the population risk:
Given a single sample the weights are updated sequentially by the gradient descent rule:
with and . The parameter is the learning rate. Despite being a simplification with respect to batch learning, one-pass gradient descent is an amenable surrogate for the theoretical analysis of non-convex optimization, since at each step the gradient is computed with a fresh data sample, which is equivalent to performing SGD directly on the population risk.
In particular, in this manuscript we assume realizability , and focus our analysis on the square loss , leading to
and the population risk is completely determined by the macroscopic state:
The training dynamics (6) defines a discrete-time stochastic process for the evolution of the overlap matrix
with fixed and and updated as:
with , , and . In what follows, we will make the concentration assumption ; this will be justified in the proof of Theorem 3.1.
We emphasize in (14) the specific role played by each term in the right hand-side. The "learning" terms are the fundamental ones, that actually drive the learning of the teacher by the student. We show in Appendix C.3 that these "learning" terms are identical to those obtained in the gradient flow approximation of SGD, whose performance is the topic of many works . Those are precisely the terms that draw the population risk towards zero. However, in our setting there is an additional variance term (so that this flow approximation is incomplete) that corresponds to the fluctuations of around its expected value . In particular, this is where the effects of the noise can be felt. These terms were sometimes denoted as () and () in . We shall see that the additional "variance" term is the one responsible for the plateau in the critical (blue) region of Figure 1(a), while its contribution vanishes in the perfect learning (green) region.
Additionally, albeit our work particularizes to Gaussian input data, we believe our conclusion, and the phase diagram discussed in Figure 1(a), to hold beyond this restricted case. Indeed, while the Gaussian assumption is crucial to reach a particular set of ODEs and their analytic expression, the approach can be applied to more complex data distribution, as long as one can track the sufficient statistics required to have a closed set of equations. For instance, obtained very similar equations for an arbitrary mixture of Gaussians – that would obey the same scaling analysis as ours – while proved that many complex distributions behave as Gaussians in high-dimensional setting, including, e.g. realistic GAN-generated data. We thus expect our conclusions to be robust in this respect.
Main results
Although seems to be the most natural time scaling in the high-dimensional limit , if and are allowed to vary with the right-hand side (RHS) of Eqs. (14) can diverge and render the ODE approximation obsolete. Instead, for a given time scaling , we can rewrite Eqs. (14) as
In Theorem 3.1 we prove that as , converges to the solution of the ODE:
the time scaling satisfies for some constant ,
the activation function is -Lipschitz,
Then, there exists a constant (depending on ) such that for any , the following inequality holds:
Our proof is based on techniques introduced in (namely, their Lemma 2) which studies a different problem with related proof techniques. The proof involves decomposing as
where the two first terms can be considered as a deterministic discrete process, and the last term is a martingale increment. The main challenge lies in showing that the martingale contribution stays bounded throughout the considered time period.
Although the method is similar to , there are a number of differences between the two approaches. First, our proof fixes a number of holes in , in particular bounding by a sufficiently slowly diverging function of . Additionally, the techniques used in this paper yield a dependency in that is nearly negligible, while the previous methods imply bounds that are much too coarse for our needs.
The function can be computed explicitly for various choices of , which allows to check Assumption 3 directly. We provide in Appendix C the necessary computations for ; those for the ReLU unit can be found in . It can be checked that in the ReLU case, the function is not Lipschitz around the matrices satisfying
for any . Since the square root function is Lipschitz whenever the eigenvalues of are bounded away from zero (see e.g. ), Assumption 3 is implied by the condition
however, this assumption is much stronger, and becomes unrealistic in the specialization phase (as well as when ).
Theorem 3.1 allows us to safely navigate through Figure 1(a) by keeping track of convergence rates of the discrete process to a set ODEs. The interplay between learning rate and hidden layer width defines the time scaling and the trade-off between the linear contribution on and the quadratic one, playing a central role on whether the network achieves perfect learning or not. Specifically, consider the following learning rate and hidden layer width scaling with :
where we have chosen without loss of generality.
When , this implies that either the learning term or the noise term scale like a negative power of , and is negligible with respect to the other term. It is then easy to check that at a finite time horizon , the resulting ODEs behave as if the negligible term was not present. We refer to Theorem B.1 in the appendix for a quantitative proof of this phenomenon. Let us now describe the different regimes depicted in Figure 1(a).
When and are scaled such that , Eqs. (21) converge to
with . This regime is an extension of for which . The convergence rate to the ODEs scales with , and the phenomenology we observe for is consistent with previous works studying the setting ; namely the existence of an asymptotic plateau proportional to the noise level. For instance, the asymptotic population risk is known to be proportional to when and the dynamics is driven by a rescaled version of Eqs. (23). Since the noise term does not vanish under this scaling, perfect learning to zero population risk is not possible. There is always an asymptotic plateau related to the noise level , and the learning rate .
Green region (perfect learning) –
If we can define the time scaling . By Theorem 3.1, Eqs. (21) converge to the following deterministic set of ODEs:
at a rate proportional to , where we have highlighted that the noise term vanishes with . Hence, as long as the noise does not play any role on the dynamics. This setting could be understood by taking an effect learning rate on , which leads to zero population risk, i.e. perfect learning, in the high dimensional limit . We validate this claim by a finite size analysis in the next section.
As discussed, the time scaling determines the number of data samples required to complete one learning step on the continuous scale. The bigger , the more attenuated the noise term, thus the closer to perfect learning. The trade-off is that the bigger , the larger the number of samples needed is, since . Given a realizable learning task, one would thus rather choose the parameters to attain the perfect learning region, but being as close as possible to the plateau line for not increasing too much the needed number of samples. We remark that provides an alternative deterministic approximation in this regime, with non-asymptotic bounds, whenever ; this is the so-called mean-field approximation, with known convergence guarantees .
Orange region (bad learning) –
We now step in the unusual situation where the learning rate grows faster with than the hidden layer width: . In this case, by (22) the noise term dominates over the dynamics. Defining the time scaling , we have
According to Theorem 3.1 the convergence rate of Eqs. (21) to Eqs. (25) scales with . Therefore the existence of the noisy ODEs above is circumscribed to the region
and presents a convergence trade-off absent in the other regimes: the faster one of the contributions of Eqs. (21) goes to zero, the worse is the convergence rate. In the present case, the more the learning term is attenuated, i.e. the more negative is , the worse the dynamics is described by Eqs. (25). Although the weights are updated, the correlation between the teacher and the student weights parametrized by the overlap matrix remains fixed on its initial value , which is a fixed point of the dynamics under this scaling. Unsurprisingly, this leads to poor generalization capacity.
Red region (no ODEs) –
If , the stochastic process driven by the weight dynamics does not converge to deterministic ODEs under the assumptions of Theorem 3.1. We are then not able to state any claim about this regime.
Initialization and convergence –
There are two additional features worth commenting on the high-dimensional dynamics and its connection to the mean-field/hydrodynamic approach, regarding initialization and the specialization transition.
In the ODE approach we discuss here, we always observe a first plateau where the teacher-student overlaps are all the same. This means all the hidden layer neurons learned the same linear separator. At this point, the two-layer network is essentially linear. This is called a unspecialized network in . In fact, this is a perfectly normal phenomenon, as with few samples even the Bayes-optimal solution would be unspecialized . Only by running the dynamics long enough the student hidden neurons start to specialize, each of them learning a different sub-function so that the two-layer network can learn the non-trivial teacher.
Let us make two comments on this phenomenon: (i) while the "linear" learning in the unspecialized regime may remind the reader of the linear learning in the lazy regime of neural nets, the two phenomena are completely different. In lazy training, the learning is linear because weights change very little, so that the effective network is a linear approximation of the initial one. Here, instead, the weights are changing considerably, but each hidden neuron learns essentially the same function. (ii) If the ODEs are initialized with weights uncorrelated with the teacher, then the unspecialized regime is a fixed point of the ODEs: the student thus never specializes, at any time. Strikingly, such condition arises as well in the analysis of mean-field equations (see e.g. Theorem 2 in that discusses the need to have spread initial conditions with a non-zero overlap with the teacher) to guarantee global convergence.
This raises the question about the precise dependence of the learning on the initialization condition in the high-dimensional regime, where a random start gets a vanishing () overlap. This is a challenging problem that only recently has been studied (though in a simpler setting) in who showed it yields an additional time-dependence. Generalizing these results for high-dimensional two-layer nets is an open question which we leave for future work.
Discussion, special cases, and simulations
To illustrate the phase diagram of Figure 1(a), we present now several special cases for which we can perform simulations or numerically solve the set of ODEs. Henceforth, we take , for which the expectations of the ODEs and of the population risk, Eq. (12), can be calculated analytically . The explicit expressions are presented in Appendix C. Teacher weights are such that . The initial student weights are chosen such that the dimension can be varied without changing the initial conditions , , and consequently the initial population risk . A detailed discussion can be found in Appendix D.
We start by recalling the well-known setting characterized by the point . The convergence of the stochastic process for fixed learning rate and hidden layer width to Eqs. (23) was first obtained heuristically by . In Figure 2 we recall this classical result by plotting the population risk dynamics for different noise levels. Dots represent simulations, while solid lines are obtained by integration of the ODEs, Eq. (23).
Learning is characterized by two phases after the initial decay. The first is the unspecialized plateau where all the teacher-student overlaps are approximately the same: . Waiting long enough, the dynamics reaches the specialization phase, where the student neurons start to specialize, i.e., their overlaps with one of the teacher neurons increase and consequently the population risk decreases. This specialization is discussed extensively in . If , the population risk goes asymptotically to zero. Instead, if , the specialization phase presents a second plateau related to the noise .
The asymptotic population risk related to the second plateau is proportional to in the high-dimensional limit with finite. As mentioned in the previous section, the expectation over in Eq. (23a) prevents one from obtaining zero population risk for a noisy teacher.
2 Perfect learning for κ=0𝜅0\kappa=0
In this section we study the line with of Figure 1(a), for which Eqs. (24) with hold. We show that perfect learning can be asymptotically achieved in the realizable setting for any finite hidden layer width . Keeping and fixed, we have done simulations increasing the input layer dimension . In Figure 3(a) we set , and vary the input layer dimension. The bigger is, the closer we are to the ODE-derived noiseless result.
Gathering the asymptotic population risk from simulations for varying and we perform a finite-size analysis to study the dependence of with . This shows that the noise term goes to zero under this setting. In Figure 3(b) we plot versus from simulations (dots) for different noise levels. We fit lines under the log-log scale showing that , as expected. Figure 4 draws the same conclusion for .
As already stated, the interplay between the exponents directly affects the time scale. We end this subsection by graphically illustrating this fact through simulations. Setting the noise to we compare the cases in Figure 5(a). All simulations are rendered on the scale to illustrate the trade-off between asymptotic performance and training time.
3 Bad learning for κ=0𝜅0\kappa=0
We now quickly discuss the uncommon case of growing with within the orange region. In Figure 5(b) we compare simulations varying with the solution of the ODEs given by Eqs. (25). Both lead to poor results compared to the green and blue regions. Moreover, this regime presents strong finite-size effects, making it harder to observe the asymptotic ODEs at small sizes. However, the trend as increases is very clear from the simulations. As discussed in Section 3, the more the learning term is attenuated on the ODEs, the worse they describe the dynamics.
4 Large hidden layer: κ>0𝜅0\kappa>0
Finishing our voyage through Figure 1(a) with examples, we briefly discuss the case where both input and hidden layer widths are large. Although Theorem 3.1 provides non-asymptotic guarantees for , the number of coupled ODEs grows quadratically with , making the task of solving them rather challenging. Thus, we present simulations that illustrate the regions of Figure 1(a). Fixing we show in Figure 6 learning curves for different values of and . The colors are chosen to match their respective regions in the phase diagram.
Due to the relatively small sizes used in Figure 6, the green dots seem to decrease towards perfect learning, even when , provided that is large enough, as is predicted by the phase diagram in Figure 1(a). Moreover, since is not large enough, when the parameters are within the orange region the finite-size effects actually dominates, similarly to Figure 5(b). The learning contribution still plays a role and the asymptotic population risk is similar to the case . Within the red region, which is out of scope of our theory, the simulation gets stuck on a plateau with larger population risk.
Conclusion
Building up on classical statistical physics approaches and extending them to a broad range of learning rate, time scales, and hidden layer width, we rendered a sharp characterisation of the performance of SGD for two-layer neural networks in high-dimensions. Our phase diagram describes the possible learning scenarios, characterizing learning regimes which had not been addressed by previous classical works using ODEs. Crucially, our key conclusions do not rely on an explicit solution, as our theory allows the characterization of the learning dynamics without solving the system of ODEs. The introduction of scaling factors is non-trivial and has deep implications. Our generalized description enlightens the trade-off between learning rate and hidden layer width, which has also been crucial in the mean-field theories.
Acknowledgements
We thank Gérard Ben Arous, Lenaïc Chizat, Maria Refinetti and Sebastian Goldt for discussions. We acknowledge funding from the ERC under the European Union’s Horizon 2020 Research and Innovation Program Grant Agreement 714608- SMiLe. RV was partially financed by the Coordenação de Aperfeiçoamento de Pessoal de Nível Superior - Brasil (CAPES) - Finance Code 001. RV is grateful to EPFL and IdePHICS lab for their generous hospitality during the realization of this project.
Appendix
Appendix A Deterministic scaling limit of stochastic processes
In order to show the deterministic scaling of online SGD under a proper chosen time scale, we will make use of a convergence result by , which is adapted below in Theorem A.1.
Consider a -dimension discrete time stochastic process sequence, for some . The increment is assumed to be decomposable into three parts,
Let , with , be a continuous stochastic process such that with . Define the deterministic ODE
with .
where is the solution of Eq.(A.2).
The reader interested in the proof is referred to the supplementary materials of . ∎
Although the theorem wasn’t originally proven in the setting, a glance at its proof shows that it still holds upon replacing by in Assumption A.1.1 and A.1.2, as well as Equation (A.3). We choose to be the norm, since it suits better the scaling. The in Theorem A.1 corresponds to , where is defined in Theorem 3.1.
The functions are similarly defined on . With that, we write
The main obstacle to bounding and is the fact that the can a priori diverge to infinity. Our first task is therefore to show that this does not happen; as a proxy we show a subgaussian-like moment bound:
Since is -Lipschitz, we have by the Cauchy-Schwarz inequality
where are absolute constants. Summing those inequalities yield
As a result, we have for any
For simplicity, let denote any of the . We have, for all ,
where the remainder term has bounded expectation. Again, we write
By Assumption 3, the are bounded from below by a constant, hence
This implies that for any and ,
A.2 Assumption A.1.1
The term in is bounded by the same techniques as the last section. For the second term,
using (A.6). Choosing shows that
A similar bound holds for the , and hence
which implies Assumption A.1.1 with and .
A.3 Assumption A.1.2
Since is Lipschitz, for any
The first expectation is the variance of a random variable, which is equal to , and the second expectation is bounded by the same methods as the above sections. The term in brackets is therefore bounded by , and
Finally, since for any we have , letting we find
hence Assumption A.1.2 is true with .
A.4 √square-root\surd-Lipschitz property
The same arguments as above show that the function is Lipschitz, and hence for some constant we have
Appendix B A lemma on ODE perturbation
In this section, we prove a proposition that bounds the difference between an ODE solution and a perturbed version, for a bounded time .
where , and with the initial condition . Then, if is fixed, we have
for any , with a constant independent from .
Before proving this proposition, we begin with a small lemma:
with . Then, for some constant , we have
Upon considering the function instead, we can assume that . Then, we have
where and are ad hoc constants. The lemma then follows from adjusting the constant as needed. ∎
We are now in a position to show Theorem B.1:
Assume for simplicity that . We begin by bounding ; we have
Applying Lemma B.2 and taking square roots on each side,
for any . Now, similarly,
having used (B.1) on the last line. This is again the setting of Lemma B.2, which gives
Appendix C Expectations over the local fields
In this appendix we present the explicit expressions from the expectations of the local fields used to compute the population risk and the ODE terms.
For the expectations in Eqs. (C.2) are in general given by
where is an element of the covariance matrix given by Eq. (C.3).
Explicitly, the population risk contributions are
C.2 ODE contributions
From the update equations, we first consider the expectations linear in :
For the expectations in Eqs. (C.6) are given by
where is an element of the covariance matrix given by Eq. (C.7). As examples, we write explicitly:
The quadratic contribution in is given by
The solution of the noise-dependent term can be constructed with the covariance matrix (C.3) and is given by
For the expectations in Eqs. (C.10) are given by
C.3 From gradient flow to local fields
Now, since for any , we have
Recalling the definition , the terms present inside the expectation are exactly those in the learning term of Eq. (14).
Appendix D Initial conditions and symmetric teacher
with . Using Eq. (D.1) one can write
We stress that we use these initial conditions to make the data comparable for varying dimension in the numerical illustrations. Our conclusions do not depend on this particular choice of initial conditions. If one simply takes random initialization for each , the full picture we have presented in this manuscript remains unchanged. In Figure 7 we present an example of curves within the blue region (see Section 3 for the characterization of this regime) with unconstrained Gaussian initialization. Dots represent simulations, while solid lines are obtained by integration of the ODEs given by Eqs. (23), with initial conditions adjusted to match simulations.
Although varying the initial population risk with slightly changes the exact position where the specialization transition starts, the particular initial conditions adopted in this work do not affect whether the specialization transition takes place or not, comparing to unconstrained Gaussian initialization.