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 \nicefrac2932\nicefrac{{29}}{{32}} 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 11; 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 1/21/2 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 kk-Gaussians for fixed k≥2k\geq 2. 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

max⁡isup⁡x∈un−1(EK)∣∣∇2uin∣∣op⁡≤CK⋅δn−1/2\max_{i}\sup_{x\in\mathbf{u}_{n}^{-1}(E_{K})}\lvert\lvert\nabla^{2}u_{i}^{n}\rvert\rvert_{\operatorname{op}}\leq C_{K}\cdot\delta_{n}^{-1/2}, and max⁡isup⁡x∈un−1(EK)∣∣∇3uin∣∣op⁡≤CK\max_{i}\sup_{x\in\mathbf{u}_{n}^{-1}(E_{K})}\lvert\lvert\nabla^{3}u_{i}^{n}\rvert\rvert_{\operatorname{op}}\leq C_{K};

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, An=⟨∇Φ,∇⟩\mathcal{A}_{n}=\langle\nabla\Phi,\nabla\rangle and Ln=12⟨V,∇2⟩\mathcal{L}_{n}=\frac{1}{2}\langle V,\nabla^{2}\rangle.

In this case we call h\mathbf{h} the effective drift, and Σ\mathbf{\Sigma} the effective volatility.

We are now ready to present our main result. For a function ff and measure μ\mu we let f∗μf_{*}\mu denote the push-forward of μ\mu.

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 un\mathbf{u}_{n}. The fact that h,Σ\mathbf{h},\mathbf{\Sigma} 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 Φ,L\Phi,L 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 ⋃Kun−1(EK)\bigcup_{K}\mathbf{u}_{n}^{-1}(E_{K}) does not cover Xn\mathcal{X}_{n}, nor does it imply uniform smoothness of HH (and in turn LL) as we may (and will) be taking δn→0\delta_{n}\to 0 with nn.

Let us lastly motivate the scalings appearing in item (3), which ensure there is some independence between HH and the values of ∇u\nabla u and ∇2u\nabla^{2}u at xx. As a testbed, suppose that ∇H(x)\nabla H(x) is a random vector with i.i.d. entries all of order 11. If uu is a rescaled linear statistic, e.g., δn−1/2⟨x,e1⟩\delta_{n}^{-1/2}\langle x,e_{1}\rangle then the first bound of item (3) is saturated, and the second of course is trivial due to the linearity of uu. The second bound is saturated by taking a rescaling of a radial statistic, e.g., δn−1/2∥x∥2\delta_{n}^{-1/2}\|x\|^{2}, again assuming for maximal simplicity that ∇H\nabla H 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 n→∞n\to\infty 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 O(δ−3)O(\delta^{-3}) by the other localizability assumptions. Then Theorem 2.3 would still apply, but the limiting diffusion matrix would be the n→∞n\to\infty limit (assuming it exists) of

as opposed to simply the limit of δJVJT\delta JVJ^{T}.

in which case, evidently (2.2) holds with h=−f+g\mathbf{h}=-\mathbf{f}+\mathbf{g}. When (2.5) and (2.6) both hold, we call f,g\mathbf{f},\mathbf{g} and Σ\Sigma the population drift, the population corrector, and the diffusion matrix of u\mathbf{u} respectively. From the fixed dimensional perspective, when (2.5) holds, one predicts u\mathbf{u} to solve

with initial data u0∼u∗μ\mathbf{u}_{0}\sim\mathbf{u}_{*}\mu. as this is its evolution under gradient descent on the population loss Φ\Phi. Evidently this perspective only applies in the high-dimensional limit of Theorem 2.3 if both the population corrector g\mathbf{g} and the diffusion matrix Σ\Sigma are zero. We find that for any triple (un,Ln,Pn)(\mathbf{u}_{n},L_{n},P_{n}), there is a scaling of the learning rate δn\delta_{n} with nn below which g=Σ=0\mathbf{g}=\Sigma=0, 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 gg and Σ\Sigma may be non-zero, and the high-dimensionality induces non-trivial corrections to f\mathbf{f}. (In the case of teacher–student networks, the terms f\mathbf{f} and g\mathbf{g} can be compared to the “learning" and “variance" terms in Eq. (14a) of .)

To see this, notice that if the triple (un,Ln,Pn)(\mathbf{u}_{n},L_{n},P_{n}) is δn\delta_{n}-localizable for some δn→0\delta_{n}\to 0, then it is also δn′\delta_{n}^{\prime}-localizable for every sequence δn′=O(δn)\delta^{\prime}_{n}=O(\delta_{n}). If furthermore (2.3) and (2.5)–(2.6) hold for δn\delta_{n} with some f,g\mathbf{f},\mathbf{g} and Σ\Sigma, then these limits also exists for δn′=o(δn)\delta_{n}^{\prime}=o(\delta_{n}) with the same f\mathbf{f} but with g=Σ=0\mathbf{g}=\Sigma=0. As such, there can be exactly one scaling of δn\delta_{n} with nn at which g\mathbf{g} or Σ\Sigma may be non-zero, and for all smaller scales of δn\delta_{n}, the fixed-dimensional perspective of (2.7) applies.Note that if δn=o(δn′)\delta_{n}=o(\delta_{n}^{\prime}), then limiting g,Σ\mathbf{g},\Sigma may not exist for δn′\delta^{\prime}_{n}, 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 h≠f\mathbf{h}\neq\mathbf{f}. 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 1/δn1/\delta_{n}. If one is then interested in the evolution of un\mathbf{u}_{n} in microscopic o(1)o(1) 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 un\mathbf{u}_{n} near u⋆\mathbf{u}_{\star}:

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 k=2k=2, 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 λ=1\lambda=1 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 λ=1\lambda=1.

The case k≥3k\geq 3 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 Φ(x)=−2λmk+(r⊥2+m2)k+c\Phi(x)=-2\lambda m^{k}+(r_{\perp}^{2}+m^{2})^{k}+c, and the law of LL only depends on them.

In our normalization with λ>0\lambda>0 fixed, the regime δn=o(1/n)\delta_{n}=o(1/n) is sub-critical and the regime δn=Θ(1/n)\delta_{n}=\Theta(1/n) is critical. Note that with different scalings of λn\lambda_{n}, the critical learning rate changes. We focus our presentation on the most interesting regime, namely the critical scaling regime of δn=cδ/n\delta_{n}=c_{\delta}/n for some constant cδc_{\delta}. 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 cδ↓0c_{\delta}\downarrow 0 limits of the ODE’s/SDE’s of the critical regime.

For notational simplicity, let R2:=m2+r⊥2R^{2}:=m^{2}+r_{\perp}^{2}. We consider the pair un=(u1,u2)=(m,r⊥2)\mathbf{u}_{n}=(u_{1},u_{2})=(m,r_{\perp}^{2}), for which Theorem 2.3 yields the following effective dynamics.

Fix k≥2k\geq 2, λ>0\lambda>0, cδ>0c_{\delta}>0 and let δn=\nicefraccδn\delta_{n}=\nicefrac{{c_{\delta}}}{{n}}. Then un(t)\mathbf{u}_{n}(t) converges as n→∞n\to\infty to the solution of the following ODE initialized from lim⁡n→∞(un)∗μn\lim_{n\to\infty}(\mathbf{u}_{n})_{*}\mu_{n}:

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 cδ=1c_{\delta}=1 where one sees from (3.1) that r⊥2→1r_{\perp}^{2}\to 1, where the problem in the matrix case is most directly related to an eigenvalue problem (see Section 7 for the generic cδc_{\delta} 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 UU for an ODE stable if for every ϵ>0\epsilon>0, for every u∈Bϵ(U)u\in B_{\epsilon}(U), the solution of the ODE with initialization uu converges to some point in UU as t→∞t\to\infty. Otherwise, we call UU unstable.

Eq. (3.1) has isolated fixed points classified as follows. Let λc(k)\lambda_{c}(k) be as in (7.6) and m†(k,λ)≤m⋆(k,λ)m_{\dagger}(k,\lambda)\leq m_{\star}(k,\lambda) be as in (7.7) (if k=2k=2, λc=1\lambda_{c}=1 and m†=m⋆=λ−1m_{\dagger}=m_{\star}=\sqrt{\lambda-1}):

An unstable fixed point at (0,0)(0,0) and a fixed point at (0,1)(0,1); if k=2k=2, (0,1)(0,1) is stable if λ<λc(2)\lambda<\lambda_{c}(2) and unstable if λ>λc(2)\lambda>\lambda_{c}(2); if k>2k>2 (0,1)(0,1) is always stable.

If λ>λc(k)\lambda>\lambda_{c}(k): when k=2k=2, two stable fixed points at (±m⋆(2),1)(\pm m_{\star}(2),1). When k≥3k\geq 3, two unstable fixed points at (±m†(k),1)(\pm m_{\dagger}(k),1) and two stable fixed points at (±m⋆(k),1)(\pm m_{\star}(k),1).

The presence of two pairs of fixed points when k≥3k\geq 3 with non-zero correlation with vv may seem surprising—indeed it indicates that even some warm starts will fail to attain good correlation with the signal when λ\lambda is finite. This is an interesting consequence of the corrector in (3.1) and if one tracks the cδc_{\delta} dependence in the above, the fixed point m†m_{\dagger} goes to zero as cδ→0c_{\delta}\to 0 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 Λc(k)=1\Lambda_{c}(k)=1, as in k=2k=2.

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 Φ(x)\Phi(x). The population loss can be taken added to the family of summary statistics in our δn\delta_{n}-localizable triple; in the case of kk-tensor PCA, this yields,

5. A finer diffusive limit theorem at a random start

Interestingly, with this double rescaling, the n→∞n\to\infty 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 Y=(y,X)Y=(y,X), where yy is a {0,1}\{0,1\}-valued Ber(1/2)Ber(1/2) random variable and, conditionally on yy, we have

It is classical that the Bayes optimal estimator in this setting is given by y^=sgn⁡(μ⋅x)\hat{y}=\operatorname{sgn}(\mu\cdot x). 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 gg is applied component wise and p(v,W):=(α/2)(∣∣v∣∣2+∣∣W∣∣2)p(v,W):=(\alpha/2)(\lvert\lvert v\rvert\rvert^{2}+\lvert\lvert W\rvert\rvert^{2}).

It can be shown (see Lemma 8.1) that the law of the loss at a given point, (v,W)∈Xn(v,W)\in\mathcal{X}_{n}, depends only on the 77 summary statistics,

where mi=Wi⋅μm_{i}=W_{i}\cdot\mu and Rij⊥=Wi⊥⋅Wj⊥R_{ij}^{\perp}=W_{i}^{\perp}\cdot W_{j}^{\perp} with Wi⊥=Wi−miμW_{i}^{\perp}=W_{i}-m_{i}\mu denoting the part of WiW_{i} orthogonal to μ\mu. For a point, (v,W)∈Xn(v,W)\in\mathcal{X}_{n}, let

By similar reasoning to Lemma 8.1, it can be seen that these are functions only of un\mathbf{u}_{n}, and we denote them as such, e.g., Aiμ=Aiμ(un)\mathbf{A}_{i}^{\mu}=\mathbf{A}_{i}^{\mu}(\mathbf{u}_{n}). See Section 8. The critical scaling for δ\delta is then of order Θ(1/n)\Theta(1/n) and we obtain the following.

Let un\mathbf{u}_{n} be as in (4.2) and fix any λ>0\lambda>0 and δn=\nicefraccδN\delta_{n}=\nicefrac{{c_{\delta}}}{{N}}. Then un(t)\mathbf{u}_{n}(t) converges to the solution of the ODE system, u˙t=−f(ut)+g(ut)\dot{\mathbf{u}}_{t}=-\mathbf{f}(\mathbf{u}_{t})+\mathbf{g}(\mathbf{u}_{t}), initialized from lim⁡n→∞(un)∗μn\lim_{n\to\infty}(\mathbf{u}_{n})_{*}\mu_{n}, with:

and correctors gvi=gmi=0{g}_{v_{i}}=g_{m_{i}}=0, gRij⊥=cδvivjλBij{g}_{R_{ij}^{\perp}}=c_{\delta}\frac{v_{i}v_{j}}{\lambda}\mathbf{B}_{ij} for i,j=1,2i,j=1,2.

3. Low variance asymptotics

Due to the Gaussian integrals defining f,g\mathbf{f},\mathbf{g}, 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 λ→∞\lambda\to\infty corresponding to a small noise regime for the Gaussian mixture. We emphasize that this limit is taken after n→∞n\to\infty and therefore is still approximately on the critical scale of λ=Θ(1)\lambda=\Theta(1) at which there is a transition in the existence of any fixed point which is a good classifier. In particular, if λ=λn\lambda=\lambda_{n} is any diverging sequence, then the limiting effective dynamics would exactly match that attained by now sending λ→∞\lambda\to\infty. In Figure 6, we demonstrate numerically that the following predicted fixed points from the λ→∞\lambda\to\infty limit match those arising at finite large NN and λ>0\lambda>0.For large λ\lambda, this is indeed a quantitative approximation as f,g\mathbf{f},\mathbf{g} exhibit locally Lipschitz dependence on λ−1\lambda^{-1}, so the corresponding dynamics converges as λ→∞\lambda\to\infty by classical well-posedness results (see, e.g., )

The λ→∞\lambda\to\infty limit of the ODE system of Proposition 4.1 is given by

and R˙ij⊥=−2αRij⊥\dot{R}_{ij}^{\perp}=-2\alpha R_{ij}^{\perp}. The fixed points of this system are classified as follows. All fixed points have Rij⊥=0R_{ij}^{\perp}=0 and mi=vim_{i}=v_{i} for i,j={1,2}i,j=\{1,2\}. In (v1,v2)(v_{1},v_{2}), the coordinates are classified by

A fixed point at (v1,v2)=(0,0)(v_{1},v_{2})=(0,0) that is stable if α>\nicefrac14\alpha>\nicefrac{{1}}{{4}};

If α<\nicefrac14\alpha<\nicefrac{{1}}{{4}}, two unstable sets of fixed points at the quarter-circles given by (v1,v2)(v_{1},v_{2}) having v1v2>0v_{1}v_{2}>0 such that v12+v22=Cαv_{1}^{2}+v_{2}^{2}=C_{\alpha} for Cα:=log⁡(1−2α)−log⁡(2α)C_{\alpha}:=\log(1-2\alpha)-\log(2\alpha).

If α<\nicefrac14\alpha<\nicefrac{{1}}{{4}}, two stable fixed points at (v1,v2)(v_{1},v_{2}) equals (Cα,−Cα)(\sqrt{C_{\alpha}},-\sqrt{C_{\alpha}}) and (−Cα,Cα)(-\sqrt{C_{\alpha}},\sqrt{C_{\alpha}}).

If μn\mu_{n} is e.g., given by (v1,v2)∼N(0,I2)(v_{1},v_{2})\sim\mathcal{N}(0,I_{2}) and W1,W2∼N(0,IN/(λN))W_{1},W_{2}\sim\mathcal{N}(0,I_{N}/(\lambda N)) then ν:=lim⁡(un)∗μn\nu:=\lim(\mathbf{u}_{n})_{*}\mu_{n} is N(0,I2)\mathcal{N}(0,I_{2}) in the v1,v2v_{1},v_{2} coordinates, and is in the basin of attraction of the quarter-circles of item (2) with probability \nicefrac12\nicefrac{{1}}{{2}} and the basin of attraction of the stable fixed points of (3) with probability \nicefrac12\nicefrac{{1}}{{2}}.

4. Convergence to spurious solutions

Let us pause to interpret this result. The stable fixed points when α<1/4\alpha<1/4 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 1/21/2 (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 NN and λ\lambda.

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 K=4K=4, W1=−W2=μW_{1}=-W_{2}=\mu, W3=W4=νW_{3}=W_{4}=\nu, and vi=cv_{i}=c for i=1,…,4i=1,\ldots,4 for some c>0c>0.

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 y^(X)\hat{y}(X) to be the natural rounding of σ(v⋅g(WX))\sigma(v\cdot g(WX)), where σ\sigma and gg are the sigmoid and ReLU as in Section 4. We take WW to be a K×NK\times N matrix and vv to be a KK-vector.

where again σ,g\sigma,g are applied component wise and again p(v,W):=(α/2)(∣∣v∣∣2+∣∣W∣∣2)p(v,W):=(\alpha/2)(\lvert\lvert v\rvert\rvert^{2}+\lvert\lvert W\rvert\rvert^{2}).

In Lemma 9.1 below, we show that the law of the loss at a point (v,W)(v,W) depends only on the following 4K+(K2)4K+\binom{K}{2} variables: for 1≤i≤j≤K1\leq i\leq j\leq K,

By similar reasoning, it can be shown that these functions are expressible as functions of un\mathbf{u}_{n} alone (see Section 9 below). We then find the following effective ballistic dynamics.

Let un\mathbf{u}_{n} be as in (5.1) and fix any λ>0\lambda>0 and δn=cδ/N\delta_{n}=c_{\delta}/N. Then un(t)\mathbf{u}_{n}(t) converges to the solution of the ODE system u˙t=−f(ut)+g(ut)\dot{\mathbf{u}}_{t}=-\mathbf{f}(\mathbf{u}_{t})+\mathbf{g}(\mathbf{u}_{t}), initialized from lim⁡n(un)∗μn\lim_{n}(\mathbf{u}_{n})_{*}\mu_{n} with

and correctors gvi=gmiμ=gmiν=0g_{v_{i}}=g_{m_{i}^{\mu}}=g_{m_{i}^{\nu}}=0, and gRij⊥=cδvivjλBijg_{R_{ij}^{\perp}}=c_{\delta}\frac{v_{i}v_{j}}{\lambda}\mathbf{B}_{ij} for 1≤i≤j≤K1\leq i\leq j\leq K.

3. Low variance asymptotics

As with the binary GMM, one can develop the large λ\lambda limit of these asymptotics after n→∞n\to\infty. 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 α>1/8\alpha>1/8, then the only fixed point is at un=0\mathbf{u}_{n}=\boldsymbol{0}.

If 0<α<1/80<\alpha<1/8, then let (I0,Iμ+,Iμ−,Iν+,Iν−)(I_{0},I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}) be any disjoint (possibly empty) subsets whose union is {1,...,K}\{1,...,K\}. Corresponding to that tuple (I0,Iμ+,Iμ−,Iν+,Iν−)(I_{0},I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}), is a set of fixed points that have Rij⊥=0R_{ij}^{\perp}=0 for all i,ji,j, and have

miμ=miν=vi=0m_{i}^{\mu}=m_{i}^{\nu}=v_{i}=0 for i∈I0i\in I_{0},

miμ=vi>0m_{i}^{\mu}=v_{i}>0 such that ∑i∈Iμ+vi2=\mboxlogit(−4α)\sum_{i\in I_{\mu}^{+}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miν=0m_{i}^{\nu}=0 for all i∈Iμ+i\in I_{\mu}^{+},

−miμ=vi>0-m_{i}^{\mu}=v_{i}>0 such that ∑i∈Iμ−vi2=\mboxlogit(−4α)\sum_{i\in I_{\mu}^{-}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miν=0m_{i}^{\nu}=0 for all i∈Iμ−i\in I_{\mu}^{-},

miν=vi<0m_{i}^{\nu}=v_{i}<0 such that ∑i∈Iν+vi2=\mboxlogit(−4α)\sum_{i\in I_{\nu}^{+}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miμ=0m_{i}^{\mu}=0 for all i∈Iν+i\in I_{\nu}^{+},

−miν=vi<0-m_{i}^{\nu}=v_{i}<0 such that ∑i∈Iν−vi2=\mboxlogit(−4α)\sum_{i\in I_{\nu}^{-}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miμ=0m_{i}^{\mu}=0 for all i∈Iν−i\in I_{\nu}^{-}.

In the K=4K=4 case, these form 3939 connected sets of fixed points, and of which 4!=244!=24 are fixed points that are stable, corresponding to the possible permutations in which each of Iμ+,Iμ−,Iν+,Iν−I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-} are singletons.

Similar to the binary GMM, in Figures 9–10, we demonstrate numerically that the following predicted fixed points from the λ→∞\lambda\to\infty limit match those arising at finite large nn and λ>0\lambda>0.

In the K=4K=4 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 μn\mu_{n} where vi∼N(0,1)v_{i}\sim\mathcal{N}(0,1) and Wi∼N(0,IN/N)W_{i}\sim\mathcal{N}(0,I_{N}/N) independently, this converges to \nicefrac332\nicefrac{{3}}{{32}}. 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 KK, we can also study the probability of ballistic convergence to a stable vs. unstable fixed point as one varies KK. This addresses the regime of overparametrization for the XOR GMM since K=4K=4 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 11 exponentially fast as KK 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 Iμ+,Iμ−,Iν+,Iν−I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-} 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 vv 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 K=4K=4, or equivalently focus on a fixed point where all indices beyond the first four have vi=miμ=miν=0v_{i}=m_{i}^{\mu}=m_{i}^{\nu}=0. 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 (a1,μ,a2,μ)(a_{1,\mu},a_{2,\mu}) be such that a1,μ2+a2,μ2=Cαa_{1,\mu}^{2}+a_{2,\mu}^{2}=C_{\alpha} and (a3,ν,a4,ν)(a_{3,\nu},a_{4,\nu}) such that a3,ν2+a4,ν2=Cαa_{3,\nu}^{2}+a_{4,\nu}^{2}=C_{\alpha}, for Cα=−\mboxlogit(4α)C_{\alpha}=-\mbox{logit}(4\alpha). Take as fixed points about which we expand to be vi=miμ=ai,μ>0v_{i}=m_{i}^{\mu}=a_{i,\mu}>0 and vi=miν=ai,ν<0v_{i}=m_{i}^{\nu}=a_{i,\nu}<0 for i=3,4i=3,4. Namely, we let

Numerical simulations in Figure 11 confirm these degenerate diffusive limits at finite λ\lambda.

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 δn\delta_{n}-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 n→∞n\to\infty limit.

Our aim is to establish un→u\mathbf{u}_{n}\to\mathbf{u} weakly as random variables on C([0,∞))C([0,\infty)) where u\mathbf{u} solves (2.4). It is equivalent to show the same on C([0,T])C([0,T]) equipped with the sup-norm for every T>0T>0.

Recalling Definition 2.1, since un\mathbf{u}_{n} are δn\delta_{n}-localizable, the error term in (6) has

and similarly bj(s)=∫0sbj′(s′)ds′b_{j}(s)=\int_{0}^{s}b^{\prime}_{j}(s^{\prime})ds^{\prime}, then recalling that un(s)\mathbf{u}_{n}(s) is the linear interpolation of (uj([s/δ]))j(u_{j}([s/\delta]))_{j}, we may write

where an(s)=(aj(s))j\mathbf{a}_{n}(s)=(a_{j}(s))_{j} and bn(s)=(bj(s))j\mathbf{b}_{n}(s)=(b_{j}(s))_{j}.

We now prove that the sequence (un(s∧τKn))(\mathbf{u}_{n}(s\wedge\tau_{K}^{n})) is tight in C([0,T])C([0,T]) with limit points which are (1/4)(1/4)-Holder for each KK. To this end, let us define

As the o(1)o(1) error above is uniform in tt, we have that

Thus it suffices to show the claimed tightness and Holder properties of limit points for vn\mathbf{v}_{n} instead of un\mathbf{u}_{n}. We aim to show that for all 0≤s,t≤T0\leq s,t\leq T,

from which we will get that the sequence vn(s∧τK)\mathbf{v}_{n}(s\wedge\tau_{K}) is uniformly 1/41/4-Hölder by Kolmogorov’s continuity theorem. Evidently, for all s,ts,t we have that

We control these terms in turn. We will do this coordinate wise and, for readability, fix some j≤kj\leq k and let u=uju=u_{j}, a=aja=a_{j}, b=bjb=b_{j} etc.

Let h=(hj)j≤k\mathbf{h}=(h_{j})_{j\leq k} be as in (2.2). Then the first term in (6) satisfies

by continuity of hjh_{j}. For the second term in (6),

which is ≲Kδ2(t−s)4\lesssim_{K}\delta^{2}(t-s)^{4} by items (1)–(2) δn\delta_{n}-localizability. (Applying this bound for s=0,t=Ts=0,t=T, the last term in aa is vanishing in the limit for each KK whenever δn=o(1)\delta_{n}=o(1).) 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 δn\delta_{n}-localizability.

For the second term in the martingale difference,

by items (1)–(2) of δn\delta_{n}-localizability. Finally, by the same reasoning, for the third term,

All of the above terms are O((t−s)2)O((t-s)^{2}) since 0≤s,t≤T0\leq s,t\leq T. Thus we have the claimed (6.2), and by Kolmogorov’s continuity theorem, (vn(s∧τK))s(\mathbf{v}_{n}(s\wedge\tau_{K}))_{s}, are uniformly \nicefrac14\nicefrac{{1}}{{4}}-Holder and thus the sequence is tight with \nicefrac14\nicefrac{{1}}{{4}}-Holder limit points. Notice furthermore that if we look at (vn(t∧τK)−an(t∧τK))t(\mathbf{v}_{n}(t\wedge\tau_{K})-\mathbf{a}_{n}(t\wedge\tau_{K}))_{t}, this sequence is also tight and the limits points are continuous martingales. Let us examine their limiting quadratic variations.

Let vnK(t)=vn(t∧τK)\mathbf{v}_{n}^{K}(t)=\mathbf{v}_{n}(t\wedge\tau_{K}) and define anK(t)\mathbf{a}_{n}^{K}(t) and bnK(t)\mathbf{b}_{n}^{K}(t) analogously. Furthermore, let vK(t)\mathbf{v}^{K}(t), aK(t)\mathbf{a}^{K}(t) and bK(t)\mathbf{b}^{K}(t) be their respective limits which we have shown to exist and be \nicefrac14\nicefrac{{1}}{{4}}-Holder.

is a martingale. We therefore need to consider the limit as n→∞n\to\infty of the integral above. Write

Consider the integrals of δ\delta times each of these four terms separately. For the first term,

goes to zero as n→∞n\to\infty by the assumption in (2.3).

We now reason that the integrals of δ\delta times the other three terms in (6.7) all go to zero as n→∞n\to\infty. The second and third are identical: by Cauchy–Schwarz,

The first expectation contributes δ−1/2\delta^{-1/2} by the first part of item (3) of localizability. Also,

The first of these terms is at most δ−1\delta^{-1} as argued in (6). The second is o(δ−3/2)o(\delta^{-3/2}) by the second part of item (3) in the definition of localizability. As such, we are able to conclude that

The integral of δ\delta times the fourth term in (6.7) is handled similarly using Cauchy–Schwarz and the bound of o(δ−3/2)o(\delta^{-3/2}) on (6.8).

Thus, if we consider the continuous martingales given by bK(t)\mathbf{b}^{K}(t), 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 α≥0\alpha\geq 0 fixed, the loss is given by

where c(Y)c(Y) only depends on YY. Note that H(x)=−2⟨W,x⊗k⟩H(x)=-2\langle W,x^{\otimes k}\rangle.

Our first aim is to establish Proposition 3.1, showing that the summary statistics un=(m,r⊥2)\mathbf{u}_{n}=(m,r_{\perp}^{2}) satisfy the conditions of Theorem 2.3 with the desired f,g\mathbf{f},\mathbf{g} and Σ\Sigma. We begin by checking localizability for un\mathbf{u}_{n}. In what follows, for ease of notation we will denote r2=r⊥2r^{2}=r_{\perp}^{2} and R2=m2+r2R^{2}=m^{2}+r^{2}. In these coordinates,

We check the items in Definition 2.1 one by one, beginning with item (1). Express the derivatives for un\mathbf{u}_{n} as

For item (2), differentiating (7.2), ∇Φ=∂1ϕ∇m+∂2ϕ∇r2\nabla\Phi=\partial_{1}\phi\nabla m+\partial_{2}\phi\nabla r^{2}, where

Notice that ⟨∇m,∇m⟩=1,⟨∇m,∇r2⟩=0,\left\langle\nabla m,\nabla m\right\rangle=1,\left\langle\nabla m,\nabla r^{2}\right\rangle=0, and ⟨∇r2,∇r2⟩=4r2\left\langle\nabla r^{2},\nabla r^{2}\right\rangle=4r^{2}. Consider

the bounding quantity is evidently a continuous function of m,r2m,r^{2} and therefore as long as xx is such that (m,r2)∈EK(m,r^{2})\in E_{K}, it is bounded by some C(K)C(K). Next, if we consider

where the bound on the operator norm of an i.i.d. Gaussian kk-tensor can be found, e.g., in [9, Lemma 4.7]. Moving on to item (3), by the same reasoning, for every ww,

If w=∇m=vw=\nabla m=v then ∥w∥=1\|w\|=1 and if w=∇r2=2(x−mv)w=\nabla r^{2}=2(x-mv) then ∥w∥≤C(K)\|w\|\leq C(K), so in both cases this is at most C(k,K)n2C(k,K)n^{2}. Finally, ∇2u\nabla^{2}u is only non-zero if u=ru=r in which case it is I−vvTI-vv^{T}. Then,

by the second item in the definition of localizability, and evidently the right-hand side is O(δ−2)O(\delta^{-2}) if δn=O(1/n)\delta_{n}=O(1/n). ∎

Having checked localizability for un\mathbf{u}_{n}, we apply Theorem 2.3. To compute f\mathbf{f}, by the above,

In particular, for δ=cδ/n\delta=c_{\delta}/n, we have δLδm=0\delta\mathcal{L}^{\delta}m=0 and

from which we obtain in the limit that n→∞n\to\infty that gm=0g_{m}=0 and gr2=4cδkR2k−2g_{r^{2}}=4c_{\delta}kR^{2k-2}.

Together, these yield the ODE system of (3.1),

which in the α=0\alpha=0 case matches Proposition 3.1. Finally, to see that Σ=0\Sigma=0, consider

which when multiplied by δ=O(1/n)\delta=O(1/n) 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 u1=0u_{1}=0, then R2=u2R^{2}=u_{2} and there are two possible fixed points: either u2=0u_{2}=0 or u2u_{2} solves

Notice that if k=2k=2, this has a nontrivial solution of the form cδ−α2=u2c_{\delta}-\frac{\alpha}{2}=u_{2}, provided α<αc(2):=2cδ\alpha<\alpha_{c}(2):=2c_{\delta}, and if k>2k>2, this has a nontrivial solution provided α≤max⁡x≥0kxk−2(2cδ−2x)\alpha\leq\max_{x\geq 0}kx^{k-2}(2c_{\delta}-2x) at cδ(k−2)xk−3−(k−1)xk−2=0c_{\delta}(k-2)x^{k-3}-(k-1)x^{k-2}=0 i.e., cδ(k−2)k−1=x\frac{c_{\delta}(k-2)}{k-1}=x. This gives

Evidently when we take α=0\alpha=0, then its non-trivial solution is at u2=1u_{2}=1 for all k≥2k\geq 2.

Alternatively, if u1≠0u_{1}\neq 0 at a fixed point, then we can simplify further and get

For simplicity of calculations, set α=0\alpha=0 as is the case in Proposition 3.1. Then, we simply get u2=cδu_{2}=c_{\delta}. In the case of k=2k=2, we also find that there is a solution if and only if λ>cδ\lambda>c_{\delta}, in which case R2=λR^{2}=\lambda, from which together with R2=u12+u2R^{2}=u_{1}^{2}+u_{2}, we also get u1=±λ−cδu_{1}=\pm\sqrt{\lambda-c_{\delta}}.

In the general case of k>2k>2, we find that R2=cδ+λ−2k−2R4(k−1)k−2R^{2}={c_{\delta}}+\lambda^{-\frac{2}{k-2}}R^{\frac{4(k-1)}{k-2}}. This has real solutions (all of which have R≥u2=cδR\geq u_{2}=c_{\delta} as required) whenever λ>λc(k)\lambda>\lambda_{c}(k) defined as

(Interpreting 00=10^{0}=1, this returns λc(2)=cδ\lambda_{c}(2)=c_{\delta}.) With this λ\lambda, whenever λ>λc(k)\lambda>\lambda_{c}(k), the equation for R2R^{2} has exactly two real solutions, both of which are at least cδc_{\delta} 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 Φ\Phi to our set of summary statistics and obtain limiting equations for its evolution from (3.4).

Recalling (7.2), the fact that Φ\Phi is a localizable summary statistic follows from the facts that ∥∇m∥,∥∇r2∥≤C(K)\|\nabla m\|,\|\nabla r^{2}\|\leq C(K), and the fact that Φ\Phi is a smooth nn-independent function of m,r2m,r^{2}.

For simplicity of calculations let us stick to α=0\alpha=0.

Next, consider the corrector for Φ\Phi. For this, notice that

Recalling VV from (7.4), and taking δ=cδ/n\delta=c_{\delta}/n, all the terms in ∑ijVij∂i∂jΦ\sum_{ij}V_{ij}\partial_{i}\partial_{j}\Phi vanish in the limit except the contribution from the ∇2r2\nabla^{2}r^{2}, which yields gΦ=lim⁡n→∞δLδΦ=4cδk2R4(k−1)g_{\Phi}=\lim_{n\to\infty}\delta\mathcal{L}^{\delta}\Phi=4c_{\delta}k^{2}R^{4(k-1)} Finally, we wish to compute the volatility for the stochastic part of the evolution of Φ\Phi. For this, consider ∇ΦV∇ΦT\nabla\Phi V\nabla\Phi^{T} and notice that all the entries of that matrix are continuous functions of un\mathbf{u}_{n} and thus go to zero when multiplied by δ=O(1/n)\delta=O(1/n).

3. Diffusive limits at the equator

Taking limits as n→∞n\to\infty, as long as λ\lambda is fixed in nn, we see that f\mathbf{f} is given by

4. Diffusive limit for the radius

where we used in the first inequality that the law of HH is rotation invariant and HH is a kk-homogenous function. For the second part of item (3),

The ∂iH\partial_{i}H are Gaussian with mean zero, and by (7.4), variance Ck′R2(k−2)xi2+CkR2(k−1)C_{k}^{\prime}R^{2(k-2)}x_{i}^{2}+C_{k}R^{2(k-1)} and covariance CkxixjR2(k−2)C_{k}x_{i}x_{j}R^{2(k-2)}. Recall the following fact about Gaussians: if X,YX,Y are Gaussians with variances σ2\sigma^{2} and covariance tt, then \mboxCov(X2,Y2)≤Ct2σ4\mbox{Cov}(X^{2},Y^{2})\leq Ct^{2}\sigma^{4} for some universal constant CC. Also, \mboxVar(X2)≤Cσ4\mbox{Var}(X^{2})\leq C\sigma^{4}. Applying this to ∂iH\partial_{i}H, we get

Combined with the above, this gives a bound of n2=O(δ−2)n^{2}=O(\delta^{-2}) on the second part of item (3).

We now calculate the resulting drifts. For f\mathbf{f}, write

Combining terms and sending n→∞n\to\infty, we obtain

Multiplying by δ=1/n\delta=1/n and taking the limit as n→∞n\to\infty, the two entries of this matrix that survive are Σ11\Sigma_{11} and Σ22\Sigma_{22}, where Σ11=4k\Sigma_{11}=4k and Σ22=4k(k−1)\Sigma_{22}=4k(k-1). 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 un\mathbf{u}_{n} from (4.2).

Let Xμ∼N(μ,I/λ)X_{\mu}\sim\mathcal{N}(\mu,I/\lambda) and X−μ∼N(−μ,I/λ)X_{-\mu}\sim\mathcal{N}(-\mu,I/\lambda). Then, notice that

Next, notice that as a vector, (W1Xμ,W2Xμ)(W_{1}X_{\mu},W_{2}X_{\mu}) is distributed as (m1+Z1,μm1+Z1,⊥,m2+Z2,μm2+Z2,⊥)(m_{1}+Z_{1,\mu}m_{1}+Z_{1,\perp},m_{2}+Z_{2,\mu}m_{2}+Z_{2,\perp}), where Z1,μ,Z2,μZ_{1,\mu},Z_{2,\mu} are i.i.d. N(0,λ−1)\mathcal{N}(0,\lambda^{-1}), and Z1,⊥,Z2,⊥Z_{1,\perp},Z_{2,\perp} are jointly Gaussian with means zero and covariance

Similarly, the distribution of WX−μWX_{-\mu} also only depends on (mi,Rij⊥)i,j(m_{i},R_{ij}^{\perp})_{i,j}. Finally,

Therefore, at any point (v,W)(v,W), the law of L((v,W))L((v,W)), and thus Φ\Phi, is simply a function of un(v,W)\mathbf{u}_{n}(v,W). To see that the summary statistics satisfy the bounds of item (1) in Definition 2.1, write ∇=(∂v1,∂v2,∇W1,∇W2)\nabla=(\partial_{v_{1}},\partial_{v_{2}},\nabla_{W_{1}},\nabla_{W_{2}}). 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 un\mathbf{u}_{n} 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 ∇L\nabla L by direct calculation from (4.1).

(Notice that if w∈{μ,Wi,Wi⊥}w\in\{\mu,W_{i},W_{i}^{\perp}\}, then Ai⋅w\mathbf{A}_{i}\cdot w is only a function of un\mathbf{u}_{n} by the same reasoning as used in Lemma 8.1.) Then, we can also easily express

Finally, the matrix VV 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 δn\delta_{n}-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 ii, for every Rii⊥<∞R_{ii}^{\perp}<\infty and every mi>0m_{i}>0, we have

For every vi,Rij⊥v_{i},R_{ij}^{\perp} and mi≠0m_{i}\neq 0 for i,j=1,2i,j=1,2, we have

Throughout this section we will take μ=e1\mu=e_{1}. By rotational invariance of the problem, this is without loss of generality, and only simplifies certain expressions. The δn\delta_{n}-localizability can be seen by application of the moment bounds listed above.

The condition on un\mathbf{u}_{n} was satisfied per Lemma 8.1. Recalling ∇Φ\nabla\Phi from (8.12), one can verify that the norm of each of the four terms in ∇Φ\nabla\Phi is individually bounded, using the Cauchy–Schwarz inequality together with the bound of Lemma 8.2 on ∥Ai∥\|\mathbf{A}_{i}\|.

When uiu_{i} is viv_{i}, this is simply a fourth moment bound on ∇viH\nabla_{v_{i}}H, which follows from the 88’th moment by Jensen’s inequality. When uiu_{i} is mim_{i}, or Rij⊥R_{ij}^{\perp}, the bound follows from

for choices of ww being either μ\mu in which case ∥w∥=1\|w\|=1 or Wi⊥W_{i}^{\perp} in which case ∥w∥=Rii⊥\|w\|=R_{ii}^{\perp}. For each KK, this is at most some constant C(K)C(K) using the two bounds of Lemma 8.2.

The convergence of the population drift to f\mathbf{f} from Proposition 4.1 follows by taking the inner products of ∇L\nabla L from (8.12) with the rows of JJ from (8.8), and noticing that Aiμ\mathbf{A}_{i}^{\mu} from (4.3) is exactly Ai⋅μ\mathbf{A}_{i}\cdot\mu and Aij⊥\mathbf{A}_{ij}^{\perp} from (4.3) is exactly Ai⋅Wj⊥\mathbf{A}_{i}\cdot W_{j}^{\perp}.

Next consider the convergence of the correctors to the claimed g\mathbf{g}. The variables u∈{v1,v2,m1,m2}u\in\{v_{1},v_{2},m_{1},m_{2}\} are linear so Lnu=0\mathcal{L}_{n}u=0 and for these, gu=0\mathbf{g}_{u}=0. For u=Rij⊥u=R_{ij}^{\perp} for i,j∈{1,2}i,j\in\{1,2\}, the relevant entries in VV are those corresponding to Wi⊥W_{i}^{\perp} and Wj⊥W_{j}^{\perp}. For ease of notation, in what follows let π=σ(v⋅g(WX))\pi=\sigma(v\cdot g(WX)).

For ease of calculation taking μ=e1\mu=e_{1}, we have LnRij⊥=∑k≠1VWik,Wjk\mathcal{L}_{n}R_{ij}^{\perp}=\sum_{k\neq 1}V_{W_{ik},W_{jk}}, which by (8), and the choice of δn=cδ/N\delta_{n}=c_{\delta}/N, is given by

which we emphasize is only a function of un\mathbf{u}_{n}. We lastly need to show that the diffusion matrix Σn\Sigma_{n} goes to zero as n→∞n\to\infty when δn=O(1/n)\delta_{n}=O(1/n). This is straightforward to see by considering any element of JVJTJVJ^{T} and using Cauchy–Schwarz together with the two bounds of Lemma 8.2 to bound it in absolute value by some C(K)C(K) independent of nn. Then when multiplying by any δn=o(1)\delta_{n}=o(1), this entire matrix will evidently vanish. ∎

2. The small-noise limit of the effective dynamics

One can now take a λ→∞\lambda\to\infty limit to arrive at the ODE system of Proposition 4.2.

We begin with considering lim⁡λ→∞Aiμ\lim_{\lambda\to\infty}\mathbf{A}_{i}^{\mu}: its limiting value will depend on the signs of both m1m_{1} and m2m_{2}. We can express Aiμ\mathbf{A}_{i}^{\mu} from (4.3) as

We claim that the two terms on the right-hand side converge to −121mi>0σ(−v⋅g(m))-\frac{1}{2}\mathbf{1}_{m_{i}>0}\sigma(-v\cdot g(m)) and −121mi<0σ(v⋅g(−m))-\frac{1}{2}\mathbf{1}_{m_{i}<0}\sigma(v\cdot g(-m)) respectively. This follows by e.g., writing the difference as

at which point, we see that if m1,m2≥0m_{1},m_{2}\geq 0, this becomes 12σ(−v⋅m)\frac{1}{2}\sigma(-v\cdot m), as it is if m1,m2≤0m_{1},m_{2}\leq 0. If m1≥0m_{1}\geq 0 and m2≤0m_{2}\leq 0, then you get lim⁡λA1μ=−12σ(−v1m1)\lim_{\lambda}\mathbf{A}_{1}^{\mu}=-\frac{1}{2}\sigma(-v_{1}m_{1}) and lim⁡λA2μ=−12σ(−v2m2)\lim_{\lambda}\mathbf{A}_{2}^{\mu}=-\frac{1}{2}\sigma(-v_{2}m_{2}) and likewise if m1≤0m_{1}\leq 0 and m2≥0m_{2}\geq 0.

Next consider the limit as λ→∞\lambda\to\infty of Aij⊥\mathbf{A}_{ij}^{\perp} from (4.3), which we claim converges to . Write

Finally, since ∣Bij∣≤1|\mathbf{B}_{ij}|\leq 1, the quantity gRij⊥=cδvivjλBijg_{R_{ij}^{\perp}}=c_{\delta}\frac{v_{i}v_{j}}{\lambda}\mathbf{B}_{ij} evidently goes to zero as λ→∞\lambda\to\infty. ∎

The above argument used mi≠0m_{i}\neq 0 for the limit of Aiμ\mathbf{A}_{i}^{\mu}. If one considers the cases when mi=0m_{i}=0, the limiting drifts still apply. For this, it suffices to show that if mi=0m_{i}=0, then Aiμ\mathbf{A}_{i}^{\mu} converges to zero. Without loss of generality, suppose m1=0m_{1}=0 and consider

This is zero independently of λ\lambda by independence of Z1,μZ_{1,\mu} from the other Gaussians in the expectation.

Evidently, every fixed point must have Rij⊥=0R_{ij}^{\perp}=0. Furthermore, if we let ui=vi−miu_{i}=v_{i}-m_{i}, then

and therefore every fixed point of the ODE system must have ui=0u_{i}=0, which is to say vi=miv_{i}=m_{i}. Therefore, it suffices to characterize the fixed points in terms of (v1,v2)(v_{1},v_{2}) as claimed. This reduces to viσ(−∥v∥2)=2αviv1v2>0v_{i}\sigma(-\|v\|^{2})=2\alpha v_{i}v_{1}v_{2}>0 if v1v2>0v_{1}v_{2}>0 and viσ(−vi2)=2αviv_{i}\sigma(-v_{i}^{2})=2\alpha v_{i} otherwise. Observe first that the point (v1,v2)=(0,0)(v_{1},v_{2})=(0,0) is a fixed point of this system. If (v1,v2)≠0(v_{1},v_{2})\neq 0, then dividing out by viv_{i}, the above reduces to σ(−∥v∥2)=2α\sigma(-\|v\|^{2})=2\alpha if v1v2>0v_{1}v_{2}>0 and σ(−vi2)=2α\sigma(-v_{i}^{2})=2\alpha otherwise. Recalling that Cα=−logit⁡(2α)=log⁡(1−2α)−log⁡(2α)C_{\alpha}=-\operatorname{logit}(2\alpha)=\log(1-2\alpha)-\log(2\alpha) we obtain the claimed set of fixed points by inverting these equations (they only have a solution if α<1/4\alpha<1/4).

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 λ=∞\lambda=\infty 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 Φ\Phi. In particular, the values of Φ\Phi at its critical points are given by Φ0=log⁡2\Phi_{0}=\log 2 at v1=v2=0v_{1}=v_{2}=0, Φ+=12(log⁡2+log⁡(1+e−Cα)+αCα\Phi_{+}=\frac{1}{2}(\log 2+\log(1+e^{-C_{\alpha}})+\alpha C_{\alpha} when v1v2>0v_{1}v_{2}>0, and Φ−=log⁡(1+e−Cα)+2αCα\Phi_{-}=\log(1+e^{-C_{\alpha}})+2\alpha C_{\alpha} when v1v2<0v_{1}v_{2}<0. It is a simple calculus exercise to show that the smallest of these is Φ0\Phi_{0} when α>1/4\alpha>1/4 and Φ−\Phi_{-} when α<1/4\alpha<1/4.

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 vi=miv_{i}=m_{i} and Rij⊥=0R_{ij}^{\perp}=0 with v1v2≤0v_{1}v_{2}\leq 0 is unstable, by showing a repelling direction arbitrarily close to the point v1=−Cαv_{1}=-\sqrt{C_{\alpha}}, v2=0v_{2}=0. If v1=−Cαv_{1}=-\sqrt{C_{\alpha}} and v2=ϵ>0v_{2}=\epsilon>0, then v˙2\dot{v}_{2} there reduces to ϵ(σ(−ϵ2)2−α)\epsilon(\frac{\sigma(-\epsilon^{2})}{2}-\alpha), and as long as α<1/4\alpha<1/4, there exists ϵ>0\epsilon>0 such that σ(−ϵ2)>2α\sigma(-\epsilon^{2})>2\alpha so v˙2>0\dot{v}_{2}>0 for all ϵ\epsilon small enough.

3. Rescaled effective dynamics around unstable fixed points

Now Taylor expanding the sigmoid function, and using the definition of CαC_{\alpha}, 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 un\mathbf{u}_{n} from (5.1). The next lemma shows that un\mathbf{u}_{n} form a good set of summary statistics.

where Zi,ιZ_{i,\iota} are i.i.d. N(0,λ−1)\mathcal{N}(0,\lambda^{-1}) and (Zi⊥)(Z_{i\perp}) are jointly Gaussian with covariance matrix

Similarly, the law of WX−ιWX_{-\iota} depends only on (miι,Rij⊥)(m_{i}^{\iota},R_{ij}^{\perp}). Finally,

Therefore, at a fixed point (v,W)(v,W) the law of L(v,W)L(v,W) is only a function of un(v,W)\mathbf{u}_{n}(v,W).

where δij\delta_{ij} is 11 if i=ji=j and otherwise. For higher derivatives, we only have second derivatives in the Rjk⊥R_{jk}^{\perp} 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 22. The third derivatives of all elements of un\mathbf{u}_{n} are zero. ∎

By the same reasoning as in Lemma 9.1, if w∈{μ,ν,Wi,Wi⊥}w\in\{\mu,\nu,W_{i},W_{i}^{\perp}\}, then w⋅Aiw\cdot\mathbf{A}_{i} is only a function of un\mathbf{u}_{n}. We then also have the conclusions of Lemma 8.2 for XX 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 un\mathbf{u}_{n} was satisfied per Lemma 9.1. Recalling ∇Φ\nabla\Phi from (8.12), one can verify that the norm of each of the four terms in ∇Φ\nabla\Phi is individually bounded, using the Cauchy–Schwarz inequality together with the bound of Lemma 8.2 on ∥Ai∥\|\mathbf{A}_{i}\|, 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 f\mathbf{f} from Proposition 4.1 follows by taking the inner products of ∇L\nabla L from (8.12) with the rows of JJ from (9.3), and noticing that Aiμ\mathbf{A}_{i}^{\mu} is exactly Ai⋅μ\mathbf{A}_{i}\cdot\mu, Aiν\mathbf{A}_{i}^{\nu} is exactly ν⋅Ai\nu\cdot\mathbf{A}_{i}, and Aij⊥\mathbf{A}_{ij}^{\perp} is exactly Ai⋅Wj⊥\mathbf{A}_{i}\cdot W_{j}^{\perp}.

We next consider the population correctors. The fact that gvi=gmiμ=gmiν=0g_{v_{i}}=g_{m_{i}^{\mu}}=g_{m_{i}^{\nu}}=0 follows from the fact that the Hessians of vi,miμ,miνv_{i},m_{i}^{\mu},m_{i}^{\nu} are zero. For the corrector gRij⊥g_{R_{ij}^{\perp}} for 1≤i≤j≤K1\leq i\leq j\leq K, the relevant entries of VV are those corresponding to Wi⊥W_{i}^{\perp} and Wj⊥W_{j}^{\perp}. For ease of notation, in what follows let π=σ(v⋅g(WX))\pi=\sigma(v\cdot g(WX)).

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 δn=o(1)\delta_{n}=o(1). 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 λ→∞\lambda\to\infty 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 λ→∞\lambda\to\infty limit, the ODE from Proposition 5.1 converges to

and R˙ij⊥=−2αRij⊥\dot{R}_{ij}^{\perp}=-2\alpha R_{ij}^{\perp} for 1≤i≤j≤K1\leq i\leq j\leq K.

Let us begin with convergence of Aiμ\mathbf{A}_{i}^{\mu}. We claim that it converges to

The point will be that when taking the inner product with μ\mu, the first two terms here contribute to the limit and the latter two vanish, while when taking the inner product with ν\nu, the first two terms vanish in the λ→∞\lambda\to\infty limit while the latter two contribute.

Consider e.g., the first of the four terms above, and inner product with μ\mu. In this case, consider

which is precisely the quantity that was exactly shown to go to zero as λ→∞\lambda\to\infty in (8.20). To see that the third and fourth terms above go to zero when taking their inner product with μ\mu, observe that they become

which by orthogonality of μ\mu and ν\nu is at most λ−1/2\lambda^{-1/2} by the reasoning of Lemma 8.2, therefore vanishing as λ→∞\lambda\to\infty. Together with its analogue for X−νX_{-\nu}, this implies the claim for the convergence of Aiμ\mathbf{A}_{i}^{\mu}, as well as its analogous limit of Aiν\mathbf{A}_{i}^{\nu}.

We next consider the limit as λ→∞\lambda\to\infty of Aij⊥\mathbf{A}_{ij}^{\perp}, which we claim goes to . Using the expansion of Ai\mathbf{A}_{i} from earlier in this proof, we can consider Aij⊥=Ai⋅Wj⊥\mathbf{A}_{ij}^{\perp}=\mathbf{A}_{i}\cdot W_{j}^{\perp} as four terms having the form of the terms in (8.21), which were there showed to go to zero as λ→∞\lambda\to\infty. Since Wj⊥W_{j}^{\perp} here is orthogonal both to μ\mu and ν\nu, the same proof applies.

Finally, in order to see that the limit as λ→∞\lambda\to\infty of gRij⊥=cδvivjλBijg_{R_{ij}^{\perp}}=c_{\delta}\frac{v_{i}v_{j}}{\lambda}\mathbf{B}_{ij} is zero, which follows from the fact that ∣Bij∣≤1|\mathbf{B}_{ij}|\leq 1. ∎

The fixed points of the ODE system of Proposition 9.1 are classified as follows. If α>1/8\alpha>1/8, then the only fixed point is at un=0\mathbf{u}_{n}=\boldsymbol{0}.

If 0<α<1/80<\alpha<1/8, then let (I0,Iμ+,Iμ−,Iν+,Iν−)(I_{0},I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}) be any disjoint (possibly empty) subsets whose union is {1,...,K}\{1,...,K\}. Corresponding to that tuple (I0,Iμ+,Iμ−,Iν+,Iν−)(I_{0},I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}), is a set of fixed points that have Rij⊥=0R_{ij}^{\perp}=0 for all i,ji,j, and have

miμ=miν=vi=0m_{i}^{\mu}=m_{i}^{\nu}=v_{i}=0 for i∈I0i\in I_{0},

miμ=vi>0m_{i}^{\mu}=v_{i}>0 such that ∑i∈Iμ+vi2=\mboxlogit(−4α)\sum_{i\in I_{\mu}^{+}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miν=0m_{i}^{\nu}=0 for all i∈Iμ+i\in I_{\mu}^{+},

−miμ=vi>0-m_{i}^{\mu}=v_{i}>0 such that ∑i∈Iμ−vi2=\mboxlogit(−4α)\sum_{i\in I_{\mu}^{-}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miν=0m_{i}^{\nu}=0 for all i∈Iμ−i\in I_{\mu}^{-},

miν=vi<0m_{i}^{\nu}=v_{i}<0 such that ∑i∈Iν+vi2=\mboxlogit(−4α)\sum_{i\in I_{\nu}^{+}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miμ=0m_{i}^{\mu}=0 for all i∈Iν+i\in I_{\nu}^{+},

−miν=vi<0-m_{i}^{\nu}=v_{i}<0 such that ∑i∈Iν−vi2=\mboxlogit(−4α)\sum_{i\in I_{\nu}^{-}}v_{i}^{2}=\mbox{logit}(-4\alpha) and miμ=0m_{i}^{\mu}=0 for all i∈Iν−i\in I_{\nu}^{-}.

In the K=4K=4 case, these form 3939 connected sets of fixed points, and of which 4!=244!=24 are fixed points that are stable, corresponding to the possible permutations in which each of Iμ+,Iμ−,Iν+,Iν−I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-} are singletons.

Evidently, any fixed point must have Rij⊥=0R_{ij}^{\perp}=0 for all i,ji,j. Furthermore, the point vi=miμ=miν=0v_{i}=m_{i}^{\mu}=m_{i}^{\nu}=0 for i=1,...,Ki=1,...,K evidently forms a fixed point of the system. Now suppose there is some fixed point with vi=0v_{i}=0 for some ii; in that case, it must be that miμ=0m_{i}^{\mu}=0 and miν=0m_{i}^{\nu}=0. Therefore, we can select a subset I0I_{0} of {1,...,K}\{1,...,K\} such that vi=miμ=miν=0v_{i}=m_{i}^{\mu}=m_{i}^{\nu}=0 for i∈I0i\in I_{0}.

For any such choice of I0I_{0}, consider next, i∉I0i\notin I_{0}. We first claim that if vi>0v_{i}>0 at a fixed point, then miμ∈{±vi}m_{i}^{\mu}\in\{\pm v_{i}\} and miν=0m_{i}^{\nu}=0, whereas if vi<0v_{i}<0 then miν∈{±vi}m_{i}^{\nu}\in\{\pm v_{i}\} and miμ=0m_{i}^{\mu}=0. To see this, notice that at any fixed point,

Since σ\sigma is non-negative, if vi>0v_{i}>0, the sign of the right-hand side of the first equation is the same as the sign of miμm_{i}^{\mu} 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 miνm_{i}^{\nu}, so any such fixed point must have miν=0m_{i}^{\nu}=0. To see that miμ=±vim_{i}^{\mu}=\pm v_{i} at such a fixed point, now set miν=0m_{i}^{\nu}=0 and take the fixed point equations for viv_{i} and miμm_{i}^{\mu}, dividing one by viv_{i} and the other by miμm_{i}^{\mu} to see that

as claimed. The fixed points having vi<0v_{i}<0 are solved symmetrically.

Our classification now reduces to understanding the possible values taken by (v1,...,vK)(v_{1},...,v_{K}) given their signs (when non-zero). Fix a partition (I0,Iμ+,Iμ−,Iν+,Iν−)(I_{0},I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}) of {1,...,K}\{1,...,K\} and consider the set of fixed points having miμ=miν=vi=0m_{i}^{\mu}=m_{i}^{\nu}=v_{i}=0 for i∈I0i\in I_{0}, miμ=vi>0m_{i}^{\mu}=v_{i}>0 on Iμ+I_{\mu}^{+} 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 viv_{i} on each of these sets are as described by the proposition.

In order to see this, fix e.g., i∈Iμ+i\in I_{\mu}^{+}. Then, miμ=vim_{i}^{\mu}=v_{i} and miν=0m_{i}^{\nu}=0, and so the fixed point equations reduce to

since the only coordinates where g(mμ)g(m^{\mu}) will be non-zero are j∈Iμ+j\in I_{\mu}^{+}, where mjμ=vjm_{j}^{\mu}=v_{j}. Inverting the sigmoid function, this implies exactly the claimed ∑j∈Iμ+vj2=\mboxlogit(−4α)\sum_{j\in I_{\mu}^{+}}v_{j}^{2}=\mbox{logit}(-4\alpha). The cases of Iμ−,Iν+,Iν−I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-} are analogous, concluding the proof.

The count of the number of connected components of fixed points this forms is sensitive to KK, so for concreteness let us perform it when K=4K=4. We first notice that the fixed point at (0,...,0)(0,...,0) is disconnected from all others. Fixed points corresponding to some (I0,...,Iν−)(I_{0},...,I_{\nu}^{-}) are part of the same connected component of fixed points if one goes from one to the other by moving an element of IιηI_{\iota}^{\eta} (for some ι∈{μ,ν}\iota\in\{\mu,\nu\} and η∈{±}\eta\in\{\pm\} to I0I_{0} without making IιηI_{\iota}^{\eta} empty, or by moving an element of I0I_{0} to a non-empty IιηI_{\iota}^{\eta}.

We turn now to studying the stability of these various sets of fixed points. Observe that in the λ→∞\lambda\to\infty limit, the dynamical system of Proposition 9.1 is a gradient system for the population loss

At a fixed point (which necessarily has vi=miv_{i}=m_{i}, Rii⊥=0R_{ii}^{\perp}=0, and is characterized by the partition of {1,...,4}\{1,...,4\} into Iμ+,Iμ−,Iν+,Iν−I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}, this reduces to

At this point, noticing that ∑i∈Iμ+vi2\sum_{i\in I_{\mu}^{+}}v_{i}^{2} is equal to Cα=−\mboxlogit(4α)C_{\alpha}=-\mbox{logit}(4\alpha) if Iμ+I_{\mu}^{+} is non-empty and if it is empty, and similarly for Iμ−,Iν+,Iν−I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}, this turns into a simple optimization problem over the number of non-empty Iμ+,Iμ−,Iν+,Iν−I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}. Just as in the binary GMM case, it becomes evident that when α>1/8\alpha>1/8, this is minimized at vi=0v_{i}=0 for all ii (i.e., they are all empty and I0={1,...,4}I_{0}=\{1,...,4\}, whereas when α<1/8\alpha<1/8 the above is minimized when every one of Iμ+,Iμ−,Iν+,Iν−I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-} are all non-empty. This yields the global minima of Φ\Phi 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 K>4K>4, 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 Φ\Phi will still be (0,...,0)(0,...,0) when α>1/8\alpha>1/8 and will be all fixed points that have all four of Iμ+,...,Iν−I_{\mu}^{+},...,I_{\nu}^{-} being non-empty if α<1/8\alpha<1/8. Notice that when α<1/8\alpha<1/8 and K>4K>4, 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 K=4K=4 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 3/323/32. Begin by noticing that if the first layer weights are initialized as Wi∼N(0,IN/N)W_{i}\sim\mathcal{N}(0,I_{N}/N) independently for i=1,...,4i=1,...,4 and the second layer weights viv_{i} are independent standard Gaussians, then the projection onto the coordinate system (vi,miμ,miν,Rij)(v_{i},m_{i}^{\mu},m_{i}^{\nu},R_{ij}) is given by

The δ\delta-functions at zero for miμ,miνm_{i}^{\mu},m_{i}^{\nu} however cause some trouble because of the indicator functions on the sign of miμm_{i}^{\mu} and miνm_{i}^{\nu} in the equations of Proposition 9.1.

Under the flow of Proposition 9.1, if vi(0)v_{i}(0) is positive, then miνm_{i}^{\nu} stays fixed at zero, and if miμ(0)=0−m_{i}^{\mu}(0)=0^{-} then miμm_{i}^{\mu} becomes negative infinitesimally quickly, whereas if miμ(0)=0+m_{i}^{\mu}(0)=0^{+} then it becomes positive infinitesimally quickly. At any rate, the sign of viv_{i} never changes to negative from such an initialization, and similarly if vi(0)v_{i}(0) is negative, the sign of viv_{i} 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 (vi(0))i(v_{i}(0))_{i} have positive sign and two of them have negative sign; evidently this has probability (42)/24=3/8\binom{4}{2}/2^{4}=3/8.

Given that two of vi(0)v_{i}(0) are positive, and two of them are negative—say without loss of generality that i=1,2i=1,2 are the coordinates in which it is positive, and i=3,4i=3,4 are the coordinates in which it is negative—then the dynamical system for (v1,v2,m1μ,m2μ)(v_{1},v_{2},m_{1}^{\mu},m_{2}^{\mu}) 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 1/21/2. Similarly, the dynamical system for (v2,v4,m3ν,m4ν)(v_{2},v_{4},m_{3}^{\nu},m_{4}^{\nu}) independently gives a further probability 1/21/2 of converging to its good classifier. Together, these yield a probability of 3/323/32 of converging to one of the 4!4! many optimal classifiers for the XOR GMM.

Generically, if K≥4K\geq 4, 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 Iμ+,Iμ−,Iν+,Iν−I_{\mu}^{+},I_{\mu}^{-},I_{\nu}^{+},I_{\nu}^{-}. This is the probability that vi(0)v_{i}(0) are positive for at least two indices, and negative for at least two indices, and then among the indices at which vi(0)v_{i}(0) is positive, there is at least one index where miμm_{i}^{\mu} is positive and one where it is negative, and similarly with vi(0)v_{i}(0) negative and miνm_{i}^{\nu}. 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 11 exponentially fast as K→∞K\to\infty since the initial choice of vi(0)v_{i}(0)’s will typically have around K/2K/2 positive and K/2K/2 negative coordinates, and with exponentially high probability those will have both positive and negative miμm_{i}^{\mu} and miνm_{i}^{\nu}.

5. Diffusive limit on critical submanifolds

Plugging these in, and taking the n→∞n\to\infty limit we find that for i=1,2i=1,2,

By a similar reasoning, for i=3,4i=3,4, we have

and if i∈{1,2}i\in\{1,2\} and j∈{3,4}j\in\{3,4\}, then

By a similar reasoning, if i,j∈{1,2}i,j\in\{1,2\}, then

and if i∈{1,2}i\in\{1,2\} and j∈{3,4}j\in\{3,4\}, 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 Z∼N(0,I)Z\sim\mathcal{N}(0,I) and consider

The quantities in the expectations are at most some universal constant times (w⋅μ)8+λ−4(w⋅Z)8(w\cdot\mu)^{8}+\lambda^{-4}(w\cdot Z)^{8}. To bound the expectation of the second term here, notice that w⋅Zw\cdot Z is distributed as N(0,∥w∥2)\mathcal{N}(0,\|w\|^{2}) implying the desired.

The bound on Ai\mathbf{A}_{i} goes as follows. Evidently it suffices to let Xμ=μ+λ−1/2ZX_{\mu}=\mu+\lambda^{-1/2}Z for Z∼N(0,I)Z\sim\mathcal{N}(0,I), and prove the bound on the norm of

Now decompose ZZ as Zμμ+Z1,⊥W1⊥+Z2,⊥W2⊥+Z3Z_{\mu}\mu+Z_{1,\perp}W_{1}^{\perp}+Z_{2,\perp}W_{2}^{\perp}+Z_{3}, where Zμ∼N(0,1)Z_{\mu}\sim\mathcal{N}(0,1) is independent of (Z1,⊥,Z2,⊥)(Z_{1,\perp},Z_{2,\perp}) which is distributed as N(0,A)\mathcal{N}(0,A) with AA given by (8.3), which is independent of Z3Z_{3} distributed as a standard Gaussian vector orthogonal to the subspace spanned by (μ,W1⊥,W2⊥)(\mu,W_{1}^{\perp},W_{2}^{\perp}). By independence of Z3Z_{3} 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 mi>0m_{i}>0 this goes to zero as λ→∞\lambda\to\infty.

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 λ\lambda, per Fact 8.1. For the second desired bound, expand evig(Wi⋅Xμ)−evig(mi)e^{v_{i}g(W_{i}\cdot X_{\mu})}-e^{v_{i}g(m_{i})} as

It suffices to show the expectation of the square of each of these goes to zero as λ→∞\lambda\to\infty. First,

If mi≠0m_{i}\neq 0, the expectation on the right goes to zero by (8.16). Second,

When mi<0m_{i}<0, this is evidently zero; when mi>0m_{i}>0, if Gλ∼N(0,I/λ)G_{\lambda}\sim\mathcal{N}(0,I/\lambda), this is

which goes to zero as O(λ−1)O(\lambda^{-1}) when λ→∞\lambda\to\infty, by the explicit formula for the moment generating function of the Gaussian Wi⋅GλW_{i}\cdot G_{\lambda}, whose variance is (mi2+Rii⊥)λ−1(m_{i}^{2}+R_{ii}^{\perp})\lambda^{-1}. ∎

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.

References