Phase diagram of Stochastic Gradient Descent in high-dimensional two-layer neural networks
Listen
Radio episode about this paper
Transcript
Introduction to the show: ident: AI Radio. Generated commentary on the latest Artificial Intelligence papers.
Tom: Next we'll be talking about the paper "Phase diagram of Stochastic Gradient Descent in high-dimensional two-layer neural networks".
Jane: The paper was written by Rodrigo Veiga, Ludovic Stephan, Bruno Loureiro, Florent Krzakala and Lenka Zdeborová from Ecole Polytechnique Fédérale de Lausanne and Universidade de São Paulo.
Tom: Stay tuned as we take you through the paper and discuss its implications.
Title: Tom: Welcome back, everyone. Today we're looking at a paper that's been making waves in the theory community — "Phase diagram of Stochastic Gradient Descent in high-dimensional two-layer neural networks." Jane, what caught your eye first about this one?
Jane: Tom, honestly, the title alone tells you they're trying to do something ambitious. They want to map out, like a phase diagram in physics, all the different behaviors you can get when you train a simple neural network with stochastic gradient descent. And they're doing it in high dimensions, which is where things get really interesting.
Tom: Right, and I love that they're borrowing that phrase "phase diagram" from physics. It's not just a metaphor — they literally draw regions on a graph showing where learning works perfectly, where it gets stuck, and where the math breaks down entirely.
Jane: Exactly. And the key players here are the learning rate and the number of hidden units in the network. The authors show that how you scale those two things as the input dimension grows determines everything about whether your network actually learns.
Tom: So it's like a recipe — if you use the wrong proportions, your cake flops?
Jane: That's actually a pretty good way to put it. If you scale the learning rate too aggressively relative to the network width, you land in what they call the "bad learning" region. The noise in your data just dominates, and the network basically learns nothing useful.
Tom: And if you get the proportions right, you get perfect learning — even when there's noise in the labels, which is pretty remarkable. The classical result from the 90s, the Saad and Solla work, only covered one specific point on this diagram. This paper generalizes that to the whole map.
Jane: Right, and that's what makes it exciting. It connects two big bodies of work — the older statistical physics approach and the more recent mean-field theory — and shows they're really talking about different regions of the same underlying picture.
Tom: So we've got a unified framework now?
Jane: A unified framework with a very clear visual. And the implications are pretty direct for anyone training neural networks in practice.
Tom: I'm curious about how they actually prove this stuff. That's what we're digging into next.
Summary: Tom: So Jane, let's get into the meat of it. The paper's summary is really about one central claim: that the dynamics of stochastic gradient descent can be captured by a set of deterministic ordinary differential equations. And that's a huge deal, right?
Jane: Huge. Because normally, when you train a neural network, the updates are random — each sample you feed in gives you a slightly different gradient. But the authors show that in the high-dimensional limit, all that randomness averages out, and the system follows a smooth, predictable path.
Tom: And that path depends on just a few numbers — what they call the macroscopic variables. Essentially, the overlaps between the student network's weights and the teacher network's weights.
Jane: Right. You don't need to track every individual weight. You just need to track how aligned the student is with the teacher, and how spread out the student's weights are. That's enough to compute the population risk — the actual error on new data.
Tom: And the key result is that whether that error goes to zero depends on where you sit in their phase diagram. The exponents matter — how the learning rate scales with dimension, and how the hidden layer width scales with dimension.
Jane: Let me give you the concrete numbers. They write the learning rate as gamma proportional to d to the minus delta, and the hidden layer width as p proportional to d to the kappa. Then everything depends on the sum kappa plus delta.
Tom: And that sum determines which of the four regions you're in. If it's positive, you get perfect learning. If it's exactly zero, you get a plateau — the error bottoms out at some nonzero value proportional to the noise. If it's between negative one-half and zero, you get bad learning.
Jane: And if it's less than negative one-half, their theory doesn't even apply. The stochastic process doesn't converge to those nice deterministic equations. That's the "no ODEs" region on their diagram.
Tom: So it's not just about whether learning works — it's about whether the math you're using to describe it is even valid.
Jane: Exactly. And that's a really important contribution, because a lot of previous work just assumed the ODE description held. This paper shows exactly when it does and when it doesn't.
Tom: So we've got the big picture. But how did they actually prove this? That's what I want to dig into next.
Improvements: Tom: So Jane, one thing I really appreciate about this paper is that they didn't just take the old proof and tweak it. They actually fixed some holes in the previous work. What did you see there?
Jane: Right, so the original proof by Goldt and collaborators from two thousand nineteen had some gaps. The authors here — Veiga, Stephan, Loureiro, Krzakala, and Zdeborová — they tightened things up considerably. They got much finer non-asymptotic guarantees, which means their bounds are sharper for finite system sizes, not just in the infinite limit.
Tom: And they also extended the result to handle arbitrary time scalings. That's crucial, because depending on where you are in the phase diagram, the natural time scale is different. Sometimes you need to rescale time by one over d, sometimes by one over d to some other power.
Jane: Exactly. And they provide convergence rates that scale like the square root of the time step times a log factor. That's a really clean result. It tells you exactly how fast the discrete stochastic process approaches the continuous ODE as the dimension grows.
Tom: There's also a nice lemma in the appendix about perturbing ODEs — if you have a small perturbation to the dynamics, the solution stays close to the unperturbed one. That's what lets them say the noise term in the update equations can be neglected in the perfect learning region.
Jane: Right, and that's a really elegant piece of math. The noise term scales like one over d to the kappa plus delta, so in the green region of the phase diagram, it just vanishes. That's why perfect learning is possible there.
Tom: But in the blue region — the plateau line — that noise term survives. It's exactly balanced with the learning term, and that's what creates the nonzero asymptotic error.
Jane: And that connects back to the classical Saad and Solla result. Their setting, with fixed learning rate and fixed hidden layer width, sits exactly on that plateau line. So the old results are a special case of this much bigger picture.
Tom: So they've really built a more complete theory. But I'm wondering — how does this connect to the mean-field approach that's been so popular recently?
First Page: Tom: Jane, let's look at the first page of the paper, because there's a really important connection there. The authors explicitly position their work relative to the mean-field limit of neural networks. Can you unpack that for our listeners?
Jane: Sure. The mean-field approach, which came from people like Mei, Montanari, and Chizat, looks at what happens when the hidden layer width goes to infinity. You get a partial differential equation describing the evolution of the weight distribution. It's a beautiful theory, and it guarantees global convergence.
Tom: But this paper is looking at a different limit — the input dimension going to infinity, with finite hidden layer width. And the authors show that these two approaches are really describing different regions of the same phase diagram.
Jane: Right. The mean-field limit corresponds to the green region, where you get perfect learning. But the classical Saad and Solla approach, with finite width, sits on the blue line. And the paper shows that the connection between them is not trivial — you can't just take one limit and get the other.
Tom: There's also a really interesting discussion about initialization on that first page. The authors point out that if you initialize the student weights uncorrelated with the teacher, the ODEs have a fixed point where the student never specializes. All the hidden units just learn the same function.
Jane: That's the unspecialized regime. And it's a real phenomenon — you see it in simulations. The network first learns a linear approximation, and only later do the hidden units start to specialize to different teacher components.
Tom: And the authors note that this is different from the "lazy training" regime that's been popular in the NTK literature. In lazy training, the weights barely move. Here, the weights are changing a lot — they're just all learning the same thing.
Jane: Right, that's a subtle but important distinction. And they also flag an open problem: what happens with random initialization, where the initial overlap with the teacher is vanishingly small? That requires additional log factors in the time scale, and it's still not fully understood for two-layer networks.
Tom: So there's real work left to do. But the framework they've built is solid. Let's bring in the rest of the team to talk about what this means in practice.
Lu: I want to jump in here. The connection between the phase diagram and the mean-field limit is actually really deep. The authors show that the mean-field approach requires both large width and a specific scaling of the learning rate. If you're in the wrong region, the mean-field description just doesn't apply.
Meng: And from a practical standpoint, that's a warning. A lot of people train wide networks with large learning rates and assume everything is fine. This paper says: check your scaling. If you're in the bad learning region, you might be wasting compute.
Lalam: I see a cultural implication as well. This work gives us a principled way to think about when neural network training is reliable. For applications where we need guarantees — medical imaging, autonomous systems — knowing which region of the phase diagram you're in could inform whether you can trust the trained model.
Tom: That's a great point, Lalam. And it's exactly why this kind of theory matters beyond just academic curiosity.
Conclusion: Tom: Alright, let's wrap up our discussion of "Phase diagram of Stochastic Gradient Descent in high-dimensional two-layer neural networks." Jane, what's the one thing you want our listeners to remember?
Jane: I think it's that the behavior of SGD is not universal. It depends critically on how you scale the learning rate and the hidden layer width relative to the input dimension. And this paper gives you the complete map — four regions, each with its own phenomenology.
Tom: Perfect learning, plateau, bad learning, and the region where the theory breaks down entirely. That's a really clean way to organize our understanding.
Lu: And the mathematical contributions are substantial. They fixed gaps in previous proofs, provided sharp convergence rates, and connected two major research traditions that had been developing in parallel.
Meng: From my side, the practical takeaway is clear: if you're training two-layer networks in high dimensions, you should know where you are on this diagram. It tells you whether your training will converge to a good solution, and how many samples you'll need.
Lalam: And I'd add that this kind of rigorous understanding is what builds trust in AI systems. When we can predict exactly when learning works and when it doesn't, we can deploy these models more responsibly.
Tom: Beautifully said. So we've got theory, practice, and societal implications all covered. That's a full episode.
Jane: It really is. And we should mention — the authors have made their code available, so you can reproduce their simulations and explore the phase diagram yourself.
Tom: That's a great note to end on. Thanks to everyone who joined us today. We'll be back with another paper soon.
Jane: Until then, keep learning. And maybe check where you are on the phase diagram.
Rodrigo Veiga, Ludovic Stephan, Bruno Loureiro, Florent Krzakala, Lenka Zdeborová
Ecole Polytechnique Fédérale de Lausanne · Universidade de São Paulo
stat.ML, cond-mat.dis-nn, cs.LG
Submitted: 2023-06-14
Updated: 2026-08-10
Comments: 20 pages
Journal ref: Advances in Neural Information Processing Systems (2022), vol 35, pages {23244--23255)
DOI: 10.52202/068431-1689
Code: https://github.com/rodsveiga/phdiag_sgd
License: http://creativecommons.org/licenses/by/4.0/
Importance score: 71/100
Key concepts
- Phase Diagram
- A graph used to map all possible behaviors when training a neural network. It shows distinct regions—like perfect learning, bad learning, or plateaus—based on how the learning rate and hidden unit width are scaled relative to the input dimension.
- Stochastic Gradient Descent (SGD)
- The method used to train neural networks. The authors show that in high dimensions, the random updates from individual data samples average out, allowing the system's dynamics to be described by smooth, deterministic equations.
- High-Dimensional Limit
- The theoretical scenario where the input dimension (d) is allowed to grow infinitely large. This limit allows researchers to simplify complex random processes into predictable mathematical equations.
- Mean-Field Approach
- A theoretical method that analyzes what happens when the hidden layer width goes to infinity. The paper compares this approach to its own findings, showing they describe different, but related, regions of the overall phase diagram.
Terminology
Summary
Summary
This paper investigates the high-dimensional dynamics of stochastic gradient descent (SGD) for training two-layer neural networks in a teacher-student setup, aiming to bridge two previously distinct lines of research: the mean-field/hydrodynamic limit (which describes wide networks via a partial differential equation in the weight space) and the seminal approach of Saad & Solla (which describes finite-width networks via a set of ordinary differential equations for macroscopic overlaps).
The authors consider a regression task where the teacher network has k hidden units with fixed weights* in R k times d and the student network has p hidden units with weights in R p times d. The data is Gaussian, P = N(0,), and the labels are generated by the teacher with additive noise: y nu = f(nu,) + sqrt zeta nu, where 0 controls the noise strength. The student is trained using one-pass SGD, updating weights sequentially with a learning rate gamma. The population risk is minimized, and the dynamics are tracked through sufficient statistics called macroscopic variables: the student-student overlap matrix nu = nu nu / d, the student-teacher overlap matrix nu = nu / d, and the fixed teacher-teacher overlap matrix =** / d. These are combined into the overlap matrix nu in R(p+k) times (p+k).
The main contribution is a rigorous derivation of a phase diagram for the learning dynamics. The authors scale the learning rate and hidden layer width with the input dimension d as gamma proportional to d-delta and p proportional to d kappa. They prove that the discrete stochastic process of SGD converges to a set of deterministic ODEs, d over dt (t) = psi((t)), provided a specific time scaling condition is met. This is formalized in Theorem 3.1, which states that for a time scaling delta t satisfying delta t c (gamma over pd, gamma squared over p squared d), and under Lipschitz conditions on the activation function and the function psi, the error between the discrete process and the ODE solution is bounded by E nu - (nu delta t) infinity e C tau (p) sqrt delta t. This result extends the proof of [6] by providing finer non-asymptotic guarantees and accommodating general time scalings.
Based on this analysis, the authors identify four distinct learning regimes in the phase diagram (Figure 1a):
-
Perfect learning region (green, kappa > -delta): The noise term in the ODEs vanishes as d-(kappa+ delta), and the dynamics converge to a noiseless set of ODEs. Perfect learning (zero population risk) can be asymptotically achieved with n about d 1+ kappa+ delta samples, even for tasks with additive noise. The time scaling is delta t kappa+ delta = 1/d 1+ kappa+ delta.
-
Plateau line (blue, kappa = -delta): This is an extension of the classical Saad & Solla setting (kappa = delta = 0). The noise term does not vanish, leading to an asymptotic plateau in the population risk proportional to the noise level and the learning rate gamma. The time scaling is delta t 0 = 1/d.
-
Bad learning region (orange, -1/2 < kappa + delta < 0): The noise term dominates the dynamics. The learning term is attenuated, and the teacher-student overlap matrix remains fixed at its initial value, leading to poor generalization. The time scaling is delta t 2(kappa+ delta) = 1/d 1+2(kappa+ delta).
-
No ODEs region (red, kappa + delta < -1/2): The stochastic process does not converge to deterministic ODEs under the assumptions of Theorem 3.1, and the authors make no claims about this regime.
The paper provides explicit analytical expressions for the ODE terms and the population risk for the activation function sigma(x) = erf(x/sqrt 2), which are detailed in Appendix C. The authors also discuss the role of initialization, noting that an unspecialized initial condition (where all student neurons have the same overlap with the teacher) is a fixed point of the ODEs, preventing specialization. They also highlight the trade-off between the learning rate and hidden layer width, noting that being closer to the plateau line reduces the number of samples needed but sacrifices the asymptotic performance.
The theoretical findings are validated through numerical simulations. For the Saad & Solla scaling (kappa = delta = 0), the simulations match the ODE predictions, showing a plateau related to noise. For the perfect learning region (kappa = 0, delta > 0), simulations show that the asymptotic population risk decays as R infinity proportional to d-delta, confirming perfect learning in the high-dimensional limit. For the bad learning region (kappa = 0, delta < 0), simulations show poor performance and strong finite-size effects. Finally, for large hidden layers (kappa > 0), simulations across different regions of the phase diagram confirm the predicted behaviors, with green curves decreasing towards perfect learning and orange/red curves getting stuck at higher plateaus.
Improvements for AI systems
Based on the paper, here are specific improvements that can be made to AI systems, along with what the improved system can do:
-
Improvement: Implement a dynamic scheduler that adjusts the learning rate (γ) and hidden layer width (p) based on the input dimension (d) to navigate the phase diagram.
-
What the improved system can do: Automatically select the optimal scaling exponents (κ, δ) to ensure the system operates in the
perfect learning
region (κ + δ > 0) while minimizing sample complexity. For example, choosing κ + δ = ε (small positive) achieves zero population risk with only n d(1+ε) samples, avoiding the need for excessively large datasets. -
Improvement: Use the theoretical prediction that asymptotic population risk scales as R∞ ∝ γΔ in the plateau region (κ = -δ) to set early-stopping criteria.
-
What the improved system can do: Given a known noise level Δ and learning rate γ, predict the achievable minimum risk and stop training once the empirical risk reaches this theoretical floor, saving computational resources without sacrificing performance.
-
Improvement: Monitor the teacher-student overlap matrix (M) to detect when the system transitions from the
unspecialized plateau
to thespecialization phase.
-
What the improved system can do: Dynamically adjust the learning rate or switch to a different optimization strategy once specialization begins, accelerating convergence. The system can also detect when it's stuck in the unspecialized fixed point (where all overlaps are equal) and reinitialize with correlated weights to escape this trap.
-
Improvement: Use the derived relationship n = τ·d(1+κ+δ) to optimally allocate data collection efforts.
-
What the improved system can do: For a given computational budget, determine the minimum number of samples needed to achieve a target risk level. The system can also trade off between increasing hidden layer width (κ) versus increasing learning rate (δ) to achieve the same performance with fewer samples.
-
Improvement: Leverage the proven convergence rate EΩ ν - Ω̄(νδt)∞ ≤ e(Cτ)log(p)√δt to correct for finite-dimensional effects.
-
What the improved system can do: Given training dynamics at a moderate dimension (e.g., d=100), extrapolate to predict performance at much higher dimensions (d=10,000) with quantified confidence bounds, enabling reliable model selection without expensive high-dimensional experiments.
-
Improvement: In the
bad learning
region (-1/2 < κ + δ < 0), where noise dominates, modify the loss function to explicitly account for the variance term E jE l that causes the plateau. -
What the improved system can do: Implement a variance-reduced loss that cancels the noise contribution, effectively moving the system into the
perfect learning
regime even when the raw SGD dynamics would be noise-dominated. -
Improvement: Based on the finding that uncorrelated initialization leads to the unspecialized fixed point, implement a
spread initialization
with non-zero teacher overlap. -
What the improved system can do: Guarantee that training will not get stuck in the unspecialized phase, ensuring global convergence to the optimal solution. The system can also detect when initialization has insufficient teacher overlap and automatically resample.
-
Improvement: Use the deterministic ODE system (Eqs. 23-25) as a fast surrogate model for training dynamics.
-
What the improved system can do: Predict the entire learning curve (population risk vs. time) in milliseconds for a given architecture and data distribution, without running actual training. This enables rapid hyperparameter search and architecture selection before committing to expensive training runs.
-
Improvement: Implement the time-scaling factor δt = max(1/d(1+κ+δ), 1/d(1+2(κ+δ))) to ensure the ODE approximation remains valid throughout training.
-
What the improved system can do: Automatically adjust the effective learning rate during training to maintain the validity of the deterministic approximation, preventing divergence or chaotic behavior that occurs when the noise term dominates.
-
Improvement: Exploit the finding that the noise term vanishes as d(-(κ+δ)) in the perfect learning region to implement a two-phase training schedule.
-
What the improved system can do: In the first phase, use a larger learning rate to quickly escape the unspecialized plateau. In the second phase, reduce the learning rate to achieve the theoretical zero-risk limit, with the transition point determined by the phase diagram boundaries.
Abstract
Despite the non-convex optimization landscape, over-parametrized shallow networks are able to achieve global convergence under gradient descent. The picture can be radically different for narrow networks, which tend to get stuck in badly-generalizing local minima. Here we investigate the cross-over between these two regimes in the high-dimensional setting, and in particular investigate the connection between the so-called mean-field/hydrodynamic regime and the seminal approach of Saad & Solla. Focusing on the case of Gaussian data, we study the interplay between the learning rate, the time scale, and the number of hidden units in the high-dimensional dynamics of stochastic gradient descent (SGD). Our work builds on a deterministic description of SGD in high-dimensions from statistical physics, which we extend and for which we provide rigorous convergence rates.
Sources
Related papers
- Behavior of prediction performance metrics with rare events
- Optimal Estimation of Generic Dynamics by Path-Dependent Neural Jump ODEs
- A Posterior-Dynamics Framework for Imaging Inverse Problems with Pretrained Diffusion Priors
- One Permutation Is All You Need: Fast, Deterministic Feature Importance and Model Stress-Testing
- Online Conformal Prediction for Non-Exchangeable Panel Data
- Deep Time-Series Forecasting in 10 Years: A Survey