Traditional convergence analyses of gradient-based algorithms assume learning rate η is set according to the basic relationship η<2/λ where λ is the largest eigenvalue of the Hessian of the objective, called sharpnessConfusingly, another traditional name for λ is smoothness.. Descent Lemma says that if this relationship holds along the trajectory of Gradient Descent, loss drops during each iteration. In deep learning where objectives are nonconvex and have multiple optima, similar analyses can show convergence towards stationary points and local minima. In practice, sharpness is unknown and η is set by trial and error. Since deep learning works, it has been generally assumed that this trial and error allows η to adjust to sharpness so that the theory applies. But recent empirical studies (Cohen et al., 2021; Ahn et al., 2022) showed compelling evidence to the contrary. On a variety of popular architectures and training datasets, GD with fairly small values of η displays following phenomena that they termed Edge of Stability (EoS): (a) Sharpness rises beyond 2/η, thus violating the above-mentioned relationship. (b) Thereafter sharpness stops rising but hovers noticeably above 2/η and even decreases a little. (c) Training loss behaves non-monotonically over individual iterations, yet consistently decreases over long timescales.
Note that (a) was already pointed out by Li et al. (2020b). Specifically, in modern deep nets, which use some form of normalization combined with weight decay, training to near-zero loss must lead to arbitrarily high sharpness. (However, Cohen et al. (2021) show that the EoS phenomenon appears even without normalization.) Phenomena (b), (c) are more mysterious, suggesting that GD with finite η is able to continue decreasing loss despite violating η<2/λ, while at the same time regulating further increase in value of sharpness and even causing a decrease. These striking inter-related phenomena suggest a radical overhaul of our thinking about optimization in deep learning. At the same time, it appears mathematically challenging to analyze such phenomena, at least for realistic settings and losses (as opposed to toy examples with 2 or 3 layers). The current paper introduces frameworks for doing such analyses.
We start by formal definition of stableness, ensuring that if a point + LR combination is stable then a gradient step is guaranteed to decrease the loss by the local version of Descent Lemma.
The second setting assumes that the loss is smooth but learning rate is effectively adaptive. We focus a concrete example, Normalized Gradient Descent, x←x−η∇L/∥∇L∥, which exhibits EoS behavior as ∇L→0. We can view Normalized GD as GD with a varying LR ηt=∥∇L(x(t))∥η, which goes to infinity when ∇L→0.
We show that Normalized GD on L (Section 4.3) and GD on L (Section 4.4) exhibit similar two-phase dynamics with sufficiently small LR η. In the first phase, GD tracks gradient flow (GF), with a monotonic decrease in loss until getting O(η)-close to the manifold (Theorems 4.3 and 4.5) and the stableness becomes larger than 2. In the second phase, GD no longer tracks GF and loss is not monotone decreasing due to the high stableness. Repeatedly overshooting, GD iterate jumps back and forth across the manifold while moving slowly along the direction in the tangent space of the manifold which decreases the sharpness. (See Figure 1 for a graphical illustration) Formally, we prove when η→0, the trajectory of GD converges to some limiting flow on the manifold. (Theorems 4.4 and 4.6) We further prove that in both settings GD in the second phase operates on EOS, and loss decreases in a non-monotone manner. Formally, we show that the average stableness over any two consecutive steps is at least 2 and that the average of L/η over two consecutive is proportional to sharpness or square root of sharpness. (Theorems 4.7 and 4.8)
Though many works have suggested (primarily via experiments and some intuition) that the training algorithm in deep learning implicitly selects out solutions of low sharpness in some way, we are not aware of a formal setting where this had ever been made precise. Note that our result requires no stochasticity as in SGD (Li et al., 2022b), though we need to inject tiny noise (e.g., of magnitude O(η100) ) to GD iterates occasionally (Algorithms 1 and 2). We believe that this is due the technical limitation of our current analysis and can be relaxed with a more advanced analysis. Indeed, in experiments, our theoretical predictions hold for the deterministic GD directly without any perturbation.
Novelty of Our Analysis:
To prove the alignment between the gradient and the top eigenvector of Hessian, it boils down to analyze Normalized GD on quadratic functions (2), which to the best of our knowledge has not been studied before. The dynamics is like chaotic version of power iteration, and we manage to show that the iterate will always align to the top eigenvector of Hessian of the quadratic loss. The proof is based on identifying a novel potential (Section 3) and might be of independent interest.
Related Works
Low sharpness has long been related to flat minima and thus to good generalization (Hochreiter and Schmidhuber, 1997; Keskar et al., 2016). Recent study on predictors of generalization (Jiang et al., 2020) does show sharpness-related measures as being good predictors, leading to SAM algorithm that improves generalization by explicitly controlling a parameter related to sharpness (Foret et al., 2021). However, Dinh et al. (2017) show that due to the positive homogeneity in the network architecture, networks with rescaled parameters can have very different sharpness yet be the same to the original one in function space. This observation weakens correlation between sharpness and and generalization gap and makes the definition of sharpness ambiguous. In face of this challenge, multiple notions of scale-invariant sharpness have been proposed (Yi et al., 2019a, b; Tsuzuku et al., 2020; Rangamani et al., 2021). Especially, Yi et al. (2021); Kwon et al. (2021) derived new algorithms with better generalization by explicitly regularizing new sharpness notions aware of the symmetry and invariance in the network. He et al. (2019) goes beyond the notion of sharpness/flatness and argues that the local minima of modern deep networks can be asymmetric, that is, sharp on one side, but flat on the other side.
Limiting Diffusion/Flow around Manifold of Minimizers:
The idea of analyzing the behavior of SGD with small LR along the the manifold originates from Blanc et al. (2020), which gives a local analysis on a special noise type named label noise, i.e. noise covariance is equal to Hessian at minimizers. Damian et al. (2021) extends this analysis and show SGD with label noise finds approximate stationary point for original loss plus some Hessian-related regularizer. The formal mathematical framework of approximating the limiting dynamics of SGD with arbitrary noise by Stochastic Differential Equations is later established by Li et al. (2022b), which is built on the convergence result for solutions of SDE with large-drift (Katzenberger, 1991).
Implicit Bias:
The notion that training algorithm plays an active role in selecting the solution (when multiple optima exist) has been termed the implicit bias of the algorithm (Gunasekar et al., 2018c) and studied in a large number of papers (Soudry et al., 2018; Li et al., 2018; Arora et al., 2018a, 2019a; Gunasekar et al., 2018b, a; Lyu and Li, 2020; Li et al., 2020a; Woodworth et al., 2020; Razin and Cohen, 2020; Lyu et al., 2021; Azulay et al., 2021; Gunasekar et al., 2021). In the infinite width limit, the implicit bias of Gradient Descent is shown to be the solution with the minimal RKHS norm with respect to the Neural Tangent Kernel (NTK) (Jacot et al., 2018; Li and Liang, 2018; Du et al., 2019; Arora et al., 2019b, c; Allen-Zhu et al., 2019b, a; Zou et al., 2020; Chizat et al., 2019; Yang, 2019). The implicit bias results from these papers are typically proved by performing a trajectory analysis for (Stochastic) Gradient Descent. Most of the results can be directly extended to the continuous limit (i.e., GD infinitesimal LR) and even some heavily relies on the conservation property which only holds for the continuous limit. In sharp contrast, the implicit bias shown in this paper – reducing the sharpness along the minimizer manifold – requires finite LR and doesn’t exist for the corresponding continuous limit. Other implicit bias results that fundamentally relies on the finiteness of LR includes stability analysis (Wu et al., 2017; Ma and Ying, 2021) and implicit gradient regularization (Barrett and Dherin, 2021), which is a special case of approximation results for stochastic modified equation by Li et al. (2017, 2019).
Non-monotone Convergence of Gradient Descent :
Warm-up: Quadratic Loss Functions
Our main result Theorem 3.1 is that the iterates of Normalized GD x(t) converge to v1 in direction, from which the loss oscillation Corollary 3.2 follows, suggesting that GD is operating in EoS. Since in quadratic case there is only one local minima, there is of course no need to talk about implicit bias. However, the observation that the GD iterates always align to the top eigenvector as well as the technique used in its proof play a very important role for deriving the sharpness-reduction implicit bias for the case of general loss functions.
Define x(t)=ηAx(t), and the following update rule (2) holds. It is clear that the convergence of xt to v1 in direction implies the convergence of xt as well.
If ∣⟨v1,x(t)⟩∣=0, ∀t≥0, then there exists 0<C<1 and s∈{±1} such that limt→∞x(2t)=Csλ1v1 and limt→∞x(2t+1)=(C−1)sλ1v1.
As a direct corollary, the loss oscillates as between time step 2t and time step 2t+1 as t→∞. This shows that the behavior of loss is not monotonic and hence indicates the edge of stability phenomena for the quadratic loss.
If ∣⟨v1,x(t)⟩∣=0, ∀t≥0, then there exists 0<C<1 such that limt→∞L(x(2t))=21C2λ1η2 and limt→∞L(x(2t+1))=21(C−1)2λ1η2.
For any j∈[D] and t≥λjλ1lnλjλ1+max{λD∥x(0)∥−λ1,0}, it holds that x(t)∈Ij.
First, we show for any j∈[D], Ij is indeed an invariant set for update rule (2) via Lemma A.1. With straightforward calculation, one can show that for any j∈[D], P(j:D)x(t) decreases by ∥x(t)∥λD∥P(j:D)x(t)∥ if P(j:D)x(t)≥λj (Lemma A.2). Setting j=1, we have ∥x(t)∥ decreases by λD if ∥x(t)∥≥λ1 (Corollary A.3). Thus for all t≥max{λD∥x(0)∥−λ1,0}, x(t)∈I1. Finally once x(t)∈I1, we can upper bound ∥x(t)∥ by λ1, and thus P(j:D)x(t) shrinks at least by a factor of λ1λD per step, which implies x(t) will be in Ij in another λjλ1lnλjλ1 steps.(Corollary A.4) ∎
Once the component of x(t) on an eigenvector becomes , it stays . So without loss of generality we can assume that after the preparation phase, the projection of x(t) along the top eigenvector v1 is non-zero, otherwise we can study the problem in the subspace excluding the top eigenvector.
If x(T)∈∩j=1DIj holds for some T, then for any t′,t such that T≤t≤t′ and ∥x(t)∥≤0.5λ1, it holds ∣⟨v1,x(t)⟩∣≤∣⟨v1,x(t′)⟩∣.
First, Lemma 3.5 (proved in Appendix A) shows that the norm of the iterate x(t) remains above 0.5λ1 for only one time-step.
For any t with x(t)∈∩j=1DIj, if ∥x(t)∥>2λ1, then ∥x(t+1)∥≤max(2λ1−2λ1λD2,λ1−∥x(t)∥).
Thus, for any t with x(t)∈∩j=1DIj and ∥x(t)∥≤2λ1, either ∥x(t+1)∥≤2λ1, or ∥x(t+1)∥>2λ1, which in turn implies that ∥x(t+2)∥≤2λ1 by Lemma 3.5. The proof of Lemma 3.4 is completed by induction on Lemma 3.6.
For any step t with ∥x(t)∥≤2λ1, for any k∈{1,2}, ∣⟨v1,x(t+k)⟩∣≥∣⟨v1,x(t)⟩∣.
Proof of case k=1 in Lemma 3.6 follows directly from plugging the assumption ∥x(t)∥≤2λ1 into (2) (See Lemma A.5). The case of k=2 in Lemma 3.6 follows from Lemma A.7. We defer the complete proof of Lemma 3.6 into Appendix A. ∎
To complete the proof for Theorem 3.1, we relate the increase in the projection along v1 at any step t, ∣⟨v1,x(t)⟩∣, to the magnitude of the angle between x(t) and the top eigenspace, θt. Briefly speaking, we show that if ∥x(t)∥≤2λ1, ∣⟨v1,x(t)⟩∣ has to increase by a factor of Θ(θt2) in two steps. Since ∣⟨v1,x(t)⟩∣ is bounded and monotone increases among {t∣∥x(t)∥≤2λ1} by Lemma 3.4, we conclude that θt gets arbitrarily small for sufficiently large t with ∥x(t)∥≤2λ1,∥x(t+2)∥≤2λ1 satisfied. Since the one-step normalized GD update Equation 2 is continuous when bounded away from origin, with a careful analysis, we conclude θt→0 for all iterates. Please see Section A.3 for details.
Below we show GD on loss L(x)=21x⊤Ax, Equation 3, follows the same update rule as Normalized GD on L(x)=21x⊤Ax, up to a linear transformation.
Denoting x(t)=η1(2A)1/2x(t), we can easily check x(t) also satisfies update rule (2).
Main Results
In this section we present the main results of this paper. Section 4.1 is for preliminary and notations. In Section 4.2, we make our key assumptions that the minimizers of the loss function form a manifold. In Sections 4.3 and 4.4 we present our main results for Normalized GD and GD on L respectively. In Section 4.5 we show the above two settings for GD do enter the regime of Edge of Statbility.
In this section, we focus on the setting where LR η goes to 0 and we fix the initialization xinit and the loss function L throughout this paper. We use O(⋅),Ω(⋅) to hide constants about xinit and L.
2 Key Assumptions on Manifold of Local Minimizers
Following Fehrman et al. (2020); Li et al. (2022b), we make the following assumption throughout the paper.
The smoothness assumption is satisfied for networks with smooth activation functions like tanh and GeLU (Hendrycks and Gimpel, 2016). The existence of manifold is due to the vast overparametrization in modern deep networks and preimage theorem. (See a discussion in section 3.1 of Li et al. (2022b)) The assumption rank(∇2L(x))=M basically say ∇2L(x) always attains the maximal rank in the normal space of the manifold, which ensures the differentiability of Φ and is crucial to our current analysis, though it’s not clear if it is necessary. We also make the following assumption to ensure that λ1(∇2L(⋅)) is differentiable, which is necessary for our main results, Theorems 4.4 and 4.6.
For any x∈Γ, ∇2L(x) has a positive eigengap, i.e., λ1(∇2L(x))>λ2(∇2L(x)).
3 Results for Normalized GD
We first denote the iterates of Normalized GD with LR η by xη(t), with xη(0)≡xinit for all η:
The first theorem demonstrates the movement in the manifold, when the iterate travels from xinit to a position that is O(η) distance closer to the manifold (more specifically, Φ(xinit)). Moreover, just like the result in the quadratic case, we have more fine-grained bounds on the projection of xη(t)−Φ(xη(t)) into the bottom-k eigenspace of ∇2L(Φ(xη(t))) for every k∈[D]. For convenience, we define the following quantity for all j∈[d] and x∈U:
In the quadratic case, Lemma 3.3 shows that Rj(x) will eventually become non-positive for normalized GD iterates. Similarly, for the general loss, the following theorem shows that Rj(xη(t)) eventually becomes approximately non-positive (smaller than O(η2)) in O(η1) steps.
Our main contribution is the analysis for the second phase (Theorem 4.4), which says just like the quadratic case, the angle between xη(t) and the top eigenspace of ∇2L(Φ(xη(t))), denoted by θt, will be O(η) on average. And as a result, the dynamics of Normalized GD tracks the riemannian gradient flow with respect to log(λ1(∇2L(⋅))) on manifold, that is, the unique solution of Equation 5, where Px,Γ⊥ is the projection matrix onto the tangent space of manifold Γ at x∈Γ.
Note Equation 5 is not guaranteed to have a global solution, i.e., a well-defined solution for all τ≥0, for the following two reasons: (1). when the multiplicity of top eigenvalue is larger than 1, λ1(∇2L(⋅)) may be not differentiable and (2). the projection matrix is only defined on Γ and the equation becomes undefined when the solution leaves Γ, i.e., moving across the boundary of Γ. For simplicity, we make 4.2 that every point on Γ has a positive eigengap. Or equivalently, we can work with a slightly smaller manifold Γ′={x∈Γ∣λ1(x)>λ2(x)}.
Towards a mathematical rigorous characterization of the dynamics in the second phase, we need to make the following modifications: (1). we add negligible noise of magnitude O(η100) every η−0.1 steps, (2). we assume for each η>0, there exist some step t=Θ(1/η) in phase I, except the guaranteed condition (1) and (2) (by Theorem 4.3, the additional condition (3) also holds. This assumption is mild because we only require (3) to hold for one step among Θ(1/η) steps from ηT1 to ηT1′, where T1 is the constant given by Theorem 4.3 and T1′ is arbitrary constant larger than T1. This assumption also holds empirically for all our experiments in Section 6.
4 Results for GD on L𝐿\sqrt{L}
In this subsection, we denote the iterates of GD on L with LR η by xη(t), with xη(0)≡xinit for all η:
Similar to Normalized GD, we will have two phases. The first theorem demonstrates the movement in the manifold, when the iterate travels from xinit to a position that is O(η) distance closer to the manifold. For convenience, we will denote the quantity ∑i=jMλi(x)⟨vi(x),x−Φ(x)⟩2−η1/2λj(x) by Rj(x) for all j∈[M] and x∈U.
The next result demonstrates that close to the manifold, the trajectory implicitly minimizes sharpness.
5 Operating on the Edge of Stability
In this section, we show that both Normalized GD on L and GD on L is on Edge of Stability in their phase II, that is, at least in one of every two consecutive steps, the stableness is at least 2 and the loss oscillates in every two consecutive steps. Interestingly, the average loss over two steps decreases over time, even when operating on the edge of Stability (see Figure 1 for illustration), as indicated by the following theorems. Note that Theorems 4.4 and 4.6 ensures that the average of θt are O(η) and O(η). We defer their proofs into Sections E.5 and G.4 respectively.
Under the setting of Theorem 4.4, by viewing Normalized GD as GD with time-varying LR ηt:=∥∇L(xη(t))∥η, we have [SL(xη(t),ηt)]−1+[SL(xη(t+1),ηt+1)]−1=1+O(θt+η). Moreover, we have L(xη(t))+L(xη(t+1))=η2λ1(∇2L(xη(t)))+O(ηθt).
Under the setting of Theorem 4.6, we have [SL(xη(t),ηt)]≥Ω(θt1). Moreover, we have L(xη(t))+L(xη(t+1))=ηλ1(∇2L(xη(t)))+O(ηθt).
Proof Overview
We sketch the proof of the Normalized GD in phase I and II respectively in Section 5.2. Then we briefly discuss how to prove the results for GD with L with same analysis in Section 5.3. We start by introducing the properties of limit map of gradient flow Φ in Section 5.1, which plays a very important role in the analysis.
The limit map of gradient flow Φ lies at the core of our analysis. When LR η is small, one can show xη(t) will be O(η) close to manifold and Φ(xη(t)). Therefore, Φ(xη(t)) captures the essential part of the implicit regularization of Normalized GD and characterization of the trajectory of Φ(xη(t)) immediately gives us that of Φ(xη(t)) up to O(η).
Below we first recap a few important properties of Φ that will be used later this section, which makes the analysis of Φ(xη(t)) convenient.
Under 4.1, Φ satisfies the following two properties:
∂Φ(x)∇L(x)=0 for any x∈U. (Lemma B.16)
For any x∈Γ, if λ1(x)>λ2(x), ∂2Φ(x)[v1(x),v1(x)]=−21Px,Γ⊥∇logλ1(x). (Lemmas B.18 and B.20)
Note that xη(t+1)−xη(t)=−η∥∇L(xη(t))∥∇L(xη(t)), using a second order taylor expansion of Φ, we have
where we use the first claim of Lemma 5.1 in the final step. Therefore, we have Φ(xη(t+1))−Φ(xη(t))=O(η2), which means Φ(xη(t)) moves slowly along the manifold, at a rate of at most O(η2) step. The Taylor expansion of Φ, (8) plays a crucial role in our analysis for both Phase I and II and will be used repeatedly.
2 Analysis for Normalized GD
The Phase I itself can be divided into two subphases: (A). Normalized GD iterate xη(t) gets O(η) close to manifold; (B). counterpart of preparation phase in the quadratic case: local movement in the O(η)-neighborhood of the manifold which decreases Rj(xη(t)) to O(η2). Below we sketch their proofs respectively:
Subphase (A): First, with a very classical result in ODE approximation theory, normalized GD with small LR will track the normalized gradient flow, which is a time-rescaled version of standard gradient flow, with O(η) error, and enter a small neighborhoods of the manifold where Polyak-Łojasiewicz (PL) condition holds. Since then, Normalized GD decreases the fast loss with PL condition and the gradient has to be O(η) small in O(η1) steps. (See details in Section C.1).
Subphase (B): The result in subphase (B) can be viewed as a generalization of Lemma 3.3 when the loss function is O(η)-approximately quadratic, in both space and time. More specifically, it means ∇2L(Φ(xη(t)))−∇2L(x)≤O(η) for all x which is O(η)-close to some Φ(xη(t′)) with t′−t≤O(1/η). This is because by Taylor expansion (8), ∥Φ(xη(t))−Φ(xη(t′))∥=O(η2(t′−t))=O(η), and again by Taylor expansion of ∇2L, we know ∇2L(x)−∇2L(Φ(xη(t)))=O(∥x−Φ(xη(t))∥)=O(η).
With a similar proof technique, we show xη(t) enters ainvariant set around the manifold Γ, that is, {x∈U∣Rj(x)≤O(η2),∀j∈[D]}. Formally, we show the following analog of Lemma 3.3:
Let {xη(t)}t≥0 be the iterates of Normalized GD (4) with LR η. If for some step t0, ∥xη(t0)−Φ(xη(t0))∥=O(η), then for sufficiently small LR η and all steps t∈[t0+Θ(1),Θ(η−2)] steps, the iterate xη(t) satisfy maxj∈[M]Rj(xη(t))≤O(η2).
Analysis for Phase II, Theorem 4.4:
Similar to the subphase (B) in the Phase I, the high-level idea here is again that xη(t) locally evolves like normalized GD with quadratic loss around Φ(xη(t)) and with an argument similar to the alignment phase of quadratic case (though technically more complicated), we show xη(t)−Φ(xη(t)) approximately aligns to the top eigenvector of ∇2L(Φ(xη(t))), denoted by v1(t) and so does ∇L(xη(t)). More specifically, it corresponds to the second claim in Theorem 4.4, that ⌊T2/η2⌋1∑t=0⌊T2/η2⌋θt≤O(η).
We now have a more detailed look at the movement in Φ. Since Φ(xη(t)) belongs to the manifold, we have ∇L(Φ(xη(t)))=0 and so ∇L(xη(t))=∇2L(Φ(xη(t)))(xη(t)−Φ(xη(t)))+O(η2) using a Taylor expansion. This helps us derive a relation between the Normalized GD update and the top eigenvector of the hessian (simplified version of Lemma B.9):
Incorporating the above into the movement in Φ(xη(t)) from Equation 8 gives:
Applying the second property of Lemma 5.1 on Equation 10 above yields Lemma 5.3.
Under the setting in Theorem 4.4, for sufficiently small η, we have at any step t≤⌊T2/η2⌋
To complete the proof of Theorem 4.4, we show that for small enough η, the trajectory of Φ(xη(τ/η2)) is O(η3⌊T2/η2⌋+η2∑t=0⌊T2/η2⌋θt)-close to X(τ) for any τ≤T2, where X(⋅) is the flow given by Equation 5. This error is O(η), since ∑t=0⌊T2/η2⌋θt=O(⌊T2/η2⌋η).
One technical difficulty towards showing the average of ηt is only O(η) is that our current analysis requires ∣⟨v1(xη(t)),xη(t)−Φ(xη(t))⟩∣ doesn’t vanish, that is, it remains Ω(η) large throughout the entire training process. This is guaranteed by Lemma 3.4 in quadratic case – since the alignment monotone increases whenever it’s smaller 2λ1, but the analysis breaks when the loss is only approximately quadratic and the alignment ∣⟨v1(xη(t)),xη(t)−Φ(xη(t))⟩∣could decrease decrease by O(θtη2) per step. Once the alignment becomes too small, even if the angle θt is small, the normalized GD dynamics become chaotic and super sensitive to any perturbation. Our current proof technique cannot deal with this case and that’s the main reason we have to make the additional assumption in Theorem 4.4.
Role of η100 noise. Fortunately, with the additional assumption that the initial alignment is at least Ω(η), we can show adding any poly(η) perturbation (even as small as Ω(η100)) suffices to prevent the aforementioned bad case, that is, ∣⟨v1(xη(t)),xη(t)−Φ(xη(t))⟩∣ stays Ω(η) large. The intuition why Ω(η100) perturbation works again comes from quadratic case – it’s clear that x=cv1 for any ∣c∣≤1 is a stationary point for two-step normalized GD updates for quadratic loss under the setting of Section 3. But if c is smaller than critical value determined by the eigenvalues of the hessian, the stationary point is unstable, meaning any deviation away from the top eigenspace will be amplified until the alignment increases above the critical threshold. Based on this intuition, the formal argument, Lemma E.11 uses the techniques from the ‘escaping saddle point’ analysis (Jin et al., 2017). Adding noise is not necessary in experiments to observe the predicted behavior (see ‘Alignment’ in Figure 4 where no noise is added). On one hand, it might be because the floating point errors served the role of noise. On the other hand, we suspect it’s not necessary even for theory, just like GD gets stuck at saddle point only when initialized from a zero measure set even without noise (Lee et al., 2016, 2017).
3 Analysis for GD on L𝐿\sqrt{L}
In this subsection we will make an additional assumption that L(x)=0 for all x∈Γ. The analysis then will follow a very similar strategy as the analysis for (Normalized) GD. However, the major difference from the analysis for Normalized GD comes from the update rule for xη(t) when it is O(η)-close to the manifold:
Thus, the effective learning rate is λ1(t)η at any step t. This shows up, when we compute the change in the function Φ. Thus, we have the following lemma showcasing the movement in the function Φ with the GD update on L:
Under the setting in Theorem 4.6, for sufficiently small η, we have at any step t≤⌊T2/η2⌋, Φ(xη(t+1))−Φ(xη(t))=−8η2Pt,Γ⊥∇λ1(t)+O(η3+η2θt).
Experiments
Though our main theorems characterizes the dynamics of Nomalized GD and GD on L for sufficiently small LR, it’s not clear if the predicted phenomena is related to the training with practical LR as the function and initialization dependent constants are hard to compute and could be huge. Neverthesless, in this section we show the phenomena predicted by our theorem does occur for real-life models like VGG-16. We further verify the predicted convergence to the limiting flow for Normalized GD on a two-layer fully-connected network trained on MNIST.
We observe that the alignment function reaches close to 1, towards the end of training. The top eigenvalue decreases over time (as predicted byTheorem 4.4 and Theorem 4.6), and the stableness hovers around 2 at the end of training.
Verifying Convergence to Limiting Flow on MNIST:
Conclusion
The recent discovery of Edge of Stability phenomenon in Cohen et al. (2021) calls for a reexamination of how we understand optimization in deep learning. The current paper gives two concrete settings with fairly general loss functions, where gradient updates can be shown to decrease loss over many iterations even after stableness is lost. Furthermore, in one setting the trajectory is shown to amount to reduce the sharpness (i.e., the maximum eigenvalue of the Hessian of the loss), thus rigorously establishing an effect that has been conjectured for decades in deep learning literature and was definitively documented for GD in Cohen et al. (2021). Our analysis crucially relies upon learning rate η being finite, in contrast to many recent results on implicit bias that required an infinitesimal LR. Even the alignment analysis of Normalized GD to the top eigenvector for quadratic loss in Section 3 appears to be new.
One limitation of our analysis is that it only applies close to the manifold of local minimizers. By contrast, in experiments the EoS phenomenon, including the control of sharpness, begins much sooner. Addressing this gap, as well as analysing the EoS for the loss L itself (as opposed to L as done here) is left for future work. Very likely this will require novel understanding of properties of deep learning losses, which we were able to circumvent by looking at L instead. Exploration of EoS-like effects in SGD setting would also be interesting, although we first need definitive experiments analogous to Cohen et al. (2021).
Acknowledgement
We thank Kaifeng Lyu for helpful discussions. The authors acknowledge support from NSF, ONR, Simons Foundation, Schmidt Foundation, Mozilla Research, Amazon Research, DARPA and SRC. ZL is also supported by Microsoft Research PhD Fellowship.
References
Appendix A Omitted Proofs for Results for Quadratic Loss Functions
Recall the loss function L is defined as L(x)=21x⊤Ax. The Normalized GD update (LR= η )is given by x(t+1)=x(t)−η∥Ax(t)∥Ax(t). A substitution x(t):=ηAx(t) gives the following update rule:
Now we recall the main theorem for Normalized GD on quadratic loss functions: See 3.1
We also note that GD on L with any LR η can also be reduced to update rule (2), as shown in the discussion at the end of Section 3.
In this subsection, we show (1). Ij is indeed an invariant set for normalized GD ∀j∈[D] and (2). from any initialization, normalized GD will eventually go into their intersection ∩j=1DIj.
Note P(j:D)A=P(j:D)AP(j:D), by definition of Normalized GD (2), we have
Note that P(j:D)A≼λjI, P(j:D)x(t)≤∥x(t)∥ and P(j:D)x(t)≤λj by assumption, we have
Therefore I−∥x(t)∥P(j:D)A≤∥P(j:D)x(t)∥λj and thus we conclude P(j:D)x(t+1)≤λj. ∎
Since λj≤P(j:D)x(t)≤∥x(t)∥, we have 0≼I−∥x(t)∥P(j:D)A≼1−∥x(t)∥λD. Therefore I−∥x(t)∥P(j:D)A≤1−∥x(t)∥λD. The proof is completed by plugging this into Equation 11. ∎
Lemma A.2 has the following two direct corollaries.
For any initialization x(0) and t≥λD∥x(0)∥−λ1, ∥x(t)∥≤λ1, that is, x(t)∈I1.
Set j=1 in Lemma A.2, it holds that ∥x(t+1)∥≤∥x(t)∥−λD whenever ∥x(t)∥≥λ1. Thus \left\|\widetilde{x}(\big{\lceil}\frac{\left\|\widetilde{x}(0)\right\|-\lambda_{1}}{\lambda_{D}}\big{\rceil})\right\|\leq\lambda_{1}. The proof is completed as I1 is an invariant set by Lemma A.1. ∎
For any coordinate j∈[D] and initial point x(0)∈I1, if t≥λDλ1lnλjλ1 then P(j:D)x(t)≤λj.
Since I1 is an invariant set, we have ∥x(t)∥≤λ1 for all t≥0. Thus let T=⌊λDλ1lnλjλ1⌋, we have
The proof is completed since Ij is a invariant set for any j∈[D] by Lemma A.1. ∎
A.2 Proofs for Alignment Phase
In this subsection, we analyze how normalized GD align to the top eigenvector once it goes through the preparation phase, meaning x(t)∈∩j=1DIj for all t in alignment phase.
Let the index k be the smallest integer such that λk+1<2∥x(t)∥−λ1. If no such index exists, then one can observe that ∥x(t+1)∥≤λ1−∥x(t)∥. Assuming that such an index exists in [D], we have λk≥2∥x(t)∥−λ1 and ∥x(t)∥−λj≤λ1−∥x(t)∥, ∀j≤k. Now consider the following vectors:
By definition of k, ∣∥x(t)∥−λj∣≤∣∥x(t)∥−λ1∣. Thus
By assumption, we have x(t)∈∩j=1DIj. Thus
where we applied AM-GM inequality multiple times in the pre-final step.
where the final step is because 2λ1≤∥x(t)∥≤λ1 and that the maximal value of a convex function is attained at the boundary of an interval.
At any step t and i∈[D], if ∥x(t)∥⪌2λi, then ∣xi(t+1)∣⪋∣xi(t)∣, where ⪌ denotes larger than, equal to and smaller than respectively. (Same for ⪋, but in the reverse order)
From the Normalized GD update rule, we have xi(t+1)=xi(t)(1−∥x(t)∥λi), for all i∈[D]. Thus
At any step t, if ∥x(t)∥≤2λ1, then
where θt=arctan∣e1⊤x(t)∣∥P(2:D)x(t)∥ and λ=min(λ1−λ2,λD).
We first show that the left side inequality holds by the following update rule for ⟨e1,x(t)⟩:
Since ∥x(t+1)∥≥∣⟨e1,x(t+1)⟩∣ and θt denotes the angle between e1 and x(t+1), we get the left side inequality.
Now, we focus on the right hand side inequality. First of all, the update in the coordinate j∈[2,D] is given by
where again in the final step, we have used ∥x(t)∥<2λ1. The above bound can be further bounded by
where we have used λ=min(λ1−λ2,λD).
If at some step t, ∥x(t+1)∥+∥x(t)∥≤λ1, then ∣x1(t+2)∣≥∣x1(t)∣, where the equality holds only when ∥x(t+1)∥+∥x(t)∥=λ1. Therefore, by Lemma A.6, we have :
where θt=arctan∣e1⊤x(t)∣∥P(2:D)x(t)∥, and λ=min(λ1−λ2,λD).
Using the Normalized GD update rule, we have
where the equality holds only when ∥x(t+1)∥+∥x(t)∥=λ1.
Moreover, with the additional condition that ∥x(t)∥<2λ1, we have from Lemma A.6, ∥x(t+1)∥≤λ1−∥x(t)∥−λ(λ1−λ)sin2θt, where λ=min(λ1−λ2,λD).
Hence, retracing the steps we followed before, we have
where the final step follows from ∥x(t+1)∥≤λ1−∥x(t)∥ and therefore ∥x(t+1)∥∥x(t)∥≤4λ12. ∎
A.3 Proof of Main theorems for Quadratic Loss
Preparation phase: x(t) enters and stays in an invariant set around the origin, that is, ∩j=1DIj, where Ij:={x∣∑i=jD⟨ei,x(t)⟩2≤λj2}. (See Lemma 3.3, which is a direct consequence of Lemmas A.1, A.3 and A.1.)
Alignment phase: The projection of x(t) on the top eigenvector, ∣⟨x(t),e1⟩∣, is shown to increase monotonically among the steps among the steps {t∣∥x(t)∥≤0.5}, up until convergence, since it’s bounded. (Lemma 3.4)
By Lemma A.7, the convergence of ∣⟨x(t),e1⟩∣ would imply the convergence of x(t) to e1 in direction.
Now we claim ∀t≥3, there is some k∈{0,1,3} such that t−k∈S′. This is because Lemma 3.5 says that if t∈/S, then both t−1,t+1∈S. Thus for any t∈/S, t−1∈S′. Therefore, for any t∈S/S′, if t−2∈/S, then t−3∈S′. Thus we conclude that ∀t≥3, there is some k∈{0,1,3} such that t−k∈S′, which implies t→∞limθt=0. Hence t→∞lim∥x(t+1)−x(t)∥=λ1, meaning for sufficiently large t, x1(t) flips its sign per step and thus t→∞limx(t+2)−x(t)=0, t→∞lim∥x(t+1)∥+∥x(t)∥=λ1.
If C=21, then we must have t→∞lim∥x(t)∥=2λ1 and we are done in this case. If C<21, note that t→∞,t∈S′lim∣x1(t)∣=Cλ1, it must hold that t→∞,t∈S′lim∥x(t+1)∥=(1−C)λ1, thus there is some large T∈S such that for all t∈S,t≥T, t+1∈/S. By Lemma 3.5, t+2∈S. Thus we conclude t→∞limx(T+2t)=Cλse1 for some s∈{−1,1} and thus t→∞limx(T+2t+1)=(C−1)λse1. This completes the proof. ∎
A.4 Some Extra Lemmas (only used in the general loss case)
For a general loss function L satisfying 4.1, the loss landscape looks like a strongly convex quadratic function locally around its minimizer. When sufficient small learning rate, the dynamics will be sufficiently close to the manifold and behaves like that in quadratic case with small perturbations. Thus it will be very useful to have more refined analysis for the quadratic case, as they allow us to bound the error in the approximate quadratic case quantitatively. Lemmas A.8, A.9, A.10 and A.11 are such examples. Note that they are only used in the proof of the general loss case, but not in the quadratic loss case.
Lemma A.8 is a slightly generalized version of Lemma 3.5.
Suppose at time t, P(j:D)x(t)≤λj(1+λ12λD2), for all j∈[D], if ∥x(t)∥>2λ1, then ∥x(t+1)∥≤2λ1.
The proof is similar to the proof of Lemma 3.5. Let the index k be the smallest integer such that λk+1<2∥x(t)∥−λ1. If no such index exists, then one can observe that ∥x(t+1)∥≤λ1−∥x(t)∥. Assuming that such an index exists in [D], we have λk≥2∥x(t)∥−λ1 and ∥x(t)∥−λj≤λ1−∥x(t)∥, ∀j≤k. With the same decomposition and estimation, since x(t)∈∩j=1D(1+λ12λD2)Ij, we have
∣⟨e1,x(t)⟩∣≤(1−2c)g(λk).
θt≤c∣⟨e1,x(t)⟩∣,
where θt=arctan∣⟨e1,x(t)⟩∣∥P(2:D)(x(t))∥.
From the quadratic update, we have the update rule as:
Thus, as long as, the following holds true:
We can use (λ1−∥x(t)∥)cosθt≤∥x(t+1)∥≤λ1−∥x(t)∥−2λ1λ(1−λ1λ)λ1sin2θt, where λ=min(λ1−λ2,λD) from Lemma A.6 to show the following with additional algebraic manipulation:
where the last step we use that ∣θt∣≤c∣⟨e1,x(t)⟩∣, we only need
The above inequality is true when ∣⟨e1,x(t)⟩∣≤(1−2c)g(λk). ∎
Then, the following must hold true at time t.
where the final step holds true for any c∈(0,1).
The result follows after substituting this bound in Equation 12.
implying ∣xi(t+1)∣<(1−∥x(t)∥1)∣xi(t)∣ for all i∈[2,D], since λi<1.
Since λi<λ1 and ∥x(t)∥≤2λ1, it holds that
Recall ∣tan(∠(v,e1))∣=∣⟨e1,v⟩∣∥P(2:D)v∥ for any vector v, the first claim follows from re-arranging the terms.
For the second claim, it suffices to apply the above inequality to t+1, which yields that
The proof is completed by noting ∥x(t+1)∥≤λ1−∥xt∥ (Lemma A.6) and tan(∠(x(t+1),e1))≤tan(∠(x(t),e1)). ∎
Appendix B Setups for General Loss Functions
Before we start the analysis for Normalized GD for general loss functions in Appendix C, we need to introduce some new notations and terminologies to complete the formal setup. We start by first recapping some core assumptions and definitions in the main paper and provide the missing proof in the main paper.
We define Φ as the limit map of gradient flow below. We summarize various properties of Φ from LABEL:ch:diffusion_on_manifold in Section B.2.
Given any two points x,y, we use xy to denote the line segment between x and y, i.e., {z∣∃λ∈,z=(1−λ)x+λy}.
The main result of this chapter focuses on the trajectory of Normalized GD from fixed initialization xinit with LR η converges to , which can be roughly split into two phases. In the first phase, Theorem 4.3 shows that the normalized GD trajectory converges to the gradient flow trajectory, ϕ(xinit,⋅). In second phase, Theorem 4.4 shows that the normalized GD trajectory converges to the limiting flow which decreases sharpness on Γ, (5). Therefore, for sufficiently small η, the entire trajectory of normalized GD will be contained in a small neighbourhood of gradient flow trajectory Z and limiting flow trajectory Y. The convergence rate given by our proof depends on the various local constants like smoothness of L and Φ in this small neighbourhood, which intuitively can be viewed as the actual ”working zone” of the algorithm. The constants are upper bounded or lower bounded from zero because this ”working zone” is compact after fixing the stopping time of (5), which is denoted by T2.
We construct the ”working zone” of the second phase, Yρ and Yϵ in Lemmas B.2 and B.5 respectively, where 0<ϵ<ρ, implying Yϵ⊂Yρ. The reason that we need the two-level nested ”working zones” is that even though we can ensure all the points in Yρ have nice properties as listed in Lemma B.2, we cannot ensure the trajectory of gradient flow from x∈Yρ to Φ(x) or the line segment xΦ(x) is in Yρ, which will be crucial for the geometric lemmas (in Section B.1) that we will heavily use in the trajectory analysis around the manifold. For this reason we further define Yϵ and Lemma B.5 guarantees the trajectory of gradient flow from x to Φ(x) or the line segment xΦ(x) whenever x∈Yρ.
A function L is said to be μ-PL in a set U iff for all x∈U,
For convenience, we define \bm{\Delta}:=\frac{1}{2}\inf_{x\in Y}\big{(}\lambda_{1}(\nabla^{2}L(x))-\lambda_{2}(\nabla^{2}L(x)))\big{)} and μ:=41infx∈YλM(∇2L(x)). By 4.1, we have μ>0. By 4.2, Δ>0.
Given Y, there are sufficiently small ρ>0 such that
We first claim for every y∈Y, for all sufficiently small ρy>0 (i.e. for all ρy smaller than some threshold depending on y), the following three properties hold (1) By(ρy)∩Γ is compact; (2) By(ρy)∩Γ⊂U and (3) L is μ-PL on By(ρy∩Γ).
Thus for any c′>0, for sufficiently small ρy, (x−p(x))⊤∇2L(p(x))(x−p(x))≥c′∥x−p(x)∥3. Combining Equations 14 and 15, we conclude that for sufficiently small ρy,
Again for sufficiently small ρy, by Taylor expansion of L at p(x), we have
Meanwhile, since λM(∇2L(p(x))) and λ1(∇2L(p(x)))−λ2(∇2L(p(x))) are continuous functions in x, we can also choose a sufficiently small ρy such that for all x∈By(ρy), λM(∇2L(p(x)))≥21λM(∇2L(p(y)))=21λM(∇2L(y))>Δ and \lambda_{1}(\nabla^{2}L(p(x)))-\lambda_{2}(\nabla^{2}L(p(x)))\geq\frac{1}{2}\big{(}\lambda_{1}(\nabla^{2}L(p(y)))-\lambda_{2}(\nabla^{2}L(p(y)))\big{)}=\frac{1}{2}\big{(}\lambda_{1}(\nabla^{2}L(y))-\lambda_{2}(\nabla^{2}L(y))\big{)}\geq\bm{\mu}. Further note Y⊂∪y∈YBy(ρy) and Y is a compact set, we can take a finite subset of Y, Y′, such that Y⊂∪y∈Y′By(ρy). Taking ρ:=miny∈Y′2ρy completes the proof. ∎
We define the following constants regarding smoothness of L and Φ of various orders over Yρ.
Given ρ as defined in Lemma B.2, there is an ϵ∈(0,ρ) such that
supx∈YϵL(x)−x∈YϵinfL(x)<min(8μρ2,ν2ζ22μ5);
∀x∈Yϵ, Φ(x)∈Y2ρ.
For every y∈Y, there is an ϵy, such that ∀x∈By(ϵy), it holds that L(x)<min(8μρ2,ν2ζ22μ4) and Φ(x)∈Y2ρ, as both L(x) and Φ(x) are continuous. Further note Y⊂∪y∈YBy(ϵy) and Y is a compact set, we can take a finite subset of Y, Y′, such that Y⊂∪y∈Y′By(ϵy). Taking ϵ:=miny∈Y′2ϵy completes the proof. ∎
Summary for Setups:
The initial point xinit is chosen from an open neighborhood of manifold Γ, U, where the infinite-time limit of gradient flow Φ is well-defined and for any x∈U, Φ(x)∈Γ. We consider normalized GD with sufficiently small LR η such that the trajectory enters a small neighborhood of limiting flow trajectory, Yρ. Moreover, L is μ-PL on Yρ and the eigengaps and smallest eigenvalues are uniformly lower bounded by positive Δ,μ respectively on Yρ. Finally, we consider a proper subset of Yρ, Yϵ, as the final ”working zone” in the second phase (defined in Lemma B.5), which enjoys more properties than Yρ, including Lemmas B.7, B.8, B.9 and B.10.
B.1 Geometric Lemmas
In this subsection we present several geometric lemmas which are frequently used in the trajectory analysis of normalized GD. In this section, O(⋅) only hides absolute constants. Below is a brief summary:
Lemma B.6: Inequalities connecting various terms: the distance between x and Φ(x), the length of GF trajectory from x to Φ(x), square root of loss and gradient norm;
Lemma B.7: For any x∈Yϵ, the gradient flow trajectory from x to Φ(x) and the line segment between x and Φ(x) are all contained in Yρ, so it’s ”safe” to use Taylor expansions along GF trajectory or xΦ(x) to derive properties;
Lemmas B.8, B.9 and B.10: for any x∈Yϵ, the normalized GD dynamics at x can be roughly viewed as approximately quadratic around Φ(x) with positive definite matrix ∇2L(Φ(x)).
Lemma B.11: In the ”working zone”, Yρ, one-step normalized GD update with LR η only changes Φ(xt) by O(η2).
Lemma B.13: In the ”working zone”, Yρ, one-step normalized GD update with LR η decreases L(x)−miny∈YL(y) by η42μ if ∥∇L(x)∥≥ηζ.
If the trajectory of gradient flow starting from x, ϕ(x,t), stays in Yρ for all t≥0, then we have
Since Φ(x) is defined as limt→∞ϕ(x,t) and ϕ(x,0)=x, the left-side inequality follows immediately from triangle inequality. The right-side inequality is by the definition of PL condition. Below we prove the middle inequality.
Since ∀t≥0, ϕ(x,t)∈Yρ, it holds that ∥∇L(ϕ(x,t))∥2≥2μ(L(ϕ(x,t))−L(Φ(x))) by the choice of ρ in Lemma B.2. Without loss of generality, we assume L(y)=0,∀y∈Γ. Thus we have
The proof is complete since ϕ(x,0)=x and we assume L(Φ(x)) is . ∎
Let ρ,ϵ be defined in Lemmas B.2 and B.5. For any x∈Yϵ, we have
The entire trajectory of gradient flow starting from x is contained in Yρ, i.e., ϕ(x,t)∈Yρ, ∀t≥0;
Moreover, ∥Φ(x)−ϕ(x,t)∥≤min(ρ,νζ2μ2), ∀t≥0.
Let time τ∗≥0 be the smallest time after which the trajectory of GF is completely contained in Yρ, that is, τ∗:=inf{t≥0∣∀t′≥t,ϕ(x,t′)∈Yρ}. Since Yρ is closed and ϕ(x,⋅) is continuous, we have ϕ(x,τ∗)∈Yρ.
Since ∀τ≥τ∗, ϕ(x,τ)∈Yρ, by Lemma B.6, it holds that ∥ϕ(x,τ∗)−Φ(x)∥≤μ2(L(ϕ(x,τ∗))−L(Φ(x))).
Note that loss doesn’t increase along GF, we have L(ϕ(x,τ∗))−L(Φ(x))≤L(x)−L(Φ(x))≤8μρ2, which implies that ∥ϕ(x,τ∗)−Φ(x)∥≤2ρ. Therefore τ∗ must be , otherwise there exists a 0<τ′<τ∗ such that ∥ϕ(x,τ)−Φ(x)∥≤ρ for all τ′<τ<τ∗ by the continuity of ϕ(x,⋅). This proves the first claim.
Given the first claim is proved, the second claim follows directly from Lemma B.6.
The following theorem shows that the projection of x in the tangent space of Φ(x) is small when x is close to the manifold. In particular if we can show that in a discrete trajectory with a vanishing learning rate η, the iterates {xη(t)} stay in Yϵ, we can interchangeably use ∥xη(t)−Φ(xη(t))∥ with ∥Pt,Γ(xη(t)−Φ(xη(t)))∥, with an additional error of O(η3), when ∥Pt,Γ(xη(t)−Φ(xη(t)))∥≤O(η).
For all x∈Yϵ, we have that
First of all, we can track the decrease in loss along the Gradient flow trajectory starting from x. At any time τ, we have
where ϕ(x,0)=x. Without loss of generality, we assume L(y)=0,∀y∈Γ. Using the fact that L is μ-PL on Yρ and the GF trajectory starting from any point in Yϵ stays inside Yρ (from Lemma B.7), we have
Moreover, we can relate L(ϕ(x,0) with ∥Φ(x)−x∥ with a second order taylor expansion:
where in the final step, we have used the fact that L(Φ(x))=0 and ∇L(Φ(x))=0. By Lemma B.7, we have xΦ(x)⊂Yρ. Thus maxs∈∇2L(sx+(1−s)Φ(x))≤ζ from Definition B.4 and it follows that
Finally we focus on the movement in the tangent space. It holds that
By Lemma B.7, we have ϕ(x,τ)Φ(x)⊂Yρ for all τ≥0 and thus
Since PΦ(x),Γ⊥ is the projection matrix for the tangent space, PΦ(x),Γ⊥∇2L(Φ(x))=0 and thus by Equation 16
Plug Equation 19 into Equation 18, we conclude that
The left-side inequality of the second inequality is proved by plugging the first claim into the above inequality Equation 20 and rearranging the terms. Note by the second claim in Lemma B.7, 4μ2νζ∥x−Φ(x)∥≤21, the right-side inequality is also proved. ∎
At any point x∈Yϵ, we have
Moreover, the normalized gradient of L can be written as
Using taylor expansion at x, we have using ∇L(Φ(x))=0:
where we use Lemma B.8 since x∈Yϵ. Thus, the normalized gradient at any step t can be written as
Consider any point x∈Yϵ. Then,
where θ=arctan∣⟨v1(x),x⟩∣PΦ(x),Γ(2:M)x, with x=∇2L(Φ(x))(x−Φ(x)).
For any xy∈Yϵ where y=x−η∥∇L(x)∥∇L(x) is the one step Normalized GD update from x, we have
Moreover, we must have for every 1≤k≤M,
By Lemma B.16, we have ∂Φ(x)∇L(x)=0 for all x∈U. Thus we have
where the final step follows from using Definition B.4.
For the second claim, we have for every 1≤k≤M,
where the first step involves Theorem F.2.
The third claim follows from using Theorem F.4. Again,
where we borrow the bound on ∇2L(Φ(x))−∇2L(Φ(y)) from our previous calculations. The final step follows from the constants defined in Definition B.4. ∎
For any xy∈Yϵ where y=x−η∥∇L(x)∥∇L(x) is the one step Normalized GD update from x, we have that
Here θ=arctan∣⟨v1(x),x⟩∣PΦ(x),Γ(2:M)x, with x=∇2L(Φ(x))(x−Φ(x)). Additionally, we have that
By Taylor expansion for Φ at x, we have
where in the pre-final step, we used the property of Φ from Lemma B.16. In the final step, we have used a second order taylor expansion to bound the difference between ∂2Φ(x) and ∂2Φ(Φ(x)). Additionally, we have used y−x=η∥∇L(x)∥∇L(x) from the Normalized GD update rule.
Applying Taylor expansion on Φ again but at Φ(x), we have that
Also, at Φ(x), since v1(x) is the top eigenvector of the hessian ∇2L, we have that from Corollary B.23,
Plug Equations 24 and 23 into Equation 22, we have that
For the second claim, continuing from Equation 22, we have that
where Σ=PΦ(x),Γ∥∇L(x)∥∇L(x)(PΦ(x),Γ∥∇L(Φ(x))∥∇L(x))⊤ and the last step is by Lemma B.9. Here PΦ(x),Γ denotes the projection matrix of the subspace spanned by v1(x),…,vM(x).
By Lemmas B.22, B.18 and B.19, we have that
Let Lmin=miny∈UL(y). For any xy∈Yϵ where y=x−η∥∇L(x)∥∇L(x) is the one step Normalized GD update from x, if ∥∇L(xη(t))∥≥ζη, we have that
Thus for ∥∇L(xη(t))∥≥ζη, we have that
where the last step is because L is μ-PL on Yϵ. In other words, we have that
where in the last step we use L(y)−L(x)≤0. This completes the proof. ∎
B.2 Properties of limiting map of gradient flow, ΦΦ\Phi
where ζK denotes supx∈K∇2L(ϕ(x,t)). This implies that ∥∇L(ϕ(x,t))∥≤eζKT∥∇L(x)∥ and ∥ϕ(x,t)−x∥≤eζKT∥∇L(x)∥ for all t∈[0,T].
The following results Lemmas B.16, B.17, B.18, B.19, B.20, B.21 and B.22 are from [Li et al., 2022b].
For any x∈U, it holds that (1). ∂Φ(x)∇L(x)=0 and (2). ∂2Φ(x)[∇L(x),∇L(x)]=−∂Φ(x)∇2L(x)∇L(x).
We will also use the following two corollaries of Lemma B.22.
For any x∈Γ, let v1 be a top eigenvector of ∇2L(x), then
Simply note that L∇2L(x)−1(v1v1⊤)=2λ1(∇2L(x))1v1v1⊤ and apply Lemma B.22. ∎
For any x∈Γ, let v1 be the unit top eigenvector of ∇2L(x), then
The proof follows from using Corollary B.23 and the derivative of λ1 from Theorem F.1. ∎
As a variant of Corollary B.24, we have the following lemma.
For any x∈Γ, let v1 be the unit top eigenvector of ∇2L(x), then
Appendix C Analysis of Normalized GD on General Loss Functions
We restate the theorem concerning Phase I for the Normalized GD algorithm. Recall the following notation for each 1≤j≤M:
See 4.3 The intuition behind the above theorem is that for sufficiently small LR η, xη(t) will track the normalized gradient flow starting from xinit, which is a time-rescaled version of the standard gradient flow. Thus the normalized GF will enter Yϵ and so does normalized GD. Since L satisfies PL condition in Yϵ, the loss converges quickly and the iterate xη(t) gets η to manifold. To finish, we need the following theorem, which is the approximately-quadratic version of Lemma 3.3 when the iterate is O(η) close to the manifold.
Suppose {xη(t)}t≥0 are iterates of Normalized GD (4) with a learning rate η and xη(0)=xinit. There is a constant C>0, such that for any constant ς>1, if at some time t′, xη(t′)∈Yϵ and satisfies η∥xη(t′)−Φ(xη(t′))∥≤ς, then for all tˉ≥t′+Cμζςlogμςζ, the following must hold true for all 1≤j≤M:
provided that for all steps t∈{t′,…,tˉ−1}, xη(t)xη(t+1)⊂Yϵ.
The proof of the above theorem is in Section D.1.
Let Tx be the length of the GF trajectory starting from x, and we know limτ→Txϕ(x,τ)=Φ(x), where ϕ(x,τ) is defined as the Normalized gradient flow starting from x. In Lemmas B.5 and B.2 we show there is a small neighbourhood around Φ(xinit), Yϵ such that L is μ-PL in Yϵ. Thus we can take some time T0<Txinit such that ϕ(xinit,T0)∈Yϵ/2 and L(ϕ(xinit),T0)≤21Lcritical, where Lcritical:=8ϵ2μ. (Without loss of generality, we assume miny∈YL(y)=0) By standard ODE approximation theory, we know there is some small η0, such that for all η≤η0, xη(⌈T0/η⌉)−ϕ(xinit,T0)=O(η), where O(⋅) hides constants depending on the initialization xinit and the loss function L.
Without loss of generality, we can assume η0 is small enough such that xη(⌈T0/η⌉)∈Yϵ and L(xη(⌈T0/η⌉))≤Lcritical. Now let tη be the smallest integer (yet still larger than ⌈T0/η⌉) such that xη(tη)xη(tη−1)⊂Yϵ and we claim that there is t∈{⌈T0/η⌉,…,tη}, ∥∇L(xη(t))∥<ζη. By the definition of tη, we know for any t∈{⌈T0/η⌉+1,…,tη−1}, by Lemma B.11 we have ∥Φ(xη(t))−Φ(xη(t−1))∥≤ξη2, and by Lemma B.13, L(xη(t))−xη(t−1)≤−η42μ if ∥∇L(xη(t))∥≥ζη. If the claim is not true, since L(xη(t)) decreases η42μ per step, we have
which implies that tη−⌈T0/η⌉−1≤ηϵ, and therefore by Lemma B.11,
Meanwhile, by Lemma B.6, we have ∥Φ(xη(tη−1))−xη(tη−1)∥≤μ2L(xη(tη−1))≤μ2L(xη(⌈T0/η⌉))=2ϵ. Thus for any κ∈, we have ∥κxη(tη)+(1−κ)xη(tη−1)−Φ(xinit)∥ is upper bounded by
which is smaller than ϵ since we can set η0 sufficiently small. In other words, Φ(xη(tη))Φ(xη(tη−1))⊂Yϵ, which contradicts with the definition of tη. So far we have proved our claim that there is some tη′∈{⌈T0/η⌉,…,tη}, ∇L(xη(tη′))<ζη. Moreover, since L(xη(t)) decreases η42μ per step before tη′, we know tη′−⌈T0/η⌉≤ηϵ. By Lemma B.6, we know xη(tη′)−Φ(xη(tη′))≤μζη.
Now we claim that for any T1′, there is some sufficiently small threshold η0, tη≥ηT1′+1 if η≤η0. Below we prove this claim by contradiction. If the claim is not true, that is, tη<ηT1′+1. if tη≤Cμζςlogμςζ+tη′ with ς=μζ, we know ∥xη(tη)−Φ(xinit)∥≤xη(tη)−xη(tη′)+xη(tη′)−Φ(xη(tη′))+Φ(xη(tη′))−Φ(xinit)=O(η), which implies that xη(tη)xη(tη−1)∈Y. If tη≥Cμζςlogμςζ+tη′, by Lemma C.1, we have ∥xη(tη)−Φ(xη(tη))∥=O(η). By Lemma B.11, we have ∥Φ(xη(tη))−Φ(xη(⌈T0/η⌉))∥≤O(η). Thus again we have that ∥xη(tη)−Φ(xinit)∥≤∥xη(tη)−Φ(xη(tη))∥+∥Φ(xη(tη))−Φ(xη(⌈T0/η⌉))∥+∥Φ(xη(⌈T0/η⌉))−Φ(xinit)∥=O(η), which implies that xη(tη)xη(tη−1)∈Y. In both cases, the implication is in contradiction to the definition of tη.
Thus for any T1′, tη≥ηT1′+1 for sufficiently small threshold η0 and η≤η0. To complete the proof of Theorem 4.3, we pick T1 to be any real number strictly larger than ϵ+T0, as ηT1>Cμζςlogμςζ+ηϵ+⌈T0/η⌉≥Cμζςlogμςζ+tη′ when η is sufficiently small with ς=μζ. By Lemma C.1 the second claim of Theorem 4.3 is proved. Using the same argument again, we know ∀ηT1≤t≤ηT1′, it holds that ∥Φ(xη(t))−Φ(xinit)∥≤O(η). ∎
C.2 Phase II, Limiting Flow
We first restate the main theorem that demonstrates that the trajectory implicitly minimizes sharpness. See 4.4
where v′(τ′+0):=limδ→0δv(τ′+δ)−v(τ′) is the right time derivative of v at τ′.
To prove the first claim, we first show the movement in the manifold for the discrete trajectory for Algorithm 1 by Lemma B.12: for each step t, provided Φ(xη(t))Φ(xη(t+1))∈Yϵ, it holds that
The high-level idea for the proof of the first claim is to bound the gap between Equation 26 and Equation 5 using Theorem C.2. And the first claim eventually boils down to upper bound the average angle by O(η), which is exactly the second claim.
Formally, let t2 be the largest integer no larger than ⌊T2/η2⌋ such that for any 0≤t≤t2, it holds that Φ(xη(t))Φ(xη(t+1))∈Yϵ.
Since we started from a point that has max1≤j≤MRj(xη(0))≤O(η2), we have from Lemma C.1, that the iterate satisfies the condition max1≤j≤MRj(xη(t))≤O(η2) at step t as well, meaning that ∥xη(t)−Φ(xη(t))∥≤O(η).
Therefore, for any τ≤t2η2, note that v′(τ+0)=v′(⌊τ/η2⌋+0) and that f(v(⌊τ/η2⌋+0))−f(v(τ))=O(Φ(xη(⌊τ/η2⌋+1))−Φ(xη(⌊τ/η2⌋)))=O(η2), we have that
where in the last step we use the second claim. This implies that t2 must be equal to ⌊T2/η2⌋ for sufficiently small η otherwise xη(t2)xη(t2+1)⊆Yϵ. This is because ∥xη(t2+1)−xη(t2)∥=O(η) and X(t2η2)∈Y. The proof is completed by noting that X(T2)−X(⌊T2/η2⌋)=O(η2). ∎
Appendix D Phase I, Omitted Proofs of the Main Lemmas
The Normalized GD update at any step t can be written as (from Lemma B.9)
From Lemma B.11, we have ∥Φ(xη(t))−Φ(xη(t+1))∥≤O(ξη2), which further implies, ∇2L(Φ(xη(t+1)))−∇2L(Φ(xη(t)))≤O(νξη2). Thus, using the notation x=∇2L(Φ(x))(x−Φ(x)), we have
Below we will show that ∥xη(t)−Φ(xη(t))∥≤O(η), and thus the trajectory of xη is similar to the trajectory in the qudratic model with an O(η2) error, with the hessian fixed at ∇2L(Φ(xη(t))), and hence we can apply the same techniques from Corollary A.4 and Lemma A.1.
First, we consider the norm of the vector xη(t) for t′+1≤t≤t. We will show the following induction hypothesis:
Base case: (t=t′). We have ∥xη(t′)∥=∇2L(Φ(xη(t′)))[xη(t′)−Φ(xη(t′))]≤ηλ1(t)ς≤ηζς.
Induction case:(t>t′). Suppose the hypothesis holds true for t−1. Then,
If ∥xη(t−1)∥≥ηλ1(t). We can directly apply Corollary A.3 on (28) to show that
where the final step follows if η is sufficiently small. Hence, ∥xη(t)∥<∥xη(t−1)∥≤ηζς.
If ∥xη(t−1)∥≤ηλ1(t). Then, we can directly apply Lemma A.1 on (28) to show that
Hence, we have shown that, ∥xη(t)−Φ(xη(t))∥≤λM(t)1∥xη(t)∥≤μ1.01ηςζ for all time t′≤t≤t.
We complete the proof of Lemma C.1 with a similar argument as that for the quadratic model (see Corollary A.4 and Lemma A.1). The major difference from the quadratic model is that here the hessian changes over time, along with its eigenvectors and eigenvalues. Hence, we need to take care of the errors introduced in each step by the change of hessian.
The high-level idea is to divide the eigenvalues at each step t into groups such that eigenvalues in the same group are O(η) close and eigenvalues from different groups are at least 2η far away from each other. Formally, we divide [M] into disjoint subsets S1(t),⋯,Sp(t)(t) (with 1≤p(t)≤M) such that
Thus for any t′≤t≤t−1 and k∈[p(t)], suppose i∈Sk(t) and j=minSk(t), we have that
and ηλi(t+1)≥ηλi(t)−O(η2)≥ηλj(t)−O(η2).
Next we will use the results from the quadratic case to upper bound ∑h=jM⟨vh(t),xη(t+1)⟩2 using ∑h=jM⟨vh(t),xη(t)⟩2. For all 1≤j≤M, we consider the following two cases for any time t′+1≤t≤t:
If ∑h=jM⟨vh(t),xη(t)⟩2>ηλj(t), then we can apply Lemma A.2 on (28) to show that
If ∑h=jM⟨vh(t),xη(t)⟩2≤ηλj(t), then we can apply Lemma A.1 on (28) to show that
and therefore following the same proof of quadratic case Corollary A.4, for t≥t′+Ω(μςζlogμζς), it holds that ∀j∈[M], ∑i=jM⟨vi(tˉ),x(tˉ)⟩2≤ηλj(tˉ)+O(η2). ∎
D.2 Properties of the condition in Equation 25
By Lemma C.1, the following condition will continue to hold true for all 1≤j≤M before xη(t) leaves Yϵ:
where xη(t)=∇2L(Φ(xη(t)))(xη(t)−Φ(xη(t))). We will call the above condition as the alignment condition from now onwards.
From the alignment condition (29), we can derive the following property that continues to hold true throughout the trajectory, once the condition is satisfied:
The proof follows from using the noisy quadratic update for Normalized GD in Lemma B.9 (Equation 21) and the behavior in a quadratic model along the non-top eigenvectors in Lemma A.5.
Appendix E Phase II, Omitted Proofs of the Main Lemmas
The main lemma in this section is Lemma E.5 in Section E.2, which says the sum of the angles across the entire trajectory in any interval [0,t2] with t2=Ω(1/η2), is at most O(ηt2). Before proving the main lemma, we will first recap and introduce some notations that will be used.
In Phase II, we start from a point xη(0), such that (1) ∥xη(0)−Φ(xinit)∥≤O(η), (2) maxj∈[D]Rj(xη(t))≤O(η2), and additionally (3) ∣⟨v1(xη(0)),xη(0)−Φ(xη(0))⟩∣=Ω(η).
The condition (29) that was shown to hold true in Phase II is:
Further, we had proved in Lemma D.1 that:
The properties of the three sets N0,N1,N2 in an interval (t,t) given directly by the algorithm in Algorithm 3 include:
∀t,t∈N1⟺t+1∈N2.
N0∪N1∪N2=[t,t], and the intersection between each pair of them is empty.
We also have the following lemmas, which is less direct:
For any step t in N0, t−1,t+2∈N2 and t−2,t+1∈N1
tanθt+1≤(1−ζmin(Δ,2μ))tanθt+O(Gtη2)
tanθt+2≤Gtηλ1tanθt+O(Gtη2)
As a direct consequence of Lemma E.3, we have the following lemma:
Given any t with θt=Ω(1), let t=maxN1∩{t∣t≤t}. If Gt≥Ω(η), then θt=Ω(1).
The claim is clearly true if t∈N1. If t∈N0, then Lemma E.1 shows that t−1∈N2,t−2∈N1 and thus t=t−2. The claim is true because of the second property of Lemma E.3. If t∈N2, then t=t−1∈N1 and the proof is completed by applying the first property of Lemma E.3. ∎
E.2 Time Average of Angles Against Top Eigenspace
provided η is set sufficiently small, and for all time 0≤t≤t2−1, xη(t)xη(t+1)⊂Yϵ.
We analyze the behavior of a general ti when it falls in any of the above cases:
Case (B).
From Lemma E.6 we have that the sum of angle over this time is
Case (C).
Now it remains to upper bound the number of occurrence of (A),(B) and (C). Since our goal is to show average angle is O(η), which is equal to the average angle in case (A), so the number of occurrence of case (A) doesn’t matter. For case (B), if it is followed by case (A), then there is an Ω(1/η2) gap before next occurrence of (B). If (B) is followed by case (C), then by Lemma E.10, it takes at least Ω(1/η2) steps to escape from (B). Thus we can have O(1) occurrence of case (B). For the same reason, there could be at most O(1) occurrence of case (C).
All in all, with probability at least 1−O(η12⋅η12)=1−O(η10), we must have
where we use t2≥Ω(η21) in the last step and and the number of occurrence of case (B) is O(1) in the second to the last step. ∎
The noisy update rule for Normalized GD, as derived in Lemma B.9, which says that the Normalized GD update is very close to the update in a quadratic model with an additional O(η2) error. Keeping this in mind, we then divide our trajectory in the interval (t,t′) as per Algorithm 3 into three subsets N0,N1,N2. (Please see Section E.1 for a summary on the properties of these 3 sets.)
Consider any t∈N1. Using the behavior of Gt from Lemma E.10, we can show that in each of the time-frames, Gt+2≥(1+Ω(sin2θt))Gt−O(η2(η+ηt))≥Gt+Ω(θt2η)−O(η2(η+θt)).
Next we want to telescope over Gt+2−Gt to get an upper bound for ∑t∈N1θt. If t+2 is also in N1 then it’s fine. If t+2∈N0, then t+3∈N1 by Lemma E.1 and we proceed in the following two cases.
If θt+2≤C for some sufficiently small constant C, since Gt+2≤λ1(t+2)η/2−Ω(η), we have ∥xη(t+2)∥≤cosθt+2Gt+2=λ1(t+2)η/2−Ω(η), and thus by Lemma E.9, we have Gt+3≤Gt+2 and therefore, Gt+3≥Gt+Ω(θt2η)−O(η2(η+θt)).
If θt+2≥C, then by Lemma D.2, we have θt=Ω(1), thus Gt+2≥Gt+Ω(η) by Lemma E.10. Again by Lemma E.9, we have Gt+3≥Gt+2−O(η2). Thus again we conclude Gt+3≥Gt+Ω(η)≥Gt+Ω(θt2η)−O(η2(η+θt)), since θt is always O(1).
Since total increase in Gt during this interval can is most O(η), we conclude that ∑t∈N1θt2=O(1)+η∑t∈N1(η+θt) and thus it holds that
Moreover, by Lemma E.3, we must have θt<θt−1+O(η) for any time t∈N2, and t−1 must be in N1. By Lemma E.9, we have θt≤Ω(θt−2) for any t∈N0 and t−2 must be in N2. That implies,
Consider any coordinate 2≤k≤M. For any constants 0<β, there is some constant α>0 such that for any time step t where xη(t) is in Yϵ, Gt≥βη, condition (29) holds and ⟨vk(t),x(t)⟩≥αη2, then there is some time t≤t+O(ln1/η) such that if for all time t≤t′<t, xη(t′)xη(t′+1)⊂Yϵ, then condition (29) holds at time t and at least one of the following two conditions hold:
Gt≥0.99gt(λk(t)).
We will prove by contradiction. Suppose neither of the two condition happens, we will show θt grows exponentially and thus the condition (2) must be false in O(log1/η) steps.
Now, we can use Equation 21 (Lemma B.9) to show that the Normalized GD update is equivalent to update in quadratic model, up to an additional O(η2) error.
Similar to Lemma A.9, consider the coordinate k, we have that
The third step follows from using the same argument as the one used for the quadratic update in Lemma A.9 and the assumption that Gt≥0.99gt(λk(t)). The final step holds true because we can pick α as a large enough constant and by assumption ∣⟨v1(t),xη(t)⟩∣∣⟨vk(t),xη(t)⟩∣≥αη.
We then bound vk(t)−vk(t) and Φ(xη(t)−Φ(xη(t)) by O(η2(t−t)) using Lemma B.11. Combining everything, we conclude that at least one of the two assumptions has to break for some t≤t+O(log1/η). ∎
Consider any coordinate 2≤k≤M. For any constants 0<β, suppose at time step t, xη(t) is in Yϵ, (1.01)gt(λk(t))η≤Gt<0.5ηλ1(t) and condition (29) holds, then there is some time t≤t+O(ln1/η) such that if for all time t≤t′<t, xη(t′)xη(t′+1)⊂Yϵ, then the following two conditions hold:
The proof of Lemma E.8 is very similar to the proof of Lemma E.7 and thus we omit the proof. The only difference will be that we need to use Lemma A.10 in place of Lemma A.9, when we use the result for the quadratic model.
E.3 Dynamics in the Top Eigenspace
Here, we will state two important lemmas that we used for the proof of Lemma E.5, which is about the behavior of the iterate along the top eigenvector. Lemma E.10 can be viewed as perturbed version for Lemma A.7 in the quadratic case, and We assume in all the lemmas, that Equation 29 holds true for the time under consideration, which we showed in Lemma C.1, and also that we start Phase II from a point where the alignment along the top eigenvector is non negligible.
The following lemmas give the properties of dynamics in the top eigenspace in Phase II for one-step and two-step updates respectively. Recall we use Gt to denote the quantity ∣⟨v1(t),x(t)⟩∣.
provided that xη(t)xη(t+1),xη(t+1)xη(t+2)⊂Yϵ.
provided that xη(t)xη(t+1),xη(t+1)xη(t+2)⊂Yϵ.
First note that ∠(xη(t),∇L(xη(t)))=O(∥xη(t)∥∥xη(t)−∇L(xη(t))∥)=O(Gtη2)=O(η), where the last step we use Lemma B.9. Let δ=∠(v1(t),∇L(xη(t)))−∠(v1(t),xη(t)) and we have ∣δ∣≤∠(xη(t),∇L(xη(t)))=O(η). Therefore, it holds that
From Lemma B.11, we have ∥Φ(xη(t))−Φ(xη(t+1))∥≤O(ξη2), which further implies, ∇2L(Φ(xη(t+1)))−∇2L(Φ(xη(t)))≤O(η2). Thus, we can use Theorem F.4 to have ∥v1(t)−v1(t+1)∥≤O(λ1(t)−λ2(t)νξη2)=O(η2). From Lemma B.12, we have ∣⟨v1(t),Φ(xη(t+1))−Φ(xη(t))⟩∣≤O(η3). Thus we have that
Therefore, we have the following inequality by applying the same argument above to t+1:
By Lemma D.1, we know ∥x(t+1)∥=Ω(η). By Lemma D.2, we know that θt+1≤θt+O(η). Thus
Next we will show ∥xη(t+1)∥−(I−∥xη(t)∥η∇2L(Φ(xη(t))))xη(t)=O(η2θt). For convenience, we denote ∇2L(Φ(xη(t))) by H. First we have that
where α is the angle between ∥∇L(xη(t))∥∇L(xη(t))−∥xη(t)∥xη(t) and 2Hxη(t)−ηH2(∥∇L(xη(t))∥∇L(xη(t))+∥xη(t)∥xη(t)). Note that and that both ∠(xη(t),v1(t)),∠(∇L(xη(t)),v1(t))=O(ηt+η), we have that the angle between ∥∇L(xη(t))∥∇L(xη(t))+∥xη(t)∥xη(t) and 2Hxη(t)−ηH2(∥∇L(xη(t))∥∇L(xη(t))+∥xη(t)∥xη(t)) is at most O(ηt+η). Further note that ∥∇L(xη(t))∥∇L(xη(t))−∥xη(t)∥xη(t) is perpendicular to ∥∇L(xη(t))∥∇L(xη(t))+∥xη(t)∥xη(t), we know cosα≤O(θt+η). Therefore we have that
Thus we have proved a perturbed version of Lemma A.6, that is,
Therefore a perturbed version of Lemma A.7 would give us:
The proof of the first inequality is completed by plugging the above equation into Equation 32.
The second inequality is immediate by noting that ηθt2+C2η3≥2Cη2θt for any C>0. ∎
E.4 Dynamics in Top Eigenspace When Dropping Below Threshold
Denote r=η100. For any constant 0<β, there is a constant α>0, such that for any step t and xη(t)∈Yϵ with the following conditions hold:
∣⟨vi(t),xη(t)⟩∣≤O(η2), for all 2≤i≤M.
Lemma E.11 is a direct consequence of the following lemma.
where Pt,Γ(2:M) denotes the subspace spanned by v2(t),…,vM(t).
We will first prove Lemma E.11 and then we turn to the proof of Lemma E.12.
For convenience, we denote ∇2L(Φ(xη(t)))[u(t)−Φ(xη(t))] by u(t) and ∇2L(Φ(xη(t)))[w(t)−Φ(xη(t))] by w(t). Suppose both Pt,Γ(2:M)u(t),Pt,Γ(2:M)v(t) are O(η2), we will show the following, which indicates the contradiction:
An important claim to note is the following:
Note that the condition has been slightly changed to use {vi(t)} as reference coordinate system and Φ(xη(t)) as reference point. The above lemma follows from the fact that both u(0) and w(0) are r-close to xη(t), which itself satisfies the alignment condition (Equation 29). Thus, both u(0) and w(0) initially follow the desired condition. Since, both the trajectories follow Normalized GD updates, the proof will follow from applying the same technique used in the proof of Lemma C.1. Another result to keep in mind is the following modified version of Corollary D.3, Lemma E.14.
The above lemma uses {vi(t)} as reference coordinate system and Φ(xη(t)) as reference point. The above lemma follows from showcasing Normalized GD updates of u(t) and w(t) as equivalent to the update in a quadratic model, with an additional noise of O(μνζη2), similar to Equation 28.
Further, applying the same technique from Lemma E.10, we can show that
Initially, because u was initialized close to xη(t), we must have
Now, we use taylor expansion of F around u(t) to get
Using taylor expansion: ∇L(z(γ))=∇2L(Φ(xη(t)))(z(γ)−Φ(xη(t)))+O(ν∥z(γ)−Φ(xη(t))∥2) and hence, we must have ∥∇L(z(γ))∥≥Ω(η).
Now we define Bt and claim At can be approximated as below with ∥Bt∥=O(η). Furthermore, ∥At∥≤O(1).
with At+1≤O(1) and Bt+1≤O(η).
Finally, we use Lemma E.15 and Lemma E.16 to handle the main and error terms in Equation 41,
which completes the proof of Lemma E.12. ∎
For simplicity of presentation, we have used Mt to define
Thus, the term under consideration can be simplified as follows,
We first recall the definition of the error term:
By Equation 37, the following property holds for all t≥0 for some ϱ>0
First, by induction hypothesis at t−2, we know
The proof is completed by picking C large enough such that ϱφC+O(η)≤C. ∎
E.5 Proof for Operating on Edge of Stability
According to the proof of Theorem 4.4, we know for all t, it holds that Rj(xη(t))≤O(η2). Thus SL(xη(t),ηt)=ηt⋅sup0≤s≤ηtλ1(∇2L(xη(t)−s∇L(xη(t))))=ηt(λ1(t)+O(η)), which implies that [SL(xη(t),ηt)]−1=ηλ1(t)∥∇L(xη(t))∥+O(η)=ηλ1(t)∥xη(t)∥+O(η). The proof for the first claim is completed by noting that η1(∥xη(t)∥+∥xη(t+1)∥)=λ1(t)+O(η+θt) as an analog of the quadratic case.
For the second claim, it’s easy to check that L(xη(t))=2λ1(t)∥xη(t)∥+O(ηθt). Thus have L(xη(t))+L(xη(t+1))=2λ1(t)∥xη(t)∥+2λ1(t+1)∥xη(t+1)∥+O(η(θt+θt+1)). Note that λ1(t)−λ1(t+1)=O(η2) and θt+1=O(θt), we conclude that L(xη(t))+L(xη(t+1))=η2λ1(∇2L(xη(t)))+O(ηθt). ∎
Appendix F Some Useful Lemmas About Eigenvalues and Eigenvectors
Moreover, the functions λ and u are C∞ on N(X0) and the differentials at X0 are
The next theorem is the Davis-Kahan sin(θ) theorem, that bounds the change in the eigenvectors of a matrix on perturbation. Before presenting the theorem, we need to define the notion of unitary invariant norms. Examples of such norms include the frobenius norm and the spectral norm.
Appendix G Analysis of L𝐿\sqrt{L}
The analysis will follow the same line of proof used for the analysis of Normalized GD. Hence, we write down the main lemmas that are different from the analysis of Normalized GD. Rest of the lemmas are nearly the same and hence, we have omitted them.
The major difference between the results of Normalized GD and GD with L is in the behavior along the manifold Γ (for comparison, see Lemma B.12 for Normalized GD and Lemma G.10 for GD with L). Another difference between the results of Normalized GD and GD with L is in the error rates mentioned in Theorem 4.4 and Theorem 4.6. The difference comes from the stronger behavior of the projection along the top eigenvector that we showed for Normalized GD in Lemma E.10, but doesn’t hold for GD with L (see Lemma G.6). This difference shows up in the sum of angles across the trajectory (for comparison, see Lemma E.5 for Normalized GD and Lemma G.4 for GD with L), and is finally reflected in the error rates.
The notations will be the same as Appendix B . However, here we will use xη(t) to denote (2∇2L(Φ(xη(t))))1/2(xη(t)−Φ(xη(t))). We will now denote Y as the limiting flow given by Equation 7.
G.2 Phase I, convergence
Here, we will show a very similar stability condition for the GD update on L as the one (Lemma C.1) derived for Normalized GD. Recall our notation xη(t)=2∇2L(Φ(xη(t)))(xη(t)−Φ(xη(t))).
Suppose {xη(t)}t≥0 are iterates of GD with L (6) with a learning rate η and xη(0)=xinit. There is a constant C>0, such that for any constant ς>0, if at some time t′, xη(t′)∈Yϵ and satisfies η∥xη(t′)−Φ(xη(t′))∥≤ς, then for all tˉ≥t′+Cμζςlogμςζ, the following must hold true for all 1≤j≤M:
provided that for all steps t∈{t,…,tˉ−1}, xη(t)xη(t+1)⊂Yϵ.
The proof exactly follows the strategy used in Lemma C.1. We can use the noisy update formulation from Lemma G.7 and the bound on the movement in Φ from Lemma G.10 to get for any time t with tˉ≥t≥t′ (similar to Equation 28):
Hence, similar to Lemma D.1, we can derive the following property that continues to hold true throughout the trajectory, once the condition Equation 42 is satisfied:
We also have the counterpart of Corollary D.3 with the same proof, which follows from using the noisy update of GD on L from Lemma G.7 and using the quadratic update result from Lemma A.5.
G.3 Phase II, limiting flow
Let T2 be the time up until which solution to the limiting flow exists.
Lemma G.10 shows the movement in Φ, which can be informally given as follows: in each step t,
provided Φ(xη(t))Φ(xη(t+1))∈Yϵ.
Motivated by this update rule, we show that the trajectory of Φ(xη(⋅)) is close to the limiting flow in Equation 7, for a small enough learning rate η. The major difference from Theorem 4.4 comes from the fact that the total error introduced in Equation 43 over an interval [0,t2] is ∑t=0t2O(η2θt+η3), which is of the order O(η1/2) using the result of Lemma G.4.
The first lemma shows that the sum of the angles in an interval [0,t2] of length Ω(1/η2) is at most O(t2η1/2).
provided η is sufficiently small and for all time 0≤t≤t2−1, xη(t)xη(t+1)⊂Yϵ.
The proof is very similar to the proof of Lemma E.5, except we replace Lemma E.6 by Lemma G.5 in the analysis of case (B). Hence the final average angle becomes O(η). ∎
The proof will follow exactly as Lemma E.6, except we replace Lemma E.10 by Lemma G.6, which changes the rate into O(t′−t+(t′−t)η) ∎
provided Gt≥Ω(η) and xη(t)xη(t+1),xη(t+1)xη(t+2)⊂Yϵ.
Here, we will follow a much simpler approach than Lemma E.10 to have a weaker error bound. The stronger error bounds in Lemma E.10 were due to the very specific update rule of Normalized GD.
First recall xη(t)=2∇2L(Φ(x))(xη(t)−Φ(xη(t))). By Lemma G.10, we have ∥Φ(xη(t+1))−Φ(xη(t))∥=O(η2), thus xη(t+1)−xη(t)=2∇2L(Φ(x))(xη(t+1)−xη(t))=η2∇2L(Φ(x))2L(xη(t))∇L(xη(t)). From Lemma G.7, we have
where we have used the fact that ∥xη(t)∥=O(η).
Hence, the update is similar to the update in a quadratic model, with ∇2L(Φ(xη(t))) guiding the updates with an additional O(η2) perturbation. As a result we also get a O(η2) perturbation in Gt. Here we use the assumption Gt=Ω(η) so that GD updates are O(1)-lipschitz. ∎
G.4 Omitted Proof for Operating on Edge of Stability
This proof is similar to that of Theorem 4.7.
If M=1, that is, the dimension of manifold Γ is D−1, we know xη(t)xη(t+1) will cross Γ, making the ∇2L diverges at the intersection and the first claim becomes trivial. If M≥2, we have ∇2L=4L32L∇2L−∇L∇L⊤ diverges at the rate of ∥∇L∥1. It turns out that using basic geometry, one can show that the distance from Φ(xη(t)) to xη(t)xη(t+1) is O(η(θt+θt+1)), thus sup0≤s≤ηλ1(∇2L(xη(t)−s∇L(xη(t))))=Ω(η(θt+θt+1)1). The proof of the first claim is completed by noting that θt+1=O(θt).
For the second claim, it’s easy to check that L(xη(t))=∥xη(t)∥+O(η). The proof for the first claim is completed by noting that ∥xη(t)∥+∥xη(t+1)∥=ηλ1(t)+O(η+θt) as an analog of the quadratic case. ∎
G.5 Geometric Lemmas for L𝐿\sqrt{L}
First recall our notations, x=2∇2L(Φ(x))(x−Φ(x)) and θ=arctan∣⟨v1(x),x⟩∣PΦ(x),Γ(2:M)x.
At any point x∈Yϵ, we have
Since Φ(x) is a local minimizer of zero loss, we have ∇L(Φ(x))=0, we have that
By Lemma B.8, we know ∂2L(Φ(x))[x−Φ(x),x−Φ(x)]≥Ω(μ∥x−Φ(x)∥2) and therefore
For the second claim, with x=2∇2L(Φ(x))(x−Φ(x)), we have that
By Lemma B.7, we have ∥x−Φ(x)∥≤ζν2μ, thus μζ1/2ν∥x−Φ(x)∥=O(ζ1/2μ)=O(ζ1/2). ∎
The following two lemmas are direct implications of Lemma G.7.
At any point x∈Yϵ, we have
Consider any point x∈Yϵ. Then,
where θ=arctan∣⟨v1(x),x⟩∣PΦ(x),Γ(2:M)x, with x=2∇2L(Φ(x))(x−Φ(x)).
For any xy∈Yϵ where y=x−η∇L(x) is the one step update on L loss from x, we have
Here θ=arctan∣⟨v1(x),x⟩∣PΦ(x),Γ(2:M)x, with x=2∇2L(Φ(x))(x−Φ(x)).
We outline the major difference from the proof of Lemma B.12. Using taylor expansion for the function Φ, we have
where in the final step, we used the property of Φ from Lemma B.16 to kill the first term and use the bound on L(x)∇L(x) from Lemma G.7 for the third term.
Since the function Φ∈C3, hence ∂2Φ(x)=∂2Φ(Φ(x))+O(χ∥x−Φ(x)∥).
Also, at Φ(x), since v1(x) is the top eigenvector of the hessian ∇2L, we have from Corollary B.23,
where recall our notation of θ=arctan∣⟨v1(x),x−Φ(x)⟩∣PΦ(x),Γ(2:M)(x−Φ(x)).
With further simplification, it turns out that
The proof is completed by using Corollary B.25. ∎
Appendix H Additional Experimental Details
For running GD on L, we start from (x,y)=(14.7,3.), and use a learning rate η=0.5. For running Normalized GD on L, we start from (x,y)=(14.7,−3), and use a learning rate η=5.
For Figure 2:
We start Normalized GD from ⟨v1,x(0)⟩=10−4,⟨v2,x(0)⟩=0.45. We use a learning rate of 1 for the optimization updates.
H.2 Implementation Details for Simulation for the Limiting Flow of Normalized GD
We provide the code for running a single step of the riemannian flow (5) corresponding to Normalized GD. The pseudocode is given in Algorithm 4.
Riemannian gradient update details:
Each update comprises of three major steps: a) computing ∇3L(x)[v1(x),v1(x)], b) a projection onto the tangent space of the manifold, and c) few steps of gradient descent with small learning rate to drop back to manifold.