High-dimensional limit theorems for SGD: Effective dynamics and critical scaling
Gerard Ben Arous, Reza Gheissari, Aukosh Jagannath
Part I Introduction and main results
Stochastic gradient descent (SGD) is the go-to method for large-scale optimization problems in modern data science. It is often used to train complex parametric models on high-dimensional data. Since its introduction in , there has been a tremendous amount of work in analyzing its evolution.
In fixed dimensions, the asymptotic theory of SGD, and stochastic approximations more broadly, is by now classical. There have been works on path-wise limit theorems, such as functional central limit theorems and even large deviations principles . At the core of this line of work is the idea that in the limit where the step-size, or learning rate, tends to zero, the trajectory of SGD with a fixed loss function (appropriately rescaled in time) converges to the solution of gradient flow for the population loss with the same initialization. Recently there has been considerable interest in quantifying the rate of this trajectory-wise convergence to higher order, in terms of a diffusion approximation. Namely, there are many works developing asymptotic expansions of the trajectory in the learning rate . Motivated by this, there is a rich line of work bounding the time to equilibrium for the associated diffusion approximation (as well as Langevin–type modifications) under uniform ellipticity assumptions . There is also an interesting line of work obtaining PDE limits in the “shallow network” regime where the dimension of the parameter space diverges but the dimension of the data remains constant: see e.g., .
In recent years, there has been considerable interest in understanding the high-dimensional setting, where one is constrained in the amount of data or the run-time of the algorithm due to the high-dimensional nature of the data and the complexity of the model being trained. In these regimes, one cannot simply take the learning rate to be arbitrarily small as this would force an unlimited sample size and run-time. This is a common issue in high-dimensional statistics and the standard analytic approach is to study regimes where the sample size scales with the dimension of the problem .
For SGD with constant learning rate, there has been recent progress on quantifying the dimension dependence of the sample complexity for various tasks on general (pseudo or quasi-) convex objectives and special classes of non-convex objectives . There has also been important work on scaling limits as the dimension tends to infinity for the specific problems of linear regression , Online PCA , and phase retrieval from random starts, and teacher-student networks and two-layer networks for XOR Gaussian mixtures from warm starts. We also note that the study of high-dimensional regimes of gradient descent and Langevin dynamics have a history from the statistical physics perspective, e.g., in .
We develop a unified approach to the scaling limits of SGD in high-dimensions with constant learning rate that allows us to understand a broad range of estimation tasks. One of course cannot develop a high-dimensional scaling limit for the full trajectory of SGD as the dimension of the underlying parameter space is growing. On the other hand, in practice, one is rarely interested in the full trajectory; instead one typically tracks the trajectory of various summary statistics of the algorithm’s evolution, such as the loss, the amplitude of various weights, or correlations between the classifier and the ground truth (in a supervised setting). We show in Theorem 2.3 that under mild regularity assumptions, the evolution of these summary statistics converges as the dimension grows to the solution of a system of (possibly stochastic) differential equations. These effective dynamics depend dramatically on the initializations (warm vs. random or cold), the parameter regions in which one is developing the scaling limit, and the scaling of the step-size with the dimension.
In practice, SGD often exhibits two types of phases in training: ballistic phases where the summary statistics macroscopically change in value, and diffusive phases, where they fluctuate microscopically. (During training, the evolution can start with either, and can even alternate multiple times between these phases.) Our approach allows us to develop scaling limits for both types of phases.
In ballistic phases, the effective dynamics are given by an ordinary differential equation (ODE) and the finite-dimensional intuition that the summary statistics evolve under the gradient flow for the population loss is correct provided the (constant) learning rate is sufficiently small in the dimension. When the learning rate follows a certain critical scaling—matching scalings commonly used in the high-dimensional statistics literature—an additional correction term appears. At this critical scaling, the phase portrait deviates significantly from that of the population gradient flow. Furthermore, in microscopic neighborhoods of the fixed points of this ODE, the effective dynamics become diffusive and are given by SDEs which can exhibit a wide range of (possibly degenerate) behaviors. We note that the appearance of the correction term in the ballistic phase was first observed in the setting of teacher-student networks in and very recently investigated in detail in .
As a simple, first example of the departure of the effective dynamics in the critical step-size regime from the classical perspective, we study estimation for spiked matrix and tensor models in Section 3. In these models, the effective dynamics are exactly solvable and when the step-size scales critically with the dimension, in the ballistic phase the dynamics have additional fixed points as compared to the population gradient flow. The stability of these fixed points exhibit sharp transitions at special signal-to-noise ratios. When initialized randomly, the SGD starts in a microscopic neighborhood of an uninformative such fixed point, within which its effective dynamics become diffusive and exhibit a sharp transition between mean-reverting and mean-repellent Ornstein–Uhlenbeck (OU) processes.
To demonstrate our approach on more complex classification tasks typically studied using neural networks, we study a Gaussian mixture model analogue of the classical XOR problem in Section 5. (The XOR problem is arguably the canonical example of a decision boundary requiring at least two-layers to represent .) Here we find that the natural summary statistics are 22 dimensional, and their (ballistic) effective dynamics exhibit a rich phenomenology between some 39 connected fixed point regions of varying topological dimension. Surprisingly, we find that if we initialize the weights of the network randomly (following a Gaussian distribution), then the algorithm will converge to a classifier with macroscopic generalization error with probability and then follow a degenerate diffusion. On the other hand, we demonstrate the benefit of overparametrization, showing that as the width of the second layer grows, the probability of ballistically converging to a Bayes optimal classifier goes to ; this is a mathematically rigorous example of the lottery ticket hypothesis of .
Before delving into the XOR problem, we first analyze the classification of a two component Gaussian mixture model in Section 4. This task is of course best solved using a one-layer network i.e., logistic regression, but with a two-layer network it exhibits some similar phenomenologies to the XOR problem while being more amenable to finer analysis. Here, we again find that if with random initial weights, with probability the SGD will first converge to a classifier with macroscopic generalization error, and then follow a degenerate diffusion in a microscopic neighborhood of that set of unstable fixed points. We demonstrate this both empirically for positive signal-to-noise ratio and theoretically in the limit where the SNR tends to zero after the dimension tends to infinity.
While the above are a few examples that we are able to solve in detail for both their ballistic and diffusive limits, we expect our main theorem to be applicable and lend new insights into a host of other problems including SGD for finite-rank matrix and tensor PCA, and one and two-layer neural networks applied to mixtures of -Gaussians for fixed . We leave this to future investigation. In this paper, we only consider the simplest variant of SGD, namely online SGD; we leave other variants involving batching and re-use to future works.
Main result
To develop a scaling limit, we need some regularity assumptions on the relationship between how the step-size scales in relation to the loss, its gradients, and the data distribution. To this end let
, and ;
We now turn to our second assumption, that the limiting evolution equations for the family of summary statistics chosen close. Define the following first and second-order differential operators,
Alternatively written, and .
In this case we call the effective drift, and the effective volatility.
We are now ready to present our main result. For a function and measure we let denote the push-forward of .
The proof of Theorem 2.3 is provided in Section 6 and can be seen as a version of the classical martingale problem (see ) for high-dimensional stochastic gradient descent. We call the solution to (2.4) the effective dynamics of the summary statistics . The fact that are locally Lipschitz ensures that this solution is unique.
We end this subsection with discussion of the various scalings appearing in Definition 2.1.
Turning to item (2) of Definition 2.1, we comment that the regularity assumptions made on here are less restrictive than uniform Lipchitz assumptions common to the literature. In particular, we do not assume the population loss is Lipschitz everywhere, as we may have that does not cover , nor does it imply uniform smoothness of (and in turn ) as we may (and will) be taking with .
Let us lastly motivate the scalings appearing in item (3), which ensure there is some independence between and the values of and at . As a testbed, suppose that is a random vector with i.i.d. entries all of order . If is a rescaled linear statistic, e.g., then the first bound of item (3) is saturated, and the second of course is trivial due to the linearity of . The second bound is saturated by taking a rescaling of a radial statistic, e.g., , again assuming for maximal simplicity that is an i.i.d. random vector with order one entries. In fact, the second part of item (3) could be dropped at the expense of more complicated diffusion coefficients in limiting SDE’s: see Remark 2.
While we discussed above the reasons for which the various scalings of Definition 2.1 were selected, it is interesting to ask what changes in Theorem 2.3 should certain of the assumptions of Definition 2.1 be violated. Most of the assumed bounds in the definition of localizability are used to establish tightness and ensure higher order terms in Taylor expansions vanish in the limit. In principle the second assumption in item (3) of Definition 2.1 could be dropped; in that case, the same quantity is still ensured to be by the other localizability assumptions. Then Theorem 2.3 would still apply, but the limiting diffusion matrix would be the limit (assuming it exists) of
as opposed to simply the limit of .
in which case, evidently (2.2) holds with . When (2.5) and (2.6) both hold, we call and the population drift, the population corrector, and the diffusion matrix of respectively. From the fixed dimensional perspective, when (2.5) holds, one predicts to solve
with initial data . as this is its evolution under gradient descent on the population loss . Evidently this perspective only applies in the high-dimensional limit of Theorem 2.3 if both the population corrector and the diffusion matrix are zero. We find that for any triple , there is a scaling of the learning rate with below which , and the effective dynamics agree with the population dynamics (2.7) (we call this the sub-critical scaling regime, where the classical perspective applies), and a critical scaling regime in which and may be non-zero, and the high-dimensionality induces non-trivial corrections to . (In the case of teacher–student networks, the terms and can be compared to the “learning" and “variance" terms in Eq. (14a) of .)
To see this, notice that if the triple is -localizable for some , then it is also -localizable for every sequence . If furthermore (2.3) and (2.5)–(2.6) hold for with some and , then these limits also exists for with the same but with . As such, there can be exactly one scaling of with at which or may be non-zero, and for all smaller scales of , the fixed-dimensional perspective of (2.7) applies.Note that if , then limiting may not exist for , so there is no super-critical regime.
2. Ballistic vs. diffusive behavior of effective dynamics
In all of our examples, the diffusion matrix for the effective dynamics of the most natural choice of summary statistics is zero even in the critical scaling regime where . We call this the ballistic limit. In this case, the effective dynamics of the summary statistics is given by the ODE system
In these settings, the phase portrait of the summary statistics is asymptotically that of this flow.
Note that by construction of the scaling limit, the phase portrait of the ballistic limit only describes the evolution of summary statistics on length-scales that are order 1 and number of iterations that are order . If one is then interested in the evolution of in microscopic neighborhoods of the fixed points of the ballistic effective dynamics of (2.8), Theorem 2.3 also allows one to develop separate diffusive limits there.
This then leads to the rescaled effective dynamics of the summary statistics near :
In , it was empirically observed that the best training for neural networks does not occur when step-sizes are small enough for the classical gradient flow approximation to be valid. Instead, it occurs at the edge of stability where the step size is just small enough for the training to remain stable. Here, the loss fluctuates for some time before eventually converging to lower values than it would with smaller step size. This critical step size scaling is defined via the sharpness, namely the largest eigenvalue of the training loss Hessian. For a selection of recent theoretical investigations of this phenomenon see, e.g., .
While sharpness and edge of stability do not have direct analogues in the context of online SGD, a qualitatively similar phenomenon can be seen by taking the population loss as a summary statistic. The critical scaling of the learning rate with dimension discussed in Sections 2.1–2.2 constrains the step size in terms of the top eigenvalue of the Hessian of the loss. With this scaling, the population loss fluctuates near critical regions of its ballistic flow, allowing it to escape the critical region, whereas with a sub-critical learning rate the population loss stays stuck. We leave more detailed investigation of this connection to edge-of-stability phenomena for SGD to future investigation.
Part II Examples
In the following sections, we demonstrate Theorem 2.3 on a range of popular examples of high-dimensional statistical tasks. We begin first in Section 3 by presenting an application to a widely studied problem of high-dimensional estimation: namely, de-noising a rank one tensor that has been corrupted additively by Gaussian noise. We then turn to classification. Our aim in these examples is to demonstrate the applicability of our result to the analysis of multi-layer neural networks. To this end we analyze the training dynamics of a two-layer neural network for two canonical classification tasks, namely classification of a symmetric, binary gaussian mixture model (Section 4) and classification of a Gaussian analogue of the XOR problem of Minsky–Papert (Section 5).
Consider the problem of de-noising a rank one tensor that has been corrupted additively by Gaussian noise via SGD. A popular statistical model of this task is the spiked tensor model . Suppose that we are given i.i.d. samples of data of the form
In the case , this is a version of the well-known spiked matrix model of PCA for which there is, by now, a substantial literature regarding the statistical thresholds. For a necessarily small selection see, e.g., . For related work on online learning in this context see, e.g., . Of particular interest in this direction is the well-known phase transition at for estimation in this problem, which was determined first for Wishart ensembles in and subsequently for this setting in . As we will see in Section 3.3 below, we find a dynamical analogue of this transition at .
The case was introduced by Montanari and Richard as a natural generalization of the spiked matrix models for estimation (and testing) problems where the data has multiple indices or requires higher moments. Here there has been a large literature on the statistical thresholds for estimation and testing, see, e.g., . In this setting, there has also been a tremendous literature on the computational aspects of this problem as it is viewed as a important example of a model with a statistical-computational gap, namely, a setting where there is a gap between the regimes of statistical and computational tractability. See, e.g., .
We begin with these examples as their effective dynamics are particularly simple to analyze. In particular, they are are exactly solvable and only require two summary statistics, a correlation observable and a radial term. Even with this relative simplicity, we encounter a wide range of ODE and SDE limits. In particular, as mentioned above, we find dynamical phase transitions corresponding to the aforementioned thresholds in these models. For our analysis we will focus exclusively on the most interesting, critical step-size scaling which corresponds to the proportional asymptotics regime from the random matrix theory literature.
2. Analysis
We take as loss the (negative) log-likelihoodNote that one might also add additional penalty terms. The case of a ridge penalty is treated in Section 7. namely,
are such that , and the law of only depends on them.
In our normalization with fixed, the regime is sub-critical and the regime is critical. Note that with different scalings of , the critical learning rate changes. We focus our presentation on the most interesting regime, namely the critical scaling regime of for some constant . Recalling the relation between number of samples and step-size, we see that this regime corresponds to the proportional asymptotics regime most studied in the random matrix theory literature where the above-mentioned transition for the top eigenvalue occurs. Note, however, that the limits in the subcritical regime are in all cases recovered by taking the limits of the ODE’s/SDE’s of the critical regime.
For notational simplicity, let . We consider the pair , for which Theorem 2.3 yields the following effective dynamics.
Fix , , and let . Then converges as to the solution of the following ODE initialized from :
We are able to identify and classify the set of fixed points of this effective dynamics. We focus on the critical step-size regime with where one sees from (3.1) that , where the problem in the matrix case is most directly related to an eigenvalue problem (see Section 7 for the generic dependencies). Throughout the following, we use the following notion of stability/unstability of a set of fixed points of an ODE.
We call a set of fixed points for an ODE stable if for every , for every , the solution of the ODE with initialization converges to some point in as . Otherwise, we call unstable.
Eq. (3.1) has isolated fixed points classified as follows. Let be as in (7.6) and be as in (7.7) (if , and ):
An unstable fixed point at and a fixed point at ; if , is stable if and unstable if ; if is always stable.
If : when , two stable fixed points at . When , two unstable fixed points at and two stable fixed points at .
The presence of two pairs of fixed points when with non-zero correlation with may seem surprising—indeed it indicates that even some warm starts will fail to attain good correlation with the signal when is finite. This is an interesting consequence of the corrector in (3.1) and if one tracks the dependence in the above, the fixed point goes to zero as and this barrier to recovery from warm starts vanishes as one approaches sub-critical step-sizes.
3. A dynamical analogue of the BBP transition
4. On the sample complexity of tensor PCA
which transitions between mean-reverting and mean-repellent at , as in .
We only considered a few specific choices of summary statistics in the above, and the strength of Theorem 2.3 derives from its general applicability. As demonstrations, let us mention a few other examples that we would expect to be of interest in the study of SGD for matrix and tensor PCA. The first example is a limiting ballistic limit theorem for the evolution of the population loss . The population loss can be taken added to the family of summary statistics in our -localizable triple; in the case of -tensor PCA, this yields,
5. A finer diffusive limit theorem at a random start
Interestingly, with this double rescaling, the limit yields a pair of OU processes that are decoupled, namely, each of their drifts are autonomous and their stochastic parts independent. This pair of independent OU processes is depicted in Figures 4–5.
Two-layer networks for classifying a binary Gaussian mixture
As a warm-up to the XOR problem that we will consider in Section 5, we consider the problem of supervised classification of a binary Gaussian mixture model (binary GMM) which is defined as follows. Suppose that we are given i.i.d. samples of the form , where is a -valued random variable and, conditionally on , we have
It is classical that the Bayes optimal estimator in this setting is given by . Furthermore, this estimator can be achieved by (a rounding of) the output of a single layer neural network trained using the binary-cross-entropy loss (4.1). This is also called logistic regression. The single-layer setting can be easily analyzed via our framework. Our focus here, however, is to demonstrate our analysis on multi-layer neural networks.
To that end, we consider now the same setting, except that we will estimate the class labels using a simple two-layer neural network. (Note that the Bayes’ optimal estimator is still expressible by this architecture.) At first glance, this may seem an elementary setting with little to say. However as we will see, even in this simple setting surprising behaviour can occur in the high-dimensional setting which runs counter to common intuition. Furthermore, as we will see in Section 5, the phenomena occurring here also appear in richer problems such as the XOR problem.
2. Analysis
where is applied component wise and .
It can be shown (see Lemma 8.1) that the law of the loss at a given point, , depends only on the summary statistics,
where and with denoting the part of orthogonal to . For a point, , let
By similar reasoning to Lemma 8.1, it can be seen that these are functions only of , and we denote them as such, e.g., . See Section 8. The critical scaling for is then of order and we obtain the following.
Let be as in (4.2) and fix any and . Then converges to the solution of the ODE system, , initialized from , with:
and correctors , for .
3. Low variance asymptotics
Due to the Gaussian integrals defining , it is difficult to analyze the ODE system defined by Proposition 4.1, let alone any rescaled effective dynamics. For ease of analysis, we next send corresponding to a small noise regime for the Gaussian mixture. We emphasize that this limit is taken after and therefore is still approximately on the critical scale of at which there is a transition in the existence of any fixed point which is a good classifier. In particular, if is any diverging sequence, then the limiting effective dynamics would exactly match that attained by now sending . In Figure 6, we demonstrate numerically that the following predicted fixed points from the limit match those arising at finite large and .For large , this is indeed a quantitative approximation as exhibit locally Lipschitz dependence on , so the corresponding dynamics converges as by classical well-posedness results (see, e.g., )
The limit of the ODE system of Proposition 4.1 is given by
and . The fixed points of this system are classified as follows. All fixed points have and for . In , the coordinates are classified by
A fixed point at that is stable if ;
If , two unstable sets of fixed points at the quarter-circles given by having such that for .
If , two stable fixed points at equals and .
If is e.g., given by and then is in the coordinates, and is in the basin of attraction of the quarter-circles of item (2) with probability and the basin of attraction of the stable fixed points of (3) with probability .
4. Convergence to spurious solutions
Let us pause to interpret this result. The stable fixed points when are the optimal classifiers, whereas the unstable set of fixed points given by item (2) misclassify half of the data. Therefore, the above indicates that when solving the above task with randomly initialized weights, one of the following two scenarios occur, each with probability (with respect to the initialization): the algorithm will converge to the optimal classifier in linear time or it will appear to have converged to a macroscopically sub-optimal classifier on the same timescale, see Figures 6–7 for numerical verification of this at finite and .
5. Degeneracy of diffusive limits
Two-layer networks for the XOR Gaussian mixture
This data model is a Gaussian mixture model analogue of the (in)famous XOR problem of Minsky–Papert . In particular, it is easy to see that the optimal decision boundary is not expressible by a single-layer neural network as the data is not linearly separable. That said, it is also straightforward to see that this decision boundary is realizable by simple two-layer networks.In the notation of the following subsection, this can be realized by taking , , , and for for some .
We focus on this example as a demonstration of the applicability of our techniques to the analysis of the training dynamics for two-layer neural networks on natural data models. While this model is arguably the simplest model requiring a multi-layer network to solve, it nevertheless exhibits very complex phenomenology. We mention that some of these complexities were also observed in a very similar setup in where ballistic limits from warm starts were derived.
2. Analysis
Consider the corresponding classification problem using a two-layer neural network, taking as our estimator of the class label to be the natural rounding of , where and are the sigmoid and ReLU as in Section 4. We take to be a matrix and to be a -vector.
where again are applied component wise and again .
In Lemma 9.1 below, we show that the law of the loss at a point depends only on the following variables: for ,
By similar reasoning, it can be shown that these functions are expressible as functions of alone (see Section 9 below). We then find the following effective ballistic dynamics.
Let be as in (5.1) and fix any and . Then converges to the solution of the ODE system , initialized from with
and correctors , and for .
3. Low variance asymptotics
As with the binary GMM, one can develop the large limit of these asymptotics after . The effective dynamics in this regime are noticeably more tractable. We defer the precise expressions of these dynamics to Proposition 9.1 below. Let us instead classify the corresponding fixed points.
The fixed points of the ODE system of Proposition 9.1 are classified as follows. If , then the only fixed point is at .
If , then let be any disjoint (possibly empty) subsets whose union is . Corresponding to that tuple , is a set of fixed points that have for all , and have
for ,
such that and for all ,
such that and for all ,
such that and for all ,
such that and for all .
In the case, these form connected sets of fixed points, and of which are fixed points that are stable, corresponding to the possible permutations in which each of are singletons.
Similar to the binary GMM, in Figures 9–10, we demonstrate numerically that the following predicted fixed points from the limit match those arising at finite large and .
In the case, we can also exactly calculate the probability that the effective dynamics in the ballistic phase converges to a stable fixed point (as opposed to an unstable one). From a Gaussian initialization where and independently, this converges to . We refer the reader to Section 9.4 for the proof.
4. Overparametrization in the XOR GMM
Since the the derivations of the ballistic limiting equations apply for general , we can also study the probability of ballistic convergence to a stable vs. unstable fixed point as one varies . This addresses the regime of overparametrization for the XOR GMM since suffices to express a Bayes–optimal classifier. In this more generic setting, the probability of being in the ballistic domain of attraction of the stable fixed points (corresponding to the Bayes optimal classifiers) is
which goes to exponentially fast as grows. This clearly demonstrates the benefits of overparametrizaiton of the landscape in a concrete two-layer network: a random initialization is more likely to to contain the “right" initial signature (corresponding to none of being empty at initialization) in order to be in the basin of a Bayes optimal classifier as the width grows, and as long as the right signature is present in the nodes at initialization, the SGD will ballistically converge to a global minimizer of the population loss. This is a rigorous example of the well-known lottery ticket hypothesis of . Roughly speaking, the lottery ticket hypothesis proposes that the reason for the success of overparametrized networks is that they give more attempts for a sufficiently expressive subnetwork to be initialized well, and succeed at the task on its own.
5. Diffusive limits at unstable fixed points
As an example of the diffusions that can arise in the rescaled effective dynamics at the unstable fixed points, let us consider the unstable fixed points in which has the correct signature (two positive, two negative) but for each of those we are at a corresponding quarter-ring. By way of example, we can set , or equivalently focus on a fixed point where all indices beyond the first four have . Here, the dynamics effectively becomes a pair of 2 two-layer GMM’s on quarter-rings (as in Section 4), that are anti-correlated. More precisely, let be such that and such that , for . Take as fixed points about which we expand to be and for . Namely, we let
Numerical simulations in Figure 11 confirm these degenerate diffusive limits at finite .
Part III Proofs
In this section, we prove our main convergence result, namely Theorem 2.3. The drift terms can be seen from a Taylor expansion out to second order, with the role played by -localizability being to justify neglecting certain negligible second order terms, as well as all higher order terms. The identification of the stochastic term is via the classical martingale problem for summary statistics of stochastic gradient descent in the high-dimensional limit.
Our aim is to establish weakly as random variables on where solves (2.4). It is equivalent to show the same on equipped with the sup-norm for every .
Recalling Definition 2.1, since are -localizable, the error term in (6) has
and similarly , then recalling that is the linear interpolation of , we may write
where and .
We now prove that the sequence is tight in with limit points which are -Holder for each . To this end, let us define
As the error above is uniform in , we have that
Thus it suffices to show the claimed tightness and Holder properties of limit points for instead of . We aim to show that for all ,
from which we will get that the sequence is uniformly -Hölder by Kolmogorov’s continuity theorem. Evidently, for all we have that
We control these terms in turn. We will do this coordinate wise and, for readability, fix some and let , , etc.
Let be as in (2.2). Then the first term in (6) satisfies
by continuity of . For the second term in (6),
which is by items (1)–(2) -localizability. (Applying this bound for , the last term in is vanishing in the limit for each whenever .) Combining the above bounds yields
For the martingale term, notice that by Burkholder’s inequality,
For the first term in that martingale difference, observe that
where in the second line we used Cauchy-Schwarz and in the last we used item (3) of -localizability.
For the second term in the martingale difference,
by items (1)–(2) of -localizability. Finally, by the same reasoning, for the third term,
All of the above terms are since . Thus we have the claimed (6.2), and by Kolmogorov’s continuity theorem, , are uniformly -Holder and thus the sequence is tight with -Holder limit points. Notice furthermore that if we look at , this sequence is also tight and the limits points are continuous martingales. Let us examine their limiting quadratic variations.
Let and define and analogously. Furthermore, let , and be their respective limits which we have shown to exist and be -Holder.
is a martingale. We therefore need to consider the limit as of the integral above. Write
Consider the integrals of times each of these four terms separately. For the first term,
goes to zero as by the assumption in (2.3).
We now reason that the integrals of times the other three terms in (6.7) all go to zero as . The second and third are identical: by Cauchy–Schwarz,
The first expectation contributes by the first part of item (3) of localizability. Also,
The first of these terms is at most as argued in (6). The second is by the second part of item (3) in the definition of localizability. As such, we are able to conclude that
The integral of times the fourth term in (6.7) is handled similarly using Cauchy–Schwarz and the bound of on (6.8).
Thus, if we consider the continuous martingales given by , its angle bracket is, by definition, given by
Proofs for matrix and tensor PCA
In this section, we prove the results of Section 3. We will state them in the more general setting where we add a ridge penalty to the loss, so that for fixed, the loss is given by
where only depends on . Note that .
Our first aim is to establish Proposition 3.1, showing that the summary statistics satisfy the conditions of Theorem 2.3 with the desired and . We begin by checking localizability for . In what follows, for ease of notation we will denote and . In these coordinates,
We check the items in Definition 2.1 one by one, beginning with item (1). Express the derivatives for as
For item (2), differentiating (7.2), , where
Notice that and . Consider
the bounding quantity is evidently a continuous function of and therefore as long as is such that , it is bounded by some . Next, if we consider
where the bound on the operator norm of an i.i.d. Gaussian -tensor can be found, e.g., in [9, Lemma 4.7]. Moving on to item (3), by the same reasoning, for every ,
If then and if then , so in both cases this is at most . Finally, is only non-zero if in which case it is . Then,
by the second item in the definition of localizability, and evidently the right-hand side is if . ∎
Having checked localizability for , we apply Theorem 2.3. To compute , by the above,
In particular, for , we have and
from which we obtain in the limit that that and .
Together, these yield the ODE system of (3.1),
which in the case matches Proposition 3.1. Finally, to see that , consider
which when multiplied by evidently vanishes. ∎
We now turn to analyzing the ODE of Proposition 3.1.
At the fixed points of the ODE in Proposition 3.1,
If , then and there are two possible fixed points: either or solves
Notice that if , this has a nontrivial solution of the form , provided , and if , this has a nontrivial solution provided at i.e., . This gives
Evidently when we take , then its non-trivial solution is at for all .
Alternatively, if at a fixed point, then we can simplify further and get
For simplicity of calculations, set as is the case in Proposition 3.1. Then, we simply get . In the case of , we also find that there is a solution if and only if , in which case , from which together with , we also get .
In the general case of , we find that . This has real solutions (all of which have as required) whenever defined as
(Interpreting , this returns .) With this , whenever , the equation for has exactly two real solutions, both of which are at least which we can denote by
2. Effective dynamics for the population loss
In practice, one is interested in tracking the loss, or ideally, the generalization error. In this subsection, we add the generalization error to our set of summary statistics and obtain limiting equations for its evolution from (3.4).
Recalling (7.2), the fact that is a localizable summary statistic follows from the facts that , and the fact that is a smooth -independent function of .
For simplicity of calculations let us stick to .
Next, consider the corrector for . For this, notice that
Recalling from (7.4), and taking , all the terms in vanish in the limit except the contribution from the , which yields Finally, we wish to compute the volatility for the stochastic part of the evolution of . For this, consider and notice that all the entries of that matrix are continuous functions of and thus go to zero when multiplied by .
3. Diffusive limits at the equator
Taking limits as , as long as is fixed in , we see that is given by
4. Diffusive limit for the radius
where we used in the first inequality that the law of is rotation invariant and is a -homogenous function. For the second part of item (3),
The are Gaussian with mean zero, and by (7.4), variance and covariance . Recall the following fact about Gaussians: if are Gaussians with variances and covariance , then for some universal constant . Also, . Applying this to , we get
Combined with the above, this gives a bound of on the second part of item (3).
We now calculate the resulting drifts. For , write
Combining terms and sending , we obtain
Multiplying by and taking the limit as , the two entries of this matrix that survive are and , where and . All in all, we obtain (3.5).
Proofs for the binary Gaussian mixture model
Recall the cross-entropy loss for the binary GMM with SGD from (4.1), and recall the set of summary statistics from (4.2).
Let and . Then, notice that
Next, notice that as a vector, is distributed as , where are i.i.d. , and are jointly Gaussian with means zero and covariance
Similarly, the distribution of also only depends on . Finally,
Therefore, at any point , the law of , and thus , is simply a function of . To see that the summary statistics satisfy the bounds of item (1) in Definition 2.1, write . Then
For the higher derivatives, evidently we only have second derivatives in the last 3 variables each of which is given by a block diagonal matrix where only one block is non-zero and is given by an identity matrix. The third derivatives of all elements of are zero. ∎
We can now express the loss, the population loss, and their respective derivatives and they (their laws at a fixed point) will evidently only depend on the summary statistics. One arrives at the following expressions for by direct calculation from (4.1).
(Notice that if , then is only a function of by the same reasoning as used in Lemma 8.1.) Then, we can also easily express
Finally, the matrix can be expressed as follows:
Let us conclude this subsection with the following simple preliminary bounds that will be useful towards establishing the conditions of -localizability from Definition 2.1, and the promised limiting equations. The proofs of these are straightforward using Gaussianity and are provided in Section 10 for completeness.
For each , for every and every , we have
For every and for , we have
Throughout this section we will take . By rotational invariance of the problem, this is without loss of generality, and only simplifies certain expressions. The -localizability can be seen by application of the moment bounds listed above.
The condition on was satisfied per Lemma 8.1. Recalling from (8.12), one can verify that the norm of each of the four terms in is individually bounded, using the Cauchy–Schwarz inequality together with the bound of Lemma 8.2 on .
When is , this is simply a fourth moment bound on , which follows from the ’th moment by Jensen’s inequality. When is , or , the bound follows from
for choices of being either in which case or in which case . For each , this is at most some constant using the two bounds of Lemma 8.2.
The convergence of the population drift to from Proposition 4.1 follows by taking the inner products of from (8.12) with the rows of from (8.8), and noticing that from (4.3) is exactly and from (4.3) is exactly .
Next consider the convergence of the correctors to the claimed . The variables are linear so and for these, . For for , the relevant entries in are those corresponding to and . For ease of notation, in what follows let .
For ease of calculation taking , we have , which by (8), and the choice of , is given by
which we emphasize is only a function of . We lastly need to show that the diffusion matrix goes to zero as when . This is straightforward to see by considering any element of and using Cauchy–Schwarz together with the two bounds of Lemma 8.2 to bound it in absolute value by some independent of . Then when multiplying by any , this entire matrix will evidently vanish. ∎
2. The small-noise limit of the effective dynamics
One can now take a limit to arrive at the ODE system of Proposition 4.2.
We begin with considering : its limiting value will depend on the signs of both and . We can express from (4.3) as
We claim that the two terms on the right-hand side converge to and respectively. This follows by e.g., writing the difference as
at which point, we see that if , this becomes , as it is if . If and , then you get and and likewise if and .
Next consider the limit as of from (4.3), which we claim converges to . Write
Finally, since , the quantity evidently goes to zero as . ∎
The above argument used for the limit of . If one considers the cases when , the limiting drifts still apply. For this, it suffices to show that if , then converges to zero. Without loss of generality, suppose and consider
This is zero independently of by independence of from the other Gaussians in the expectation.
Evidently, every fixed point must have . Furthermore, if we let , then
and therefore every fixed point of the ODE system must have , which is to say . Therefore, it suffices to characterize the fixed points in terms of as claimed. This reduces to if and otherwise. Observe first that the point is a fixed point of this system. If , then dividing out by , the above reduces to if and otherwise. Recalling that we obtain the claimed set of fixed points by inverting these equations (they only have a solution if ).
In order to study the stability of the various fixed points, notice first that the ODE system of Proposition 4.2 is a gradient system for the population loss,
Since it is a gradient system, with only the specified fixed points, the stability of a fixed point can be deduced by showing it is the minimizer of . In particular, the values of at its critical points are given by at , when , and when . It is a simple calculus exercise to show that the smallest of these is when and when .
To show that each of the other critical points are all unstable, one can find a direction along which the dynamical system is locally repelled from it. For instance, we will show that the ring of fixed points with and with is unstable, by showing a repelling direction arbitrarily close to the point , . If and , then there reduces to , and as long as , there exists such that so for all small enough.
3. Rescaled effective dynamics around unstable fixed points
Now Taylor expanding the sigmoid function, and using the definition of , we get
Proofs for the XOR Gaussian mixture model
We could also have added a bias at each layer, however the Bayes classifier in this problem is an “X” centered at the origin so we can safely take the biases to be 0.
Recall the set of summary statistics from (5.1). The next lemma shows that form a good set of summary statistics.
where are i.i.d. and are jointly Gaussian with covariance matrix
Similarly, the law of depends only on . Finally,
Therefore, at a fixed point the law of is only a function of .
where is if and otherwise. For higher derivatives, we only have second derivatives in the variables, each of which is given by a block diagonal matrix where only one block is non-zero and it is twice an identity matrix. Thus the operator norm of these second derivatives is . The third derivatives of all elements of are zero. ∎
By the same reasoning as in Lemma 9.1, if , then is only a function of . We then also have the conclusions of Lemma 8.2 for distributed according to the XOR GMM by simply decomposing it into two mixtures, and we will therefore appeal to this lemma meaning its analogue for the XOR GMM.
The condition on was satisfied per Lemma 9.1. Recalling from (8.12), one can verify that the norm of each of the four terms in is individually bounded, using the Cauchy–Schwarz inequality together with the bound of Lemma 8.2 on , naturally adapted to XOR. The remaining estimates are also analogous to the proof of Lemma 8.4 with the analogue of Lemma 8.2 applied. ∎
2. Effective dynamics for the XOR GMM
The convergence of the population drift to from Proposition 4.1 follows by taking the inner products of from (8.12) with the rows of from (9.3), and noticing that is exactly , is exactly , and is exactly .
We next consider the population correctors. The fact that follows from the fact that the Hessians of are zero. For the corrector for , the relevant entries of are those corresponding to and . For ease of notation, in what follows let .
By the same arguments on the concentration of the norm of Gaussian vectors as used in the binary GMM case, then we deduce from this that
Finally, let us establish that the limiting diffusion matrix is all-zero whenever . This follows exactly as it did in the proof of Proposition 4.1. ∎
3. Small noise limit of the effective dynamics
The aim of this section is to establish the following small-noise limit of the effective dynamics ODE of Proposition 5.1. This will again be quite similar to the analogous proofs for the binary GMM in Section 8, and when these similarities are clear we will omit details.
In the limit, the ODE from Proposition 5.1 converges to
and for .
Let us begin with convergence of . We claim that it converges to
The point will be that when taking the inner product with , the first two terms here contribute to the limit and the latter two vanish, while when taking the inner product with , the first two terms vanish in the limit while the latter two contribute.
Consider e.g., the first of the four terms above, and inner product with . In this case, consider
which is precisely the quantity that was exactly shown to go to zero as in (8.20). To see that the third and fourth terms above go to zero when taking their inner product with , observe that they become
which by orthogonality of and is at most by the reasoning of Lemma 8.2, therefore vanishing as . Together with its analogue for , this implies the claim for the convergence of , as well as its analogous limit of .
We next consider the limit as of , which we claim goes to . Using the expansion of from earlier in this proof, we can consider as four terms having the form of the terms in (8.21), which were there showed to go to zero as . Since here is orthogonal both to and , the same proof applies.
Finally, in order to see that the limit as of is zero, which follows from the fact that . ∎
The fixed points of the ODE system of Proposition 9.1 are classified as follows. If , then the only fixed point is at .
If , then let be any disjoint (possibly empty) subsets whose union is . Corresponding to that tuple , is a set of fixed points that have for all , and have
for ,
such that and for all ,
such that and for all ,
such that and for all ,
such that and for all .
In the case, these form connected sets of fixed points, and of which are fixed points that are stable, corresponding to the possible permutations in which each of are singletons.
Evidently, any fixed point must have for all . Furthermore, the point for evidently forms a fixed point of the system. Now suppose there is some fixed point with for some ; in that case, it must be that and . Therefore, we can select a subset of such that for .
For any such choice of , consider next, . We first claim that if at a fixed point, then and , whereas if then and . To see this, notice that at any fixed point,
Since is non-negative, if , the sign of the right-hand side of the first equation is the same as the sign of so it can have a non-zero solution, while the sign of the right-hand side of the second equation is the opposite of the sign of , so any such fixed point must have . To see that at such a fixed point, now set and take the fixed point equations for and , dividing one by and the other by to see that
as claimed. The fixed points having are solved symmetrically.
Our classification now reduces to understanding the possible values taken by given their signs (when non-zero). Fix a partition of and consider the set of fixed points having for , on and so on as designated by Proposition 9.2; by the above any fixed point is of this form. It remains to check that the values of on each of these sets are as described by the proposition.
In order to see this, fix e.g., . Then, and , and so the fixed point equations reduce to
since the only coordinates where will be non-zero are , where . Inverting the sigmoid function, this implies exactly the claimed . The cases of are analogous, concluding the proof.
The count of the number of connected components of fixed points this forms is sensitive to , so for concreteness let us perform it when . We first notice that the fixed point at is disconnected from all others. Fixed points corresponding to some are part of the same connected component of fixed points if one goes from one to the other by moving an element of (for some and to without making empty, or by moving an element of to a non-empty .
We turn now to studying the stability of these various sets of fixed points. Observe that in the limit, the dynamical system of Proposition 9.1 is a gradient system for the population loss
At a fixed point (which necessarily has , , and is characterized by the partition of into , this reduces to
At this point, noticing that is equal to if is non-empty and if it is empty, and similarly for , this turns into a simple optimization problem over the number of non-empty . Just as in the binary GMM case, it becomes evident that when , this is minimized at for all (i.e., they are all empty and , whereas when the above is minimized when every one of are all non-empty. This yields the global minima of in these coordinates, and ensures the fixed points we claimed were stable are indeed stable.
To show the instability of any other connected set of fixed points, the reasoning goes just as in the binary GMM case: consider a small perturbation of the specified critical region in the direction of the stable fixed points and it can be seen by examining the drifts directly, that the dynamical system has a repelling direction. ∎
When , the counting of connected components of fixed points of course changes. However, what is still clear by an identical calculation is that the sets of fixed points minimizing will still be when and will be all fixed points that have all four of being non-empty if . Notice that when and , even the set of stable fixed points become connected to form a single stable manifold.
4. 3/32332\nicefrac{{3}}{{32}}-probability of ballistic convergence to an optimal classifier
We now reason that when the ballistic effective dynamics of Proposition 9.1 is such that under an uninformative Gaussian initialization, the probability of being in a basin of attraction of one of the 24 stable fixed points is . Begin by noticing that if the first layer weights are initialized as independently for and the second layer weights are independent standard Gaussians, then the projection onto the coordinate system is given by
The -functions at zero for however cause some trouble because of the indicator functions on the sign of and in the equations of Proposition 9.1.
Under the flow of Proposition 9.1, if is positive, then stays fixed at zero, and if then becomes negative infinitesimally quickly, whereas if then it becomes positive infinitesimally quickly. At any rate, the sign of never changes to negative from such an initialization, and similarly if is negative, the sign of will never change to positive. As such, in order to have a chance at being in the basin of attraction of one of the stable fixed points outlined in Proposition 9.2, it must be the case that two of have positive sign and two of them have negative sign; evidently this has probability .
Given that two of are positive, and two of them are negative—say without loss of generality that are the coordinates in which it is positive, and are the coordinates in which it is negative—then the dynamical system for is exactly the ballistic limit of the two-layer GMM studied in Section 4, for which we found that the probability of converging to a good classifier is . Similarly, the dynamical system for independently gives a further probability of converging to its good classifier. Together, these yield a probability of of converging to one of the many optimal classifiers for the XOR GMM.
Generically, if , by a similar reasoning to the above, in order to fall in the basin of attraction of the stable fixed points, it must be the case that the initialization has some four indices each of which initially belong to . This is the probability that are positive for at least two indices, and negative for at least two indices, and then among the indices at which is positive, there is at least one index where is positive and one where it is negative, and similarly with negative and . Doing this combinatorial calculation out, we find that the probability of being in a good initialization is exactly the expression in (5.4). This is easily seen to go to exponentially fast as since the initial choice of ’s will typically have around positive and negative coordinates, and with exponentially high probability those will have both positive and negative and .
5. Diffusive limit on critical submanifolds
Plugging these in, and taking the limit we find that for ,
By a similar reasoning, for , we have
and if and , then
By a similar reasoning, if , then
and if and , then
Proofs of technical lemmas for Gaussian mixtures
In this section, we establish the technical bounds on Gaussian moments in Lemmas 8.2–8.3.
For the first bound, let and consider
The quantities in the expectations are at most some universal constant times . To bound the expectation of the second term here, notice that is distributed as implying the desired.
The bound on goes as follows. Evidently it suffices to let for , and prove the bound on the norm of
Now decompose as , where is independent of which is distributed as with given by (8.3), which is independent of distributed as a standard Gaussian vector orthogonal to the subspace spanned by . By independence of from the indicator and the argument of the sigmoid, all those terms contribute nothing to the expectation, and therefore,
Here, we used the first inequality of the lemma. This yields the desired. ∎
The proof of (8.16) is easily seen by rewriting the probability in question as
so that as long as this goes to zero as .
Applying Cauchy–Schwarz to the first term, it suffices to establish the following bounds
To demonstrate the first of these inequalities, notice that
uniformly over , per Fact 8.1. For the second desired bound, expand as
It suffices to show the expectation of the square of each of these goes to zero as . First,
If , the expectation on the right goes to zero by (8.16). Second,
When , this is evidently zero; when , if , this is
which goes to zero as when , by the explicit formula for the moment generating function of the Gaussian , whose variance is . ∎
The authors thank the anonymous referees for their useful comments and suggestions. The authors thank F. Krzakala, L. Zdeborova, and B. Loureiro for interesting conversations and suggestions, especially suggesting we investigate the role of overparametrization in the XOR GMM. The authors thank M. Sellke for pointing out the relationship to the lottery ticket hypothesis. The authors also thank M. Glasgow for a careful reading and helpful suggestions. R.G. acknowledges the support of NSF DMS-2246780 and the Miller Institute for Basic Research in Science. A.J. acknowledges the support of the Natural Sciences and Engineering Research Council of Canada (NSERC) and the Canada Research Chairs programme. Cette recherche a été enterprise grâce, en partie, au soutien financier du Conseil de recherches en sciences naturelles et en génie du Canada (CRSNG), [RGPIN-2020-04597, DGECR-2020-00199], et du Programme des chaires de recherche du Canada.