Understanding Edge-of-Stability Training Dynamics with a Minimalist Example
1Duke University2Tsinghua University
*Equal contribution.ICLR 2023.
Summary
- With a fixed step size \(\eta\), gradient descent on deep networks raises the sharpness, the largest eigenvalue of the loss Hessian, to the stability threshold \(2/\eta\) and keeps it near there while the loss keeps decreasing, though not monotonically. This is the edge of stability [1]. Two stronger regularities hold at the end of training. From one initialization, the final sharpness sits just below \(2/\eta\) for a range of step sizes (sharpness adaptivity), and at one step size it sits there for a wide range of initializations (sharpness concentration). For shallow models such as matrix factorization and two-layer networks, the gap to \(2/\eta\) is often much larger.
- We study \(L(x,y,z,w)=\tfrac12(1-xyzw)^2\), a four-layer network with one neuron per layer, started with \(z=x\) and \(w=y\). Gradient descent on it shows both regularities (Figure 1). We prove that for small \(\eta\), every initialization in a constant-size region converges to a global minimum with sharpness in \((2/\eta-20\eta/3,\ 2/\eta)\), which implies adaptivity over a range of step sizes. A rank-1 matrix version holds as well.
- The mechanism is visible in coordinates along and across the valley of minima. Every two steps, gradient descent follows the parabola \(b^2=a\kappa/2+\kappa^4/16\) with \(\kappa=\sqrt\eta\), and the tip of the parabola, at \(a\approx-\kappa^3/8\), fixes where training stops. Degree-2 objectives lack the cubic term that creates this parabola. Their two-step paths are ellipses, so the endpoint roughly mirrors the start and a sharper start ends flatter.
- Far from the minimum the dynamics are chaotic. The boundary between converging and diverging initializations is fractal, and the two-step map \(b\mapsto b(1+8a\kappa-16b^2)\) is a cubic map that period-doubles. A 5-layer ELU network trained on 50 CIFAR-10 images ends at sharpness 199.97 with \(2/\eta=200\), and near the end its trajectory, projected onto two directions, is well fit by a parabola.
Sharpness that tracks the step size
Write \(\lambda\) for the sharpness, the top eigenvalue of the Hessian. On an objective with a fixed Hessian, gradient descent with step size \(\eta\) is stable only when \(\lambda<2/\eta\), which is why \(2/\eta\) is called the stability threshold. Cohen et al. [1] observed that full-batch gradient descent on neural networks does not stay below it. The sharpness rises to about \(2/\eta\) and hovers there, and the loss still decreases in the long run, with local increases along the way.
Deeper networks also end training in a regular place, in two senses:
- Sharpness adaptivity: from a fixed initialization, changing \(\eta\) moves the final sharpness to just below the new \(2/\eta\).
- Sharpness concentration: at a fixed \(\eta\), a wide range of initializations all end at sharpness just below \(2/\eta\).
The paper's first figure contrasts a deep network with a shallow one. A ReLU network with five fully connected layers of 50 neurons, trained with \(\eta=2/200,\ 2/300,\ 2/400\), ends at sharpness close to \(2/\eta\) each time. A linear two-layer network with 10 neurons per layer, whose layers are rescaled at initialization by factors 4 and 0.1 (without rescaling it does not reach the threshold), ends clearly below \(2/\eta\) for \(\eta=2/30,\ 2/35,\ 2/40\). Depth seems to matter. We look for the smallest model in which the endpoint is pinned to \(2/\eta\), and for the mechanism that pins it.
A four-parameter model
The model is a product of four scalars, a network with four layers and one neuron each, fit to the target 1 with the squared loss:
We run gradient descent with a fixed step size \(\eta\) from \(z_0=x_0\), \(w_0=y_0\). The objective is unchanged when \(x\) is swapped with \(z\) and \(y\) with \(w\), so the coupling persists at every step, and the dynamics are gradient descent on \(L(x,y)=\tfrac14(1-x^2y^2)^2\):
The global minima form the hyperbola \(xy=1\), and we take \(x>y>0\). With \(S=x^2+y^2\) and \(\gamma=xy\), the Hessian at \((x,y,x,y)\) has the eigenvalue
At a minimum the other three eigenvalues vanish and \(\lambda_1=2(x^2+y^2)\). Sharpness therefore grows along the hyperbola away from the point \((1,1),\) and on the branch \(x>y\) exactly one minimum has sharpness \(2/\eta\). Away from the minima the four-parameter Hessian has other eigenvalues, and \(\lambda_1\) is not always the largest; the paper's experiments compute the top eigenvalue of the exact Hessian numerically. Near the minima the other eigenvalues vanish, so our figures plot \(\lambda_1\) throughout.
Definition 1η-EoS minimum
For \(\eta<\tfrac12\), the minimum with sharpness exactly \(2/\eta\), that is \(x^2+y^2=1/\eta\) and \(xy=1\), is
For \(2/\eta=8,\ 10,\ 12\) it is \((1.932,\,0.518),\) \((2.189,\,0.457)\) and \((2.414,\,0.414).\)
Figure 1 runs gradient descent from \((2.5,\,0.4001),\) where \(\lambda_1=12.83\) is above all three thresholds \(2/\eta=8,\ 10,\ 12\). In each run the sharpness spends a while above its threshold (146, 181 and 236 steps, all within the first 450) and then settles just below it. After 1000 steps the simulation reads 7.925, 9.927 and 11.933, while the loss has fallen to about \(10^{-11}\) or less. In the plane, successive iterates jump back and forth across the valley, while every second iterate traces a smooth curve along it and stops next to the matching EoS minimum, slightly on its flat side (smaller \(x\)).
The same figure can switch to the degree-2 objective \(\tfrac12(1-xy)^2\). From \((3.2,\,0.25)\) with \(2/\eta=10,\ 9.5,\ 9\), it ends after 100 steps at 9.511, 8.472 and 7.310, below each threshold by 0.49 to 1.69 (simulation readouts).
Loss
Sharpness \(\lambda_1\)
(x, y) trajectory
One sharpness for many starts
Adaptivity fixes the start and varies \(\eta\). Concentration fixes \(\eta\) and varies the start. Figure 2 runs gradient descent at \(\eta=0.2\) from every cell of a grid of initializations, for up to 50,000 steps as in the paper's Figure 2b, and colors each start by the sharpness of the minimum it reaches.
Next to each of the two EoS minima, \((2.189,\,0.457)\) and its mirror image \((0.457,\,2.189),\) a large region of starts ends within 0.1 of \(2/\eta=10\). In our checks of the grid, this band holds about a quarter of all converging starts in \([0,4]^2\), and the status line under the figure gives the live count. Many of these starts are sharper than \(2/\eta\), so gradient descent first enters the edge of stability and then settles just below the threshold. The step-size menu repeats the experiment at \(\eta=0.3\) and \(0.4\) (paper Figure 10). In our grid the band then covers a larger share of the converging starts, about 32% and 45% against 24% at \(\eta=0.2\).
For the degree-2 objective at the same step size, the band shrinks to isolated cells (paper Figure 5b). Only starts very close to its EoS minimum end near 10 (paper Figure 12). The fractal-edge view zooms into \(x_0\in[3.0,\,3.4]\), \(y_0\in[0.10,\,0.50]\) (paper Figure 6a). Nearly every converging start there ends in the band, but converging and diverging starts interleave in ever finer strips. We return to this boundary below.
Two steps at a time
To see why the endpoint is pinned, we change coordinates. Let \(c=(x^2-y^2)^{1/2}\) and \(d=xy\). Level sets of \(d\) are hyperbolas parallel to the valley of minima, so changing \(d\) moves across the valley. Level sets of \(c\) are the orthogonal hyperbolas \(x^2-y^2=\text{const}\), so changing \(c\) moves along it. Arora et al. [2] use a similar split. The EoS minimum is \((c,d)=\big((\eta^{-2}-4)^{1/4},\,1\big)\), and we measure offsets from it:
Minima are the line \(b=0\). Positive \(a\) is sharper than the EoS minimum and negative \(a\) is flatter. For small \(\kappa\), \(a\) and \(b\), one step and two steps of gradient descent are, to leading order,
One step flips the sign of \(b\): the iterate jumps across the valley. Over two steps the flip cancels, and so do the even powers of \(b\). What remains moves \(a\) steadily toward flatter minima, and moves \(b\) by a balance of two terms. The term \(8ab\kappa\) pushes \(b\) away from zero while \(a>0\), and the cubic term \(-16b^3\) pulls it back.
Treating the two-step change as a derivative gives an ODE with explicit solutions:
Along a trajectory \(a\) decreases, so the last term decays exponentially on the scale \(\kappa^3\), and every path that starts at a positive, not too small \(a\) joins the parabola \(b^2=a\kappa/2+\kappa^4/16\). The parabola meets \(b=0\) at \(a=-\kappa^3/8\), just on the flat side of the EoS minimum. The start only decides where the path joins the parabola, and the endpoint is the tip. Near the EoS minimum, the sharpness of a minimum changes by about \(4a/\kappa\), so \(a=-\kappa^3/8\) corresponds to a gap of about \(\eta/2\) below \(2/\eta\). This last step is our own calculation, not a statement in the paper.
Figure 3 runs the actual gradient descent in these coordinates, from \(b_0=0.01\) and a chosen \(a_0\). In the degree-4 panel every second iterate joins the parabola and follows it to its tip, and the endpoint barely moves when \(a_0\) does. At \(\eta=0.1\) the runs end at \(a\) between \(-0.00373\) and \(-0.00376\) for \(a_0\) from 0.02 to 0.3, against \(-\kappa^3/8=-0.00395\). At \(\eta=0.05\) they end at \(a=-0.00138\) for every \(a_0\) from 0.05 to 0.8, with sharpness 39.9755 against \(2/\eta=40\). For \(a_0=0.1\), the ratio of the final \(a\) to \(-\kappa^3/8\) approaches 1 as \(\eta\) shrinks: 0.83 at \(\eta=0.2\), 0.95 at 0.1 and 0.99 at 0.05, since \(-\kappa^3/8\) is the small-step limit. These are simulation readouts. The degree-2 panel is the subject of the next section.
Degree 4: ¼(1 − x²y²)²
Degree 2: ½(1 − xy)²
The theorem makes this picture rigorous in a region of initializations that is large compared with the precision of the endpoint.
Theorem 1sharpness concentration, paper Theorem 3.1 and Corollary 3.1
Many starts, one final sharpness.
Let \(K\) be a large enough absolute constant and \(\eta<K^{-2}/8{,}000{,}000\). From every \((x_0,y_0)\) with \(x_0\in\big(\breve x+13\eta^{5/4},\ \breve x+K^{-2}\eta^{-1/2}/5\big)\) and \(0<|x_0y_0-1|<K^{-1}\), gradient descent with step size \(\eta\) converges to a global minimum with sharpness in \((2/\eta-20\eta/3,\ 2/\eta)\).
In \((a,b)\) coordinates: if \(\kappa<K^{-1}/(2000\sqrt2)\), \(a_0\in(12\kappa^{5/2},\ K^{-2}\kappa^{-1}/4)\) and \(b_0\in(-K^{-1},K^{-1})\setminus\{0\}\), then for every \(\varepsilon>0\) there is a time \(T\) such that \(|b_t|<\varepsilon\) and \(a_t\in(-5\kappa^3/3,\ -\kappa^3/10)\) for all \(t>T\), with \(T=O(K^{-2}\kappa^{-15/2})\) plus \(O(\log(1/\varepsilon))\) plus \(O(\log(1/|b_0|)\,\kappa^{-7/2})\).
Corollary 2sharpness adaptivity, paper Corollary 3.2
One region of starts, many step sizes.
Fix \(\alpha<K^{-1}/(2000\sqrt2)\). Every start with \(x_0\in\big(\alpha^{-1}+K^{-2}\alpha^{-1}/15,\ \alpha^{-1}+K^{-2}\alpha^{-1}/6\big)\) and \(0<|x_0y_0-1|<K^{-1}\) converges, for every step size \(\eta\in(\alpha^2-K^{-2}\alpha^2/10,\ \alpha^2)\), to a minimum with sharpness in \((2/\eta-20\eta/3,\ 2/\eta)\).
In the \((x,y)\) plane the region of Theorem 1 contains a box of width \(\Theta(K^{-2}\eta^{-1/2})\) and height \(\Theta(K^{-1}\eta^{1/2})\), so many starts are far from the EoS minimum, while the final sharpness is fixed to within \(20\eta/3\). The proof follows the two-step iterates through six regions of the \((a,|b|)\) plane:
- Region I, above \(b^2=2a\kappa\): the cubic term dominates, and \(|b|\) decreases until the iterate enters region II.
- Region III, below \(b^2=a\kappa/4\): the term \(ab\kappa\) dominates, and \(|b|\) grows by a factor of at least \(1+a\kappa\) every two steps until it enters region II.
- Region II, between those curves, keeps the iterate while \(a\) decreases into region IV. There the residual \(\xi=b^2-a\kappa/2-\kappa^4/16\) contracts, roughly as \(\xi''\approx(1-32b^2)\,\xi\), to \(|\xi|<\kappa^4/200\), and it stays that small in region V as the iterate follows the parabola.
- In region VI, near the tip, \(a\) is negative, so \(|b''|<(1+a\kappa)\,|b|\) with \(1+a\kappa<1\): \(|b|\) shrinks geometrically while \(a\) barely moves.
The same result holds for vectors. For \(\min_{x,y\in\mathbb R^d}\tfrac14\|I-xy^\top xy^\top\|_F^2\), from random initializations on spheres of suitable radii and with one multiplicative perturbation \(y\leftarrow y(1+2K^{-1})\) during training, for every \(\varepsilon>0\) gradient descent eventually keeps the loss below \(\varepsilon\) with \(\|x\|^2+\|y\|^2\in(1/\eta-10\eta/3,\ 1/\eta)\), with probability at least \(1-2\delta_0-2e^{-\Omega(d)}\) (paper Theorem 3.2). The vectors first align geometrically, and then their norms follow the scalar dynamics. The perturbation keeps the run from converging to an unstable point with sharpness above \(2/\eta\).
Why degree two behaves differently
Prior analyses of gradient descent beyond the threshold on \((\mu-xy)^2\), \((\mu-x^\top y)^2\) and \(\|\mu I-xy^\top\|_F^2\) [3, 4] prove convergence to a minimum with sharpness at most \(2/\eta\). Empirically, these runs overshoot. They end at minima clearly flatter than the EoS minimum, and the sharpness does not oscillate around the threshold (paper Figure 15, with \(\eta=0.05\) on \((1-x^\top y)^2\) and \(\|I-xy^\top\|_F^2\)).
For \(\tfrac12(1-xy)^2\), the same coordinates, centered at its own EoS minimum where \(x^2+y^2=2/\eta\), give the two-step map
There is no cubic term to hold \(b\) near a parabola. The ODE \(db/da=-4a/(b\kappa^2)\) has the solutions \(b^2=4(C-a^2)/\kappa^2\), ellipses centered at the EoS minimum. A start at \(a_0>0\) with small \(b_0\) travels along half an ellipse and ends near \(-a_0\), so a sharper start ends at a flatter minimum. The degree-2 panel of Figure 3 shows this. At \(\eta=0.1\), starts with \(a_0=0.02,\ 0.1,\ 0.3\) end at \(a=-0.0204,\ -0.109,\ -0.393\), roughly the mirror image of the start. In the \((x,y)\) plane with \(\eta=0.2\), the starts \((3.2,\,0.25),\) \((3.35,\,0.25)\) and \((3.5,\,0.25)\) of the paper's Figure 11 have sharpness 10.28, 11.26 and 12.30 and end at 9.51, 8.39 and 6.82 (simulation readouts).
The paper's Section 4 writes the objective as \((1-xy)^2\), but its two-step map and experiments correspond to \(\tfrac12(1-xy)^2\), whose sharpness at a minimum is \(x^2+y^2\). Its ODE is also printed as \(db/da=4a/(b\kappa^2)\); the ellipses it reports require the minus sign used here. We use \(\tfrac12(1-xy)^2\) throughout.
A degree-3 model separates the two behaviors. For \(\tfrac12(1-xyz)^2\) with \(z=y\), there are two EoS minima in the positive quadrant. Around the one where the small entry is duplicated (small \(y=z\)), starts concentrate just below \(2/\eta\), as in degree 4. Around the one where the only small entry is \(x\), they do not, as in degree 2 (paper Appendix A.2.3). What matters appears to be how many times the small entry is repeated, rather than the total degree.
Far from the minimum: bifurcation
The theorem is local, and the global picture is complicated. The boundary between converging and diverging starts is fractal (the fractal-edge view of Figure 2). From the asymmetric start \((12.5,\,0.05)\) with \(\eta=0.01\), close to the boundary of divergence, gradient descent first oscillates chaotically, then settles into a clean alternation across the valley, and then follows the parabola to the minimum near \((10,\,0.1)\) (paper Figure 6b). Run to convergence in our simulation, it ends at \((9.99938,\,0.100006),\) with sharpness 199.995 and \(a=-1.249\times10^{-4}\approx-\kappa^3/8\).
The two-step map explains the order of these stages. Its update of \(b\) factors as
Since \(a\) changes slowly, we can hold it fixed and read this as a one-dimensional cubic map with parameter \(a\), much like the logistic map. The regimes below are our own derivation from this map, not statements in the paper. With \(s=8a\kappa\) and \(u=4b\), the map is \(u\mapsto u(1+s-u^2)\):
- For \(-2<s<0\), that is for \(-1/(4\kappa)<a<0\), the fixed point \(b=0\) is stable, and the run converges to a minimum.
- For \(0<s<1\), the orbit settles at \(b^2=a\kappa/2\), the leading-order parabola of the previous section. Its fixed points have multiplier \(1-2s\). In gradient descent this is a steady period-two oscillation across the valley.
- At \(s=1\), that is \(a=1/(8\kappa)\) (1.25 for \(\kappa=0.1\)), the fixed points period-double, and further doublings and chaotic bands follow as \(s\) grows.
- For \(s>2\) (\(a>2.5\) for \(\kappa=0.1\)) orbits escape, which corresponds to divergence.
Figure 4 overlays the gradient descent run on the bifurcation diagram of this map. The run starts at \(a\approx2.50\), at the edge of the escape regime, and \(a\) only decreases, so it crosses the diagram from right to left. It passes through period doubling in reverse, which the paper calls de-bifurcation, joins the parabola and converges once \(a<0\). The two sides of the oscillation differ, because the one-step update of \(b\) has a \(-3b^2\) term. Close to the minimum they bracket the parabola, with \(|b|=0.078\) and 0.063 at \(a=0.1\) against 0.071 from the parabola. Further right the match is looser, 0.276 and 0.141 at \(a=1.0\) against 0.224. These are readouts from our simulation. Ruiz-Garcia et al. [5] observed similar oscillations in neural networks when the learning rate is raised, and attributed them to cascades across several large eigendirections. This model has only one oscillating direction, so such a cascade is not needed to produce them.
Two-step map \(b \mapsto b\,(1 + 8a\kappa - 16b^2)\)
κ = 0.1 (η = 0.01), 500 values of a
Orbit at this a
bk from b0 = 0.05, k = 0 to 80
Beyond the toy model
Coupling emerges on its own
The coupled start \(z_0=x_0\), \(w_0=y_0\) may look artificial. In scalar networks \(\tfrac12(1-x_1\cdots x_n)^2\) whose entries all start different, the small entries often become equal by themselves (paper Appendix A.4), and Figure 5 reruns the paper's three examples at \(\eta=0.2\). From \((6,\,0.1,\,0.4),\) the gap \(|x_2-x_3|\) falls from 0.3 to \(3.2\times10^{-7}\) within 50 steps. From the seven-layer start \((2,\,2.5,\,3,\,0.1,\,0.2,\,0.3,\,0.4),\) the four small entries agree to within \(4\times10^{-5}\) by step 200, which gives a model like ours with a small entry repeated four times. The sharpness peaks at 36.8 and ends at 9.987. From \((6,\,0.7,\,0.3,\,0.2)\) the small entries never become equal (\(|x_2-x_4|=0.107\) at step 200), yet the sharpness still ends at 9.986, just below \(2/\eta=10\). These are simulation readouts. There is no sharp line between large and small entries.
Loss
Sharpness
Distance to the last entry
A five-layer network on CIFAR-10
We train a fully connected network with five layers, width 200 and no biases, with ELU or tanh activations, by full-batch gradient descent on the mean squared error over 50 CIFAR-10 images [6]: 25 airplanes labeled \(-1\) and 25 automobiles labeled \(+1\). With \(\eta=0.01\) the loss goes to zero, and for ELU after 18,500 iterations the sharpness converges to 199.97, with \(2/\eta=200\).
To compare with the scalar model, we project the trajectory onto two directions: the top Hessian eigenvector at the end of training (the oscillation direction), and the direction in which the two-step average moved between iteration 5000 and the end, orthogonalized against the first (the movement direction). After about 3000 iterations, the part of the trajectory outside this plane is small. The projected trajectory first oscillates in a bifurcation-like way and then moves along a smooth curve. From iteration 5000 on it is well fit by the parabola \(x=7500y^2\) for ELU and \(x=14000y^2\) for tanh, with \(x\) the movement coordinate and \(y\) the oscillation coordinate.
Noise and minibatches
With label noise (\(\sigma=0.01\), \(\eta=0.2\), start \((2.5,\,0.41)\)), the scalar model first follows the parabola to near the EoS minimum and then, over \(10^6\) iterations, drifts along the minima to the flattest one, \((1,1),\) with sharpness 4. With noise added to the gradient it instead wanders along the minima between the two EoS minima. Minibatch SGD on the networks above (\(\eta=0.005,\ 0.01,\ 0.02\), batch sizes 1 to 50, 10 initializations each) ends close to \(2/\eta\) for large batches and clearly below it for small ones, with the final sharpness concentrated for each batch size. The scalar model fits a single data point, so it does not explain this.
Related work
The phenomenon. Cohen et al. [1] formalized the edge of stability and showed that the loss can decrease non-monotonically while \(\lambda>2/\eta\). Non-monotone loss has also been reported in other settings [4, 7, 8, 9, 10, 11]. Lewkowycz et al. [9] described a catapult phase in which the loss does not diverge although the sharpness exceeds \(2/\eta\).
Mechanisms. Ahn et al. [12] study this unstable convergence and discuss its possible causes. Ma et al. [13] prove the edge of stability for losses with a subquadratic property; their model shows the phenomenon but not sharpness adaptivity. Arora et al. [2] and Lyu et al. [14] analyze sharpness reduction near the manifold of minima, with the loss \(\sqrt L\) or normalized gradient descent, or with a scale-invariant objective. There the effective step size changes during training, so adaptivity and concentration do not apply.
Degree two. Wang et al. [4] prove convergence of matrix factorization with step sizes beyond \(2/\lambda\), Chen and Bruna [3] study convergence beyond the edge of stability, and Li et al. [15] analyze sharpness along the trajectory of two-layer linear networks. These degree-2 settings do not show adaptivity or concentration.
Noise. The slow drift of the scalar model toward \((1,1)\) under label noise falls in the regime of sharpness reduction near the manifold of minima, studied by Damian et al. [16] and in [14, 15, 17].
Limitations
The rigorous result is local. The region of starts and the admissible step sizes depend on an unspecified absolute constant \(K\) (for example \(\eta<K^{-2}/8{,}000{,}000\)), and the global dynamics, with their fractal boundary, remain unanalyzed. The model assumes coupled entries, and coupling from uncoupled starts is shown only empirically and is not always complete. The vector result needs an explicit perturbation to avoid an unstable point. The link to real networks is empirical: a two-dimensional projection and a fitted parabola, on a 50-image, two-class regression. The scalar model fits one data point, so it does not explain why minibatch SGD ends below \(2/\eta\).
References
- Cohen et al. Gradient Descent on Neural Networks Typically Occurs at the Edge of Stability. arXiv preprint arXiv:2103.00065, 2021.
- Arora, Li and Panigrahi. Understanding Gradient Descent on Edge of Stability in Deep Learning. arXiv preprint arXiv:2205.09745, 2022.
- Chen and Bruna. On Gradient Descent Convergence beyond the Edge of Stability. arXiv preprint arXiv:2206.04172, 2022.
- Wang et al. Large Learning Rate Tames Homogeneity: Convergence and Balancing Effect. International Conference on Learning Representations, 2022.
- Ruiz-Garcia et al. Tilting the Playing Field: Dynamical Loss Functions for Machine Learning. International Conference on Machine Learning, 2021.
- Krizhevsky. Learning Multiple Layers of Features from Tiny Images. Technical report, 2009.
- Jastrzebski et al. The Break-Even Point on Optimization Trajectories of Deep Neural Networks. arXiv preprint arXiv:2002.09572, 2020.
- Xing et al. A Walk with SGD. arXiv preprint arXiv:1802.08770, 2018.
- Lewkowycz et al. The Large Learning Rate Phase of Deep Learning: the Catapult Mechanism. arXiv preprint arXiv:2003.02218, 2020.
- Arora, Li and Lyu. Theoretical Analysis of Auto Rate-Tuning by Batch Normalization. arXiv preprint arXiv:1812.03981, 2018.
- Li et al. Robust Training of Neural Networks Using Scale Invariant Architectures. International Conference on Machine Learning, 2022.
- Ahn, Zhang and Sra. Understanding the Unstable Convergence of Gradient Descent. arXiv preprint arXiv:2204.01050, 2022.
- Ma, Wu and Ying. The Multiscale Structure of Neural Network Loss Functions: The Effect on Optimization and Origin. arXiv preprint arXiv:2204.11326, 2022.
- Lyu, Li and Arora. Understanding the Generalization Benefit of Normalization Layers: Sharpness Reduction. arXiv preprint arXiv:2206.07085, 2022.
- Li, Wang and Li. Analyzing Sharpness along GD Trajectory: Progressive Sharpening and Edge of Stability. arXiv preprint arXiv:2207.12678, 2022.
- Damian, Ma and Lee. Label Noise SGD Provably Prefers Flat Global Minimizers. Advances in Neural Information Processing Systems, 2021.
- Li, Wang and Arora. What Happens after SGD Reaches Zero Loss? A Mathematical Framework. arXiv preprint arXiv:2110.06914, 2021.