Xingyu Zhu

Understanding Edge-of-Stability Training Dynamics with a Minimalist Example

1Duke University2Tsinghua University

*Equal contribution.ICLR 2023.

Summary

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:

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:

\[L(x,y,z,w)=\tfrac12\,(1-xyzw)^2 .\]

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\):

\(\displaystyle x_{t+1}=x_t-\eta\,x_ty_t^2\,(x_t^2y_t^2-1),\)\(\displaystyle y_{t+1}=y_t-\eta\,x_t^2y_t\,(x_t^2y_t^2-1).\)

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

\(\displaystyle \lambda_1=\tfrac12\Big[S\,(3\gamma^2-1)\)\(\displaystyle {}+\sqrt{S^2(1-3\gamma^2)^2+4\gamma^2(3-10\gamma^2+7\gamma^4)}\,\Big].\)

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

\(\displaystyle \breve x=\tfrac1{\sqrt2}\Big(\big(\eta^{-2}-4\big)^{1/2}+\eta^{-1}\Big)^{1/2},\)\(\displaystyle \breve y=\sqrt2\,\Big(\big(\eta^{-2}-4\big)^{1/2}+\eta^{-1}\Big)^{-1/2}.\)

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).

step size

Loss

Sharpness \(\lambda_1\)

(x, y) trajectory

 

Figure 1. Gradient descent on the 4-layer scalar network \(\tfrac12(1-xyzw)^2\) with coupled entries \((z=x,\ w=y)\), which is gradient descent on \(\tfrac14(1-x^2y^2)^2\), simulated live in the browser. The loss and sharpness panels (sharpness is \(\lambda_1\) of the paper's Eq. 3, which is the top Hessian eigenvalue near the minima) show runs at \(2/\eta = 8, 10, 12\) from the same start, 1000 steps each, with dashed lines at \(2/\eta\). The trajectory panel shows the iterates in the \((x, y)\) plane, even steps filled and odd steps hollow. The curve \(xy = 1\) is the set of global minima, its ticks give their sharpness, and diamonds mark the minimum whose sharpness is exactly \(2/\eta\). Click the plane (or focus it and use the arrow keys) to choose another start, or switch to the degree-2 objective \(\tfrac12(1-xy)^2\), 100 steps at \(2/\eta = 10, 9.5, 9\), to see the gap to \(2/\eta\) grow instead.

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.

view
ends within 0.1 of 2/η ends flatter (0 → 2/η) diverges minima xy = 1 minima with sharpness 2/η

Figure 2. Each cell of a 160 × 160 grid is a starting point \((x_0, y_0)\) for gradient descent on \(\tfrac14(1-x^2y^2)^2\) (degree 4) or \(\tfrac12(1-xy)^2\) (degree 2), run live in your browser for up to 50,000 steps and colored by the sharpness of the minimum it reaches. Blue cells end within 0.1 of \(2/\eta\) (the band used in the paper's Fig. 2b), grey cells end flatter (the stronger the grey, the closer to \(2/\eta\)), and red-tinted cells diverge; a run counts as converged once \(|x^2y^2-1|\) (or \(|xy-1|\)) falls below \(10^{-12}\) and as diverged once \(|x|\) or \(|y|\) exceeds \(10^4\). The dashed curve is the set of minima \(xy = 1\), and diamonds mark the two minima whose sharpness is exactly \(2/\eta\). Hover for a cell's outcome and click to draw its first 400 steps; the fractal-edge view zooms into the boundary between converging and diverging starts (the paper's Fig. 6a).

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:

\(\displaystyle a=(x^2-y^2)^{1/2}-(\eta^{-2}-4)^{1/4},\)\(\displaystyle b=xy-1,\qquad \kappa=\sqrt\eta .\)

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,

\(\displaystyle a'\approx a-2b^2\kappa^3,\)\(\displaystyle b'\approx -b-4ab\kappa-3b^2-b^3;\)\(\displaystyle a''\approx a-4b^2\kappa^3,\)\(\displaystyle b''\approx b+8ab\kappa-16b^3.\)

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:

\(\displaystyle \frac{db}{da}=\frac{16b^3-8ab\kappa}{4b^2\kappa^3},\)\(\displaystyle b^2=\tfrac12a\kappa+\tfrac1{16}\kappa^4+C\,e^{8a/\kappa^3}.\)

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.

view

Degree 4: ¼(1 − x²y²)²

 

Degree 2: ½(1 − xy)²

 

Figure 3. The same gradient descent in coordinates along the valley of minima (\(a\)) and across it \((b = xy - 1)\), with the minimum of sharpness exactly \(2/\eta\) at the origin (paper Definition 2). Both runs are simulated live from \(a = a_0\), \(b = 0.01\) until \(|b| < 10^{-10}\); filled dots are every second iterate and hollow dots the ones in between. In the degree-4 panel every start joins the parabola \(b^2 = a\kappa/2 + \kappa^4/16\) (\(\kappa = \sqrt\eta\), paper Eq. 9) and stops near its tip, which tends to \(a = -\kappa^3/8\) as \(\eta\) shrinks, just on the flat side of the threshold. In the degree-2 panel, for \(\tfrac12(1 - xy)^2\), the iterates follow an ellipse around the origin (dashed, the local approximation of paper Sec. 4) and stop near \(a = -a_0\), so a sharper start ends flatter. Move \(a_0\) to compare, and hover or use the arrow keys on a panel to read an iterate's step, position and sharpness.

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:

b² = 2aκ b² = aκ/2 + κ⁴/16 b² = aκ/4 I II III IV V VI 2κ5/2 κ5/2 −κ³/8 0 a (along the minima) → |b|
Schematic of the proof, after the paper's Figure 4: the \((a,|b|)\) plane with the EoS minimum at the origin (diamond). Distances are not to scale. Two-step iterates move left as \(a\) decreases. Starts above or below the band between \(b^2=2a\kappa\) and \(b^2=a\kappa/4\) are drawn into it (regions I–III), then into the thin tube V around the parabola \(b^2=a\kappa/2+\kappa^4/16\) (blue), and they stop at its tip, \(a\approx-\kappa^3/8\) (region VI). The vertical lines mark \(a=2\kappa^{5/2}\), \(a=\kappa^{5/2}\) and the start of region VI.

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

\(\displaystyle a''\approx a-\sqrt2\,b^2\kappa^3,\)\(\displaystyle b''\approx b+4\sqrt2\,ab\kappa .\)

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

\[b''=b\,(1+8a\kappa-16b^2).\]

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)\):

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

Figure 4. The two-step map panel shows the bifurcation diagram of the paper's approximate two-step map \(b \mapsto b(1+8a\kappa-16b^2)\) with the along-valley coordinate \(a\) held fixed (\(\kappa = 0.1\), \(\eta = 0.01\)); grey dots are iterates 601 to 700 from \(b_0 = 0.05\) for 500 values of \(a\), with their negatives, since gradient descent flips the sign of \(b\) every step. Blue is the actual gradient-descent run from \((12.5, 0.05)\) (the paper's Fig. 6b, simulated live for 150,000 steps): it starts in the chaotic band on the right, moves left as \(a\) decreases, passes through period doubling in reverse, settles on the branch \(b^2 = a\kappa/2\) (dashed) and approaches the minimum once \(a < 0\); its two sides differ because the one-step update of \(b\) has a \(-3b^2\) term. The orbit panel shows the map's orbit at the \(a\) set with the slider, by dragging on the diagram, or with play, which sweeps \(a\) from right to left; the first 20 iterates are grey. The diagram (as in the paper's Fig. 6c), the dashed branch and the regime thresholds come from the approximate map, not from gradient descent; the branch and the thresholds are our derivation rather than statements in the paper.

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

 

Figure 5. Gradient descent (\(\eta=0.2\)) on \(n\)-layer scalar networks \(\tfrac12(1-x_1\cdots x_n)^2\) with all entries initialized differently, from the three initializations of the paper's Appendix A.4 (its Figures 16, 17 and 20), simulated live for 5000 steps. The distance \(|x_i-x_n|\) of each entry to the last one, over the first 200 steps, shows the small entries (blue) becoming equal geometrically, which recovers the coupled model. The sharpness, the top eigenvalue of the \(n\times n\) Hessian, rises above \(2/\eta=10\) and then settles just below it, even in the 4-layer case where the small entries never become exactly equal, while the loss decreases non-monotonically. Hover a panel, or focus it and use the arrow keys, to read the step, loss, sharpness and every 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.

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

  1. Cohen et al. Gradient Descent on Neural Networks Typically Occurs at the Edge of Stability. arXiv preprint arXiv:2103.00065, 2021.
  2. Arora, Li and Panigrahi. Understanding Gradient Descent on Edge of Stability in Deep Learning. arXiv preprint arXiv:2205.09745, 2022.
  3. Chen and Bruna. On Gradient Descent Convergence beyond the Edge of Stability. arXiv preprint arXiv:2206.04172, 2022.
  4. Wang et al. Large Learning Rate Tames Homogeneity: Convergence and Balancing Effect. International Conference on Learning Representations, 2022.
  5. Ruiz-Garcia et al. Tilting the Playing Field: Dynamical Loss Functions for Machine Learning. International Conference on Machine Learning, 2021.
  6. Krizhevsky. Learning Multiple Layers of Features from Tiny Images. Technical report, 2009.
  7. Jastrzebski et al. The Break-Even Point on Optimization Trajectories of Deep Neural Networks. arXiv preprint arXiv:2002.09572, 2020.
  8. Xing et al. A Walk with SGD. arXiv preprint arXiv:1802.08770, 2018.
  9. Lewkowycz et al. The Large Learning Rate Phase of Deep Learning: the Catapult Mechanism. arXiv preprint arXiv:2003.02218, 2020.
  10. Arora, Li and Lyu. Theoretical Analysis of Auto Rate-Tuning by Batch Normalization. arXiv preprint arXiv:1812.03981, 2018.
  11. Li et al. Robust Training of Neural Networks Using Scale Invariant Architectures. International Conference on Machine Learning, 2022.
  12. Ahn, Zhang and Sra. Understanding the Unstable Convergence of Gradient Descent. arXiv preprint arXiv:2204.01050, 2022.
  13. Ma, Wu and Ying. The Multiscale Structure of Neural Network Loss Functions: The Effect on Optimization and Origin. arXiv preprint arXiv:2204.11326, 2022.
  14. Lyu, Li and Arora. Understanding the Generalization Benefit of Normalization Layers: Sharpness Reduction. arXiv preprint arXiv:2206.07085, 2022.
  15. Li, Wang and Li. Analyzing Sharpness along GD Trajectory: Progressive Sharpening and Edge of Stability. arXiv preprint arXiv:2207.12678, 2022.
  16. Damian, Ma and Lee. Label Noise SGD Provably Prefers Flat Global Minimizers. Advances in Neural Information Processing Systems, 2021.
  17. Li, Wang and Arora. What Happens after SGD Reaches Zero Loss? A Mathematical Framework. arXiv preprint arXiv:2110.06914, 2021.