Surrogate Gap Minimization Improves Sharpness-Aware Training

Juntang Zhuang, Boqing Gong, Liangzhe Yuan, Yin Cui, Hartwig Adam, Nicha Dvornek, Sekhar Tatikonda, James Duncan, Ting Liu

Introduction

Modern neural networks are typically highly over-parameterized and easy to overfit to training data, yet the generalization performances on unseen data (test set) often suffer a gap from the training performance (Zhang et al., 2017a). Many studies try to understand the generalization of machine learning models, including the Bayesian perspective (McAllester, 1999; Neyshabur et al., 2017), the information perspective (Liang et al., 2019), the loss surface geometry perspective (Hochreiter & Schmidhuber, 1995; Jiang et al., 2019) and the kernel perspective (Jacot et al., 2018; Wei et al., 2019). Besides analyzing the properties of a model after training, some works study the influence of training and the optimization process, such as the implicit regularization of stochastic gradient descent (SGD) (Bottou, 2010; Zhou et al., 2020), the learning rate’s regularization effect (Li et al., 2019), and the influence of the batch size (Keskar et al., 2016).

These studies have led to various modifications to the training process to improve generalization. Keskar & Socher (2017) proposed to use Adam in early training phases for fast convergence and then switch to SGD in late phases for better generalization. Izmailov et al. (2018) proposed to average weights to achieve a wider local minimum, which is expected to generalize better than sharp minima. A similar idea was later used in Lookahead (Zhang et al., 2019). Entropy-SGD (Chaudhari et al., 2019) derived the gradient of local entropy to avoid solutions in sharp valleys. Entropy-SGD has a nested Langevin iteration, inducing much higher computation costs than vanilla training.

The recently proposed Sharpness-Aware Minimization (SAM) (Foret et al., 2020) is a generic training scheme that improves generalization and has been shown especially effective for Vision Transformers (Dosovitskiy et al., 2020) when large-scale pre-training is unavailable (Chen et al., 2021). Suppose vanilla training minimizes loss f(w)f(w) (e.g., the cross-entropy loss for classification), where ww is the parameter. SAM minimizes a perturbed loss defined as fp(w)≜max⁡∣∣δ∣∣≤ρf(w+δ)f_{p}(w)\triangleq\operatorname{max}_{||\delta||\leq\rho}f(w+\delta), which is the maximum loss within radius ρ\rho centered at the model parameter ww. Intuitively, vanilla training seeks a single point with a low loss, while SAM searches for a neighborhood within which the maximum loss is low. However, we show that a low perturbed loss fpf_{p} could appear in both flat and sharp minima, implying that only minimizing fpf_{p} is not always sharpness-aware.

Although the perturbed loss fp(w)f_{p}(w) might disagree with sharpness, we find a surrogate gap defined as h(w)≜fp(w)−f(w)h(w)\triangleq f_{p}(w)-f(w) agrees with sharpness — Lemma 3.3 shows that the surrogate gap hh is an equivalent measure of the dominant eigenvalue of Hessian at a local minimum. Inspired by this observation, we propose the Surrogate Gap Guided Sharpness Aware Minimization (GSAM) which jointly minimizes the perturbed loss fpf_{p} and the surrogate gap hh: a low perturbed loss fpf_{p} indicates a low training loss within the neighborhood, and a small surrogate gap hh avoids solutions in sharp valleys and hence narrows the generalization gap between training and test performances (Thm. 5.3). When both criteria are satisfied, we find a generalizable model with good performances.

GSAM consists of two steps for each update: 1) descend gradient ∇fp(w)\nabla f_{p}(w) to minimize the perturbed loss fpf_{p} (this step is exactly the same as SAM), and 2) decompose gradient ∇f(w)\nabla f(w) of the original loss f(w)f(w) into components that are parallel and orthogonal to ∇fp(w)\nabla f_{p}(w), i.e., ∇f(w)=∇∥f(w)+∇⊥f(w)\nabla f(w)=\nabla_{\parallel}f(w)+\nabla_{\perp}f(w), and perform an ascent step in ∇⊥f(w)\nabla_{\perp}f(w) to minimize the surrogate gap h(w)h(w). Note that this ascent step does not change the perturbed loss fpf_{p} because ∇f⊥(w)⊥∇fp(w)\nabla f_{\perp}(w)\perp\nabla f_{p}(w) by construction.

We summarize our contribution as follows:

We define surrogate gap, which measures the sharpness at local minima and is easy to compute.

We propose the GSAM method to improve the generalization of neural networks. GSAM is widely applicable and incurs negligible computation overhead compared to SAM.

We demonstrate the convergence of GSAM and its provably better generalization than SAM.

We empirically validate GSAM over image classification tasks with various neural architectures, including ResNets (He et al., 2016), Vision Transformers (Dosovitskiy et al., 2020), and MLP-Mixers (Tolstikhin et al., 2021).

Preliminaries

wtadv≜wt+ρt∇f(wt)∣∣∇f(wt)∣∣+ϵw_{t}^{adv}\triangleq w_{t}+\rho_{t}\frac{\nabla f(w_{t})}{||\nabla f(w_{t})||+\epsilon}: The solution to max⁡∣∣w′−wt∣∣≤ρtf(w′)\operatorname{max}_{||w^{\prime}-w_{t}||\leq\rho_{t}}f(w^{\prime}) when ρt\rho_{t} is small.

fp(wt)≜max⁡∣∣δ∣∣≤ρtf(wt+δ)≈f(wtadv)f_{p}(w_{t})\triangleq\operatorname{max}_{||\delta||\leq\rho_{t}}f(w_{t}+\delta)\approx f(w_{t}^{adv}): The perturbed loss induced by f(wt)f(w_{t}). For each wtw_{t}, fp(wt)f_{p}(w_{t}) returns the worst possible loss ff within a ball of radius ρt\rho_{t} centered at wtw_{t}. When ρt\rho_{t} is small, by Taylor expansion, the solution to the maximization problem is equivalent to a gradient ascent from wtw_{t} to wtadvw_{t}^{adv}.

h(w)≜fp(w)−f(w)h(w)\triangleq f_{p}(w)-f(w): The surrogate gap defined as the difference between fp(w)f_{p}(w) and f(w)f(w).

∇f(wt)=∇f∥(wt)+∇f⊥(wt)\nabla f(w_{t})=\nabla f_{\parallel}(w_{t})+\nabla f_{\perp}(w_{t}): Decompose ∇f(wt)\nabla f(w_{t}) into parallel component ∇f∥(wt)\nabla f_{\parallel}(w_{t}) and vertical component ∇f⊥(wt)\nabla f_{\perp}(w_{t}) by projection ∇f(wt)\nabla f(w_{t}) onto ∇fp(wt)\nabla f_{p}(w_{t}).

2 Sharpness-Aware Minimization

Conventional optimization of neural networks typically minimizes the training loss f(w)f(w) by gradient descent w.r.t. ∇f(w)\nabla f(w) and searches for a single point ww with a low loss. However, this vanilla training often falls into a sharp valley of the loss surface, resulting in inferior generalization performance (Chaudhari et al., 2019). Instead of searching for a single point solution, SAM seeks a region with low losses so that small perturbation to the model weights does not cause significant performance degradation. SAM formulates the problem as:

where ρ\rho is a predefined constant controlling the radius of a neighborhood. This perturbed loss fpf_{p} induced by f(w)f(w) is the maximum loss within the neighborhood. When the perturbed loss is minimized, the neighborhood corresponds to low losses (below the perturbed loss). For a small ρ\rho, using Taylor expansion around ww, the inner maximization in Eq. 1 turns into a linear constrained optimization with solution

As a result, the optimization problem of SAM reduces to

where ϵ\epsilon is a scalar (default: 1e-12) to avoid division by 0, and wadvw^{adv} is the “perturbed weight” with the highest loss within the neighborhood. Equivalently, SAM seeks a solution on the surface of the perturbed loss fp(w)f_{p}(w) rather than the original loss f(w)f(w) (Foret et al., 2020).

The surrogate gap measures the sharpness at a local minimum

Despite that SAM searches for a region of low losses, we show that a solution by SAM is not guaranteed to be flat. Throughout this paper we measure the sharpness at a local minimum of loss f(w)f(w) by the dominant eigenvalue σmax\sigma_{max} (eigenvalue with the largest absolute value) of Hessian. For simplicity, we do not consider the influence of reparameterization on the geometry of loss surfaces, which is thoroughly discussed in (Laurent & Massart, 2000; Kwon et al., 2021).

For some fixed ρ\rho, consider two local minima w1w_{1} and w2w_{2}, fp(w1)≤fp(w2)\centernot  ⟹  σmax(w1)≤σmax(w2)f_{p}(w_{1})\leq f_{p}(w_{2})\centernot\implies\sigma_{max}(w_{1})\leq\sigma_{max}(w_{2}), where σmax\sigma_{max} is the dominant eigenvalue of the Hessian.

We leave the proof to Appendix. Fig. 1 illustrates Lemma 3.1 with an example. Consider three local minima denoted as w1w_{1} to w3w_{3}, and suppose the corresponding loss surfaces are flatter from w1w_{1} to w3w_{3}. For some fixed ρ\rho, we plot the perturbed loss fpf_{p} and surrogate gap h≜fp−fh\triangleq f_{p}-f around each solution. Comparing w2w_{2} with w3w_{3}: Suppose their vanilla losses are equal, f(w2)=f(w3)f(w_{2})=f(w_{3}), then fp(w2)>fp(w3)f_{p}(w_{2})>f_{p}(w_{3}) because the loss surface is flatter around w3w_{3}, implying that SAM will prefer w3w_{3} to w2w_{2}. Comparing w1w_{1} and w2w_{2}: fp(w1)<fp(w2)f_{p}(w_{1})<f_{p}(w_{2}), and SAM will favor w1w_{1} over w2w_{2} because it only cares about the perturbed loss fpf_{p}, even though the loss surface is sharper around w1w_{1} than w2w_{2}.

2 The surrogate gap agrees with sharpness

We introduce the surrogate gap that agrees with sharpness, defined as:

Intuitively, the surrogate gap represents the difference between the maximum loss within the neighborhood and the loss at the center point. The surrogate gap has the following properties.

Suppose the perturbation amplitude ρ\rho is sufficiently small, then the approximation to the surrogate gap in Eq. 4 is always non-negative, h(w)≈f(wadv)−f(w)≥0,∀wh(w)\approx f(w^{adv})-f(w)\geq 0,\forall w.

For a local minimum w∗w^{*}, consider the dominate eigenvalue σmax\sigma_{max} of the Hessian of loss ff as a measure of sharpness. Considering the neighborhood centered at w∗w^{*} with a small radius ρ\rho, the surrogate gap h(w∗)h(w^{*}) is an equivalent measure of the sharpness: σmax≈2h(w∗)/ρ2.\sigma_{max}\approx 2h(w^{*})/\rho^{2}.

The proof is in Appendix. Lemma 3.2 tells that the surrogate gap is non-negative, and Lemma 3.3 shows that the loss surface is flatter as hh gets closer to 0. The two lemmas together indicate that we can find a region with a flat loss surface by minimizing the surrogate gap h(w)h(w).

Surrogate Gap Guided Sharpness-Aware Minimization

Inspired by the analysis in Section 3, we propose Surrogate Gap Guided Sharpness-Aware Minimzation (GSAM) to simultaneously minimize two objectives, the perturbed loss fpf_{p} and the surrogate gap hh:

Intuitively, by minimizng fpf_{p} we search for a region with a low perturbed loss similar to SAM, and by minimizing hh we search for a local minimum with a flat surface. A low perturbed loss implies low training losses within the neighborhood, and a flat loss surface reduces the generalization gap between training and test performances (Chaudhari et al., 2019). When both are minimized, the solution gives rise to high accuracy and good generalization.

Potential caveat in optimization It is tempting and yet sub-optimal to combine the objectives in Eq. 5 to arrive at min⁡wfp(w)+λh(w)\operatorname{min}_{w}f_{p}(w)+\lambda h(w), where λ\lambda is some positive scalar. One caveat when solving this weighted combination is the potential conflict between the gradients of the two terms, i.e., ∇fp(w)\nabla f_{p}(w) and ∇h(w)\nabla h(w). We illustrate this conflict by Fig. 2, where ∇h(w)=∇fp(w)−∇f(w)\nabla h(w)=\nabla f_{p}(w)-\nabla f(w) (the grey dashed arrow) has a negative inner product with ∇fp(w)\nabla f_{p}(w) and ∇f(w)\nabla f(w). Hence, the gradient descent for the surrogate gap could potentially increase the loss fpf_{p}, harming the model’s performance. We empirically validate this argument in Sec. 6.4.

2 Gradient decomposition and ascent for the multi-objective optimization

Our primary goal is to minimize fpf_{p} because otherwise a flat solution of high loss is meaningless, and the minimization of hh should not increase fpf_{p}. We propose to decompose ∇f(wt)\nabla f(w_{t}) and ∇h\nabla h into components that are parallel and orthogonal to ∇fp(wt)\nabla f_{p}(w_{t}), respectively (see Fig. 2):

The key is that updating in the direction of ∇h⊥(wt)\nabla h_{\perp}(w_{t}) does not change the value of the perturbed loss fp(wt)f_{p}(w_{t}) because ∇h⊥⊥∇fp\nabla h_{\perp}\perp\nabla f_{p} by construction. Therefore, we propose to perform a descent step in the ∇h⊥(wt)\nabla h_{\perp}(w_{t}) direction, which is equivalent to an ascent step in the ∇f⊥(wt)\nabla f_{\perp}(w_{t}) direction (because ∇h⊥=−∇f⊥\nabla h_{\perp}=-\nabla f_{\perp} by the definition of hh), and achieve two goals simultaneously — it keeps the value of fp(wt)f_{p}(w_{t}) intact and meanwhile decreases the surrogate gap h(wt)=fp(wt)−f(wt)h(w_{t})=f_{p}(w_{t})-f(w_{t}) (by increasing f(wt)f(w_{t}) and not affect fp(wt)f_{p}(w_{t})).

The full GSAM Algorithm is shown in Algo. 1 and Fig. 2, where g(t),gp(t)g^{(t)},g_{p}^{(t)} are noisy observations of ∇f(wt)\nabla f(w_{t}) and ∇fp(wt)\nabla f_{p}(w_{t}), respectively, and g∥(t),g⊥(t)g_{\parallel}^{(t)},g_{\perp}^{(t)} are noisy observations of ∇f∥(wt)\nabla f_{\parallel}(w_{t}) and ∇f⊥(wt)\nabla f_{\perp}(w_{t}), respectively, by projecting g(t)g^{(t)} onto gp(t)g_{p}^{(t)}. We introduce a constant α\alpha to scale the stepsize of the ascent step. Steps 1) to 2) are the same as SAM: At current point wtw_{t}, step 1) takes a gradient ascent to wtadvw_{t}^{adv} followed by step 2) evaluating the gradient gp(t)g_{p}^{(t)} at wtadvw_{t}^{adv}. Step 3) projects g(t)g^{(t)} onto gp(t)g_{p}^{(t)}, which requires negligible computation compared to the forward and backward passes. In step 4), −ηtgp(t)-\eta_{t}g_{p}^{(t)} is the same as in SAM and minimizes the perturbed loss fp(wt)f_{p}(w_{t}) with gradient descent, and αηtg⊥(t)\alpha\eta_{t}g_{\perp}^{(t)} performs an ascent step in the orthogonal direction of gp(t)g_{p}^{(t)} to minimize the surrogate gap h(wt)h(w_{t}) ( equivalently increase f(wt)f(w_{t}) and keep fp(wt)f_{p}(w_{t}) intact). In coding, GSAM feeds the “surrogate gradient” ∇ftGSAM≜gp(t)−αg⊥(t)\nabla f_{t}^{GSAM}\triangleq g_{p}^{(t)}-\alpha g_{\perp}^{(t)} to first-order gradient optimizers such as SGD and Adam.

The ascent step along g⊥(t)g_{\perp}^{(t)} does not harm convergence SAM demonstrates that minimizing fpf_{p} makes the network generalize better than minimizing ff. Even though our ascent step along g⊥(t)g_{\perp}^{(t)} increases f(w)f(w), it does not affect fp(w)f_{p}(w), so GSAM still decreases the perturbed loss fpf_{p} in a way similar to SAM. In Thm. 5.1, we formally prove the convergence of GSAM. In Sec. 6 and Appendix C, we empirically validate that the loss decreases and accuracy increases with training.

Illustration with a toy example We demonstrate different algorithms by a numerical toy example shown in Fig. 3. The trajectory of GSAM is closer to the ridge and tends to find a flat minimum. Intuitively, since the loss surface is smoother along the ridge than in sharp local minima, the surrogate gap h(w)h(w) is small near the ridge, and the ascent step in GSAM minimizes hh to pushes the trajectory closer to the ridge. More concretely, ∇f(wt)\nabla f(w_{t}) points to a sharp local solution and deviates from the ridge; in contrast, wtadvw_{t}^{adv} is closer to the ridge and ∇f(wtadv)\nabla f(w_{t}^{adv}) is closer to the ridge descent direction than ∇f(wt)\nabla f(w_{t}). Note that ∇ftGSAM\nabla f_{t}^{GSAM} and ∇f(wt)\nabla f(w_{t}) always lie at different sides of ∇fp(wt)\nabla f_{p}(w_{t}) by construction (see Fig. 2), hence ∇ftGSAM\nabla f_{t}^{GSAM} pushes the trajectory closer to the ridge than ∇fp(wt)\nabla f_{p}(w_{t}) does. The trajectory of GSAM is like descent along the ridge and tends to find flat minima.

Theoretical properties of GSAM

Consider a non-convex function f(w)f(w) with Lipschitz-smooth constant LL and lower bound fminf_{min}. Suppose we can access a noisy, bounded observation g(t)g^{(t)} (∣∣g(t)∣∣2≤G,∀t||g^{(t)}||_{2}\leq G,\forall t) of the true gradient ∇f(wt)\nabla f(w_{t}) at the tt-th step. For some constant α\alpha, with learning rate ηt=η0/t\eta_{t}=\eta_{0}/\sqrt{t}, and perturbation amplitude ρt\rho_{t} proportional to the learning rate, e.g., ρt=ρ0/t\rho_{t}=\rho_{0}/\sqrt{t}, we have

where C1,C2,C3,C4C_{1},C_{2},C_{3},C_{4} are some constants.

Thm. 5.1 implies both fpf_{p} and ff converge in GSAM at rate O(log⁡T/T)O(\log T/\sqrt{T}) for non-convex stochastic optimization, matching the convergence rate of first-order gradient optimizers like Adam.

2 Generalization of GSAM

In this section, we show the surrogate gap in GSAM is provably lower than SAM’s, so GSAM is expected to find a smoother minimum with better generalization.

Suppose the training set has mm elements drawn i.i.d. from the true distribution, and denote the loss on the training set as f^(w)=1m∑i=1mf(w,xi),\widehat{f}(w)=\frac{1}{m}\sum_{i=1}^{m}f(w,x_{i}), where we use xix_{i} to denote the (input, target) pair of the ii-th element. Let ww be learned from the training set. Suppose ww is drawn from posterior distribution Q\mathcal{Q}. Denote the prior distribution (independent of training) as P\mathcal{P}, then

where C=f^(w)C=\widehat{f}(w) is the empirical training loss, and h^\widehat{h} is the surrogate gap evaluated on the training set.

Corollary 5.2.1 implies that minimizing h^\widehat{h} (right hand side of Eq. 7) is expected to achieve a tighter upper bound of the generalization performance (left hand side of Eq. 7). The third term on the right of Eq. 7 is typically hard to analyze and often simplified to L2L2 regularization (Foret et al., 2020). Note that fp=C+h^f_{p}=C+\widehat{h} only holds when ρtrain\rho_{train} (the perturbation amplitude specified by users during training) equals ρtrue\rho_{true} (the ground truth value determined by underlying data distribution); when ρtrain≠ρtrue\rho_{train}\neq\rho_{true}, min(fp,h^)min(f_{p},\widehat{h}) is more effective than min(fp)min(f_{p}) in terms of minimizing generalization loss. A detailed discussion is in Appendix A.7.

Under the assumption in Thm. 5.1, Thm. 5.2 and Corollary 5.2.1, we assume the Hessian has a lower-bound ∣σ∣min|\sigma|_{min} on the absolute value of eigenvalue, and the variance of noisy observation g(t)g^{(t)} is lower-bounded by c2c^{2}. The surrogate gap hh can be minimized by the ascent step along the orthogonal direction g⊥(t)g_{\perp}^{(t)}. During training we minimize the sample estimate of hh. We use Δh^t\Delta\widehat{h}_{t} to denote the amount that the ascent step in GSAM decreases h^\widehat{h} for the tt-th step. Compared to SAM, the proposed method generates a total decrease in surrogate gap ∑t=1TΔh^t\sum_{t=1}^{T}\Delta\widehat{h}_{t}, which is bounded by

We provide proof in the appendix. The lower-bound of ∑t=1TΔh^t\sum_{t=1}^{T}\Delta\widehat{h}_{t} indicates that GSAM achieves a provably non-trivial decrease in the surrogate gap. Combined with Corollary 5.2.1, GSAM provably improves the generalization performance over SAM.

Experiments

We conduct experiments with ResNets (He et al., 2016), Vision Transformers (ViTs) (Dosovitskiy et al., 2020) and MLP-Mixers (Tolstikhin et al., 2021). Following the settings by Chen et al. (2021), we train on the ImageNet-1k (Deng et al., 2009) training set using the Inception-style (Szegedy et al., 2015) pre-processing without extra training data or strong augmentation. For all models, we search for the best learning rate and weight decay for vanilla training, and then use the same values for the experiments with SAM and GSAM. For ResNets, we search for ρ\rho from 0.01 to 0.05 with a stepsize 0.01. For ViTs and Mixers, we search for ρ\rho from 0.05 to 0.6 with a stepsize 0.05. In GSAM, we search for α\alpha in {0.01,0.02,0.03}\{0.01,0.02,0.03\} for ResNets and α\alpha in {0.1,0.2,0.3}\{0.1,0.2,0.3\} for ViTs and Mixers. Considering that each step in SAM and GSAM requires twice the computation of vanilla training, we experiment with the vanilla training for twice the epochs of SAM and GSAM, but we observe no significant improvements from the longer training (Table 5 in appendix). We summarize the best hyper-parameters for each model in Appendix B.

We report the performances on ImageNet (Deng et al., 2009), ImageNet-v2 (Recht et al., 2019) and ImageNet-Real (Beyer et al., 2020) in Table 1. GSAM consistently improves over SAM and vanilla training (with SGD or AdamW): on ViT-B/32, GSAM achieves +5.4% improvement over AdamW and +3.2% over SAM in top-1 accuracy; on Mixer-B/32, GSAM achieves +11.1% over AdamW and +1.2% over SAM. We ignore the standard deviation since it is typically negligible (<0.1%<0.1\%) compared to the improvements. We also test the generalization performance on out-of-distribution data (ImageNet-R and ImageNet-C), and the observation is consistent with that on ImageNet, e.g., +5.1% on ImageNet-R and +5.9% on ImageNet-C for Mixer-B/32.

2 GSAM finds a minimum whose Hessian has small dominant eigenvalues

Lemma 3.3 indicates that the surrogate gap hh is an equivalent measure of the dominant eigenvalue of the Hessian, and minimizing hh equivalently searches for a flat minimum. We empirically validate this in Fig. 4. As shown in the left subfigure, for some fixed ρ\rho, increasing α\alpha decreases the dominant value and improves generalization (test accuracy). In the middle subfigure, we plot the dominant eigenvalues estimated by the surrogate gap, σmax≈2h/ρ2\sigma_{max}\approx 2h/\rho^{2} (Lemma 3.3). In the right subfigure, we directly calculate the dominant eigenvalues using the power-iteration (Mises & Pollaczek-Geiringer, 1929). The estimated dominant eigenvalues (middle) match the real eigenvalues σmax\sigma_{max} (right) in terms of the trend that σmax\sigma_{max} decreases with α\alpha and ρ\rho. Note that the surrogate gap hh is derived over the whole training set, while the measured eigenvalues are over a subset to save computation. These results show that the ascent step in GSAM minimizes the dominant eigenvalue by minimizing the surrogate loss, validating Thm 5.3.

3 Comparison with methods in the literature

Section 6.1 compares GSAM to SAM and vanilla training. In this subsection, we further compare GSAM against Entropy-SGD (Chaudhari et al., 2019) and Adaptive-SAM (ASAM) (Kwon et al., 2021), which are designed to improve generalization. Note that Entropy-SGD uses SGD in the inner Langevin iteration and can be combined with other base optimizers such as AdamW as the outer loop. For Entropy-SGD, we find the hyper-parameter “scope” from 0.0 and 0.9, and search for the inner-loop iteration number between 1 and 14. For ASAM, we search for ρ\rho between 1 and 7 (10×10\times larger than in SAM) as recommended by the ASAM authors. Note that the only difference between ASAM and SAM is the derivation of the perturbation, so both can be combined with the proposed ascent step. As shown in Fig. 5, the proposed ascent step increases test accuracy when combined with both SAM and ASAM and outperforms Entropy-SGD and vanilla training.

4 Additional studies

GSAM outperforms a weighted combination of the perturbed loss and surrogate gap With an example in Fig. 2, we demonstrate that directly minimizing fp(w)+λh(w)f_{p}(w)+\lambda h(w) as discussed in Sec. 4.1 is sub-optimal because ∇h(w)\nabla h(w) could conflict with ∇fp(w)\nabla f_{p}(w) and ∇f(w)\nabla f(w). We empirically validate this argument on ViT-B/32. We search for λ\lambda between 0.0 and 0.5 with a step 0.1 and search for ρ\rho in the same grid as SAM and GSAM. We report the best accuracy of each method. Top-1 accuracy in Table 2 show the superior performance of GSAM, validating our analysis.

min⁡(fp,h)\boldsymbol{\operatorname{min}(f_{p},h)} vs. min⁡(f,h)\boldsymbol{\operatorname{min}(f,h)} GSAM solves min⁡(fp,h)\operatorname{min}(f_{p},h) by descent in ∇fp\nabla f_{p}, decomposing ∇f\nabla f onto ∇fp\nabla f_{p}, and an ascent step in the orthogonal direction to increase ff while keep fpf_{p} intact. Alternatively, we can also optimize min⁡(f,h)\operatorname{min}(f,h) by descent in ∇f\nabla f, decomposing ∇fp\nabla f_{p} onto ∇f\nabla f, and a descent step in the orthogonal direction to decrease fpf_{p} while keep ff intact. The two GSAM variations perform similarly (see Fig. 6, right). We choose min⁡(fp,h)\operatorname{min}(f_{p},h) mainly to make the minimal change to SAM.

GSAM benefits transfer learning Using weights trained on ImageNet-1k, we finetune models with SGD on downstream tasks including the CIFAR10/CIFAR100 (Krizhevsky et al., 2009), Oxford-flowers (Nilsback & Zisserman, 2008) and Oxford-IITPets (Parkhi et al., 2012). Results in Table 3 shows that GSAM leads to better transfer performance than vanilla training and SAM.

GSAM remains effective under various data augmentations We plot the top-1 accuracy of a ViT-B/32 model under various Mixup (Zhang et al., 2017b) augmentations in Fig. 6 (left subfigure). Under different augmentations, GSAM consistently outperforms SAM and vanilla training.

GSAM is compatible with different base optimizers GSAM is generic and applicable to various base optimizers. We compare vanilla training, SAM and GSAM using AdamW (Loshchilov & Hutter, 2017) and AdaBelief (Zhuang et al., 2020) with default hyper-parameters. Fig. 6 (middle subfigure) shows that GSAM performs the best, and SAM improves over vanilla training.

Conclusion

We propose the surrogate gap as an equivalent measure of sharpness which is easy to compute and feasible to optimize. We propose the GSAM method, which improves the generalization over SAM at negligible computation cost. We show the convergence and provably better generalization of GSAM compared to SAM, and validate the superior performance of GSAM on various models.

Acknowledgement

We would like to thank Xiangning Chen (UCLA) and Hossein Mobahi (Google) for discussions, Yi Tay (Google) for help with datasets, and Yeqing Li, Xianzhi Du, and Shawn Wang (Google) for help with TensorFlow implementation.

Ethics Statement

This paper focuses on the development of optimization methodologies and can be applied to the training of different deep neural networks for a wide range of applications. Therefore, the ethical impact of our work would primarily be determined by the specific models that are trained using our new optimization strategy.

Reproducibility Statement

We provide the detailed proof of theoretical results in Appendix A and provide the data pre-processing and hyper-parameter settings in Appendix B. Together with the references to existing works and public codebases, we believe the paper contains sufficient details to ensure reproducibility. We plan to release the models trained by using GSAM upon publication.

References

Appendix A Proofs

Suppose ρ\rho is small, perform Taylor expansion around the local minima ww, we have:

where HH is the Hessian, and is positive semidefinite at a local minima. At a local minima, ∇f(w)=0\nabla f(w)=0, hence we have

where σmax\sigma_{max} is the dominate eigenvalue (eigenvalue with the largest absolute value). Now consider two local minima w1w_{1} and w2w_{2} with dominate eigenvalue σ1\sigma_{1} and σ2\sigma_{2} respectively, we have

We have fp(w1)>fp(w2)\centernot  ⟹  σ1>σ2f_{p}(w_{1})>f_{p}(w_{2})\centernot\implies\sigma_{1}>\sigma_{2} and σ1>σ2\centernot  ⟹  fp(w1)>fp(w2)\sigma_{1}>\sigma_{2}\centernot\implies f_{p}(w_{1})>f_{p}(w_{2}) because the relation between f(w1)f(w_{1}) and f(w2)f(w_{2}) is undetermined. □\square

A.2 Proof of Lemma. 3.2

Since ρ\rho is small, we can perform Taylor expansion around ww,

where the last line is because δ\delta is approximated as δ=ρ∇f(w)∣∣∇f(w)∣∣2+ϵ\delta=\rho\frac{\nabla f(w)}{||\nabla f(w)||_{2}+\epsilon}, hence has the same direction as ∇f(w)\nabla f(w). □\square

A.3 Proof of Lemma. 3.3

Since ρ\rho is small, we can approximate f(w)f(w) with a quadratic model around a local minima ww:

where HH is the Hessian at ww, assumed to be positive semidefinite at local minima. Normalize δ\delta such that ∣∣δ∣∣2=ρ||\delta||_{2}=\rho, Hence we have:

where σmax\sigma_{max} is the dominate eigenvalue of the hessian HH, and first order term is 0 because the gradient is 0 at local minima. Therefore, we have σmax≈2h(w)/ρ2\sigma_{max}\approx 2h(w)/\rho^{2}. □\square

A.4 Proof of Thm. 5.1

For simplicity we consider the base optimizer is SGD. For other optimizers such as Adam, we can derive similar results by applying standard proof techniques in the literature to our proof.

For simplicity of notation, we denote the update at step tt as

By L−L-smoothness of ff and the definition of fp(wt)=f(wtadv)f_{p}(w_{t})=f(w^{adv}_{t}), and definition of dt=wt+1−wtd_{t}=w_{t+1}-w_{t} and wtadv=wt+δtw^{adv}_{t}=w_{t}+\delta_{t} we have

Step 1.0: Bound Eq. 18

Step 1.1: Bound Eq. 19

where g(t)g^{(t)} is the gradient of ff at wtw_{t} evaluated with a noisy data sample. When learning rate ηt\eta_{t} is small, the update in weight dtd_{t} is small, and expected gradient is

where HH is the Hessian at wtw_{t}. Therefore, we have

where the first inequality is due to (1) ρt\rho_{t} is monotonically decreasing with tt, and (2) triangle inequality that ⟨a,b⟩≤∣∣a∣∣⋅∣∣b∣∣\langle a,b\rangle\leq||a||\cdot||b||. ϕt\phi_{t} is the angle between the unit vector in the direction of ∇f(wt)\nabla f(w_{t}) and ∇f(wt+1)\nabla f(w_{t+1}). The second inequality comes from that (1) \Big{|}\Big{|}\frac{g}{||g||+\epsilon}\Big{|}\Big{|}<1 strictly, so we can replace δt\delta_{t} in Eq. 25 with a unit vector in corresponding directions multiplied by ρt\rho_{t} and get the upper bound, (2) the norm of difference in unit vectors can be upper bounded by the arc length on a unit circle.

When learning rate ηt\eta_{t} and update stepsize dtd_{t} is small, ϕt\phi_{t} is also small. Using the limit that

Plug into Eq. 27, also note that the perturbation amplitude ρt\rho_{t} is small so wtw_{t} is close to wtadvw^{adv}_{t}, then we have

Step 1.2: Total bound

Reuse results from Eq. 21 (replace LpL_{p} with 2L2L) and plug into Eq. 18, and plug Eq. 31 and Eq. 34 into Eq. 19, we have

Note that ηT=η0T\eta_{T}=\frac{\eta_{0}}{\sqrt{T}}, we have

which implies that GSAM enables fpf_{p} to converge at a rate of O(log⁡T/T)O(\log T/\sqrt{T}), and all the constants here are well-bounded.

Step 2: Convergence w.r.t. function f​(w)𝑓𝑤f(w)

We prove the risk for f(w)f(w) convergences for non-convex stochastic optimization case using SGD. Denote the update at step tt as

For simplicity, we introduce a scalar βt\beta_{t} such that

where ∇f∥(wt)\nabla f_{\parallel}(w_{t}) is the projection of ∇f(wt)\nabla f(w_{t}) onto ∇fp(wt)\nabla f_{p}(w_{t}). When perturbation amplitude ρ\rho is small, we expect βt\beta_{t} to be very close to 1.

Take expectation conditioned on observations up to step tt for both sides of Eq. 42, we have:

Also note when perturbation amplitude ρt\rho_{t} is small, we have

where δt=ρt∇f(wt)∣∣∇f(wt)∣∣2\delta_{t}=\rho_{t}\frac{\nabla f(w_{t})}{||\nabla f(w_{t})||_{2}} by definition, H(wt)H(w_{t}) is the Hessian. Hence we have

where LL is the Lipschitz constant of ff, and L−L-smoothness of ff indicates the maximum absolute eigenvalue of HH is upper bounded by LL. Plug Eq. 49 into Eq. 47, we have

perform telescope sum and taking expectations on each step, we have

Take the schedule to be ηt=η0t\eta_{t}=\frac{\eta_{0}}{\sqrt{t}} and ρt=ρ0t\rho_{t}=\frac{\rho_{0}}{\sqrt{t}}, then we have

where C1,C4C_{1},C_{4} are some constants. This implies the convergence rate w.r.t f(w)f(w) is O(log⁡T/T)O(\log T/\sqrt{T}).

Step 3: Convergence w.r.t. surrogate gap h​(w)ℎ𝑤h(w)

Note that we have proved convergence for fp(w)f_{p}(w) in step 1, and convergence for f(w)f(w) in step 3. Also note that

also converges at rate O(log⁡T/T)O(\log T/\sqrt{T}) because each item in the RHS converges at rate O(log⁡TT)O(\log T\sqrt{T}). □\square

A.5 Proof of Corollary. 5.2.1

Using the results from Thm. 5.2, with probability at least 1−a1-a, we have

Assume δ∼N(0,b2Ik)\delta\sim\mathcal{N}(0,b^{2}I_{k}) where kk is the dimension of model parameters, hence δ2\delta^{2} (element-wise square) follows a a Chi-square distribution. By Lemma.1 in Laurent & Massart (2000), we have

hence with probability at least 1−1/n1-1/\sqrt{n}, we have

Therefore, with probability at least 1-1/\sqrt{n}=1-exp\Big{(}-\big{(}\frac{\rho}{\sqrt{2}b}-\sqrt{k}\big{)}^{2}\Big{)}

A.6 Proof of Thm. 5.3

Take Taylor expansion, then the expected change of loss gap caused by descent step is

where θt\theta_{t} is the angle between vector ∇fp(wt)\nabla f_{p}(w_{t}) and ∇f(wt)\nabla f(w_{t}). The expected change of loss gap caused by ascent step is

Next we give an estimate of the decrease in h^\widehat{h} caused by our ascent step. We refer to Eq. 69 and Eq. 70 to analyze the change in loss gap caused by the descent and ascent (orthogonally) respectively. It can be seen that gradient descent step might not decrease loss gap, in fact they often increase loss gap in practice; while the ascent step is guaranteed to decrease the loss gap.

Hence we derive an upper bound for ∑t=1TΔh^t\sum_{t=1}^{T}\Delta\widehat{h}_{t}.

Next we derive a lower bound for ∑t=1TΔh^t\sum_{t=1}^{T}\Delta\widehat{h}_{t} Note that when ρt\rho_{t} is small, by Taylor expansion

where H^(wt)\widehat{H}(w_{t}) is the Hessian evaluated on training samples. Also when ρt\rho_{t} is small, the angle θt\theta_{t} between ∇f^p(wt)\nabla\widehat{f}_{p}(w_{t}) and ∇f^(wt)\nabla\widehat{f}(w_{t}) is small, by the limit that

where GG is the upper-bound on norm of gradient, ∣σ∣min|\sigma|_{min} is the minimum absolute eigenvalue of the Hessian. The intuition is that as perturbation amplitude decreases, the angle θt\theta_{t} decreases at a similar rate, though the scale constant might be different. Hence we have

where c2c^{2} is the lower bound of ∣∣∇f^∣∣2||\nabla\widehat{f}||^{2} (e.g. due to noise in data and gradient observation). Results above indicate that the decrease in loss gap caused by the ascent step is non-trivial, hence our proposed method efficiently improves generalization compared with SAM. □\square

A.7 Discussion on Corollary 5.2.1

The comment “‘The corollary gives a bound on the risk in terms of the perturbed training loss if one removes CC from both sides”’ is correct. But there is a misunderstanding in the statement “‘the perturbed training loss is small then the model has a small risk”’: it’s only true when ρtrain\rho_{train} for training equals its real value ρtrue\rho_{true} determined by the data distribution; in practice, we never know ρtrue\rho_{true}. In the following we show that the minimization of both hh and fpf_{p} is better than simply minimizing fpf_{p} when ρtrue≠ρtrain\rho_{true}\neq\rho_{train}.

1. First, we re-write the conclusion of Corollary 5.2.1 as

where RR is the regularization term, CC is the training loss, σ\sigma is the dominant eigenvalue of Hessian. As in lemma 3.3, we perform Taylor-expansion and can ignore the high-order term O(ρ3)O(\rho^{3}). We focus on

2. When ρtrue≠ρtrain\rho_{true}\neq\rho_{train}, minimizing hh achieves a lower risk than only minimizing fpf_{p}. (1) Note that after training, CC (training loss) is fixed, but hh could vary with ρ\rho (e.g. when training on dataset A and testing on an unrelated dataset B, the training loss remains unchanged, but the risk would be huge and a large ρ\rho is required for a valid bound). (2) With an example, we show a low fpf_{p} is insufficient for generalization, and a low σ\sigma is necessary:

Suppose we use ρtrain\rho_{train} for training, and consider two solutions with C1,σ1C_{1},\sigma_{1} (SAM) and C2,σ2C_{2},\sigma_{2} (GSAM). Suppose they have the same fpf_{p} during training for some ρtrain\rho_{train}, so

Suppose C1<C2C_{1}<C_{2} so σ1>σ2\sigma_{1}>\sigma_{2}.

This implies that a small σ\sigma helps generalization, but only a low fp1f_{p1} (caused by a low C1C_{1} and high σ1\sigma_{1}) is insufficient for a good generalization.

Note that ρtrain\rho_{train} is fixed during training, so minimizing htrainh_{train} during training is equivalently minimizing σ\sigma by Lemma 3.3

3. Why we are often unlucky to have ρtrue>ρtrain\rho_{true}>\rho_{train} (1) First, the test sets are almost surely outside the convex hull of the training set because “‘interpolation almost surely never occurs in high-dimensional (>100>100) cases”’ Balestriero et al. (2021). As a result, the variability of (train + test) sets is almost surely larger than the variability of (train) set. Since ρ\rho increases with data variability (see point 4 below), we have ρtrue>ρtrain_set\rho_{true}>\rho_{train\\ \_set} almost surely. (2) Second, we don’t know the value of ρtrue\rho_{true} and can only guess it. In practice, we often guess a small value because training often diverges with large ρ\rho (as observed in Foret et al. (2020); Chen et al. (2021)).

4. Why ρ\rho increases with data variability. In Corollary 5.2.1, we assume weight perturbation δ∼N(0,b2Ik)\delta\sim\mathcal{N}(0,b^{2}I^{k}). The meaning of bb is the following. If we can randomly sample a fixed number of samples from the underlying distribution, then training the model from scratch (with a fixed seed for random initialization) gives rise to a set of weights. Repeating this process, we get many sets of weights, and their standard deviation is bb. Since the number of training samples is limited and fixed, the more variability in data, the more variability in weights, and the larger bb. Note that Corollary stated that the bound holds with probability proportional to [1−e−(ρ2b−k)2][1-e^{-(\frac{\rho}{\sqrt{2}b}-\sqrt{k})^{2}}]. In order for the result to hold with a fixed probability, ρ\rho must stay proportional to bb, hence ρ\rho also increases with the variability of data.

Appendix B Experimental Details

For ViT and Mixer, we search the learning rate in {1e-3, 3e-3, 1e-2, 3e-3}, and search weight decay in {0.003, 0.03, 0.3}. For ResNet, we search the learning rate in {1.6, 0.16, 0.016}, and search the weight decay in {0.001, 0.01,0.1}. For ViT and Mixer, we use the AdamW optimizer with β1=0.9,β2=0.999\beta_{1}=0.9,\beta_{2}=0.999; for ResNet we use SGD with momentum=0.9=0.9. We train ResNets for 90 epochs, and train ViTs and Mixers for 300 epochs following the settings in (Chen et al., 2021) and (Dosovitskiy et al., 2020). Considering that SAM and GSAM uses twice the computation of vanilla training for each step, for vanilla training we try 2×2\times longer training, and does not find significant improvement as in Table. 5.

We first search the optimal learning rate and weight decay for vanilla training, and keep these two hyper-parameters fixed for SAM and GSAM. For ViT and Mixer, we search ρ\rho in {0.1, 0.2, 0.3, 0.4, 0.5, 0.6} for SAM and GSAM; for ResNet, we search ρ\rho from 0.01 to 0.05 with a stepsize 0.01. For ASAM, we amplify ρ\rho by 10×10\times compared to SAM, as recommended by Kwon et al. (2021). For GSAM, we search α\alpha in {0.1, 0.2, 0.3} throughout the paper. We report the best configuration of each individual model in Table. 4.

B.2 Transfer learning experiments

Using weights trained on ImageNet-1k, we finetune models with SGD on downstream tasks including the CIFAR10/CIFAR100 (Krizhevsky et al., 2009), Oxford-flowers (Nilsback & Zisserman, 2008) and Oxford-IITPets (Parkhi et al., 2012). For all experiments, we use the SGD optimizer with no weight decay under a linear learning rate schedule and gradient clipping with global norm 1. We search the maximum learning rate in {0.001, 0.003, 0.01, 0.03}. On Cifar datasets, we train models for 10k steps with a warmup step of 500; on Oxford datasets, we train models for 500 steps with a wamup step of 100.

B.3 Experimental setup with ablation studies on data augmentation

We follow the settings in (Tolstikhin et al., 2021) to perform ablation studies on data augmentation. In the left subfigure of Fig. 6, “Light” refers to Inception-style data augmentation with random flip and crop of images, “Medium” refers to the mixup augmentation with probability 0.2 and RandAug magnitude 10; “Strong” refers to the mixup augmentation with probability 0.2 and RandAug magnitude 15.

Appendix C Ablation studies and discussions

We plot the performance of a ViT-B/32 model varying with ρ\rho (Fig. 7(a)) and α\alpha (Fig. 7(b)). We empirically validate that fine-tuning ρ\rho in SAM can not achieve comparable performance with GSAM, as shown in Fig. 7(a). Considering that GSAM has one more parameter α\alpha, we plot the accuracy varying with α\alpha in Fig. 7(b), and show that GSAM consistently outperforms SAM and vanilla training.

Note that Thm. 5.1 assumes ρt\rho_{t} to decay with tt in order to prove the convergence, while SAM uses a constant ρ\rho during training. To eliminate the influence of ρt\rho_{t} schedule, we conduct ablation study as in Table. 6. The ascent step in GSAM can be applied to both constant ρ\rho or a decayed ρt\rho_{t} schedule, and improves accuracy for both cases. Without ascent step, constant ρ\rho and decayed ρt\rho_{t} achieve similar performance. Results in Table. 6 implies that the ascent step in GSAM is the main reason for improvement of generalization performance.

C.3 Visualize the training process

In the proof of Thm. 5.3, our analysis relies on assumption that θt\theta_{t} is small. We empirically validated this assumption by plotting cos⁡θt\operatorname{cos}\theta_{t} in Fig. 9, where θt\theta_{t} is the angle between ∇f(wt)\nabla f(w_{t}) and ∇fp(wt)\nabla f_{p}(w_{t}). Note that the cosine value is calculated in the parameter space of dimension 8.8×1078.8\times 10^{7}, and in high-dimensional space two random vectors are highly likely to be perpendicular. In Fig. 9 the cosine value is always above 0.9, indicating that ∇f(wt)\nabla f(w_{t}) and ∇fp(wt)\nabla f_{p}(w_{t}) point to very close directions considering the high dimension of parameters. This empirically validates our assumption that θt\theta_{t} is small during training.

We also plot the surrogate gap during training in Fig. 9. As α\alpha increases, the surrogate gap decreases, validating that the ascent step in GSAM efficiently minimizes the surrogate gap. Furthermore, the surrogate gap increases with training steps for any fixed α\alpha, indicating that the training process gradually falls into local minimum in order to minimize the training loss.

Appendix D Related works

Besides SAM and ASAM, other methods were proposed in the literature to improve generalization: Lin et al. (2020) proposed extrapolation of gradient, Xie et al. (2021) proposed to manipulate the noise in gradient, and Damian et al. (2021) proved label noise improves generalization, Yue et al. (2020) proposed to adjust learning rate according to sharpness, and Zheng et al. (2021) proposed model perturbation with similar idea to SAM. Izmailov et al. (2018) proposed averaging weights to improve generalization, and Heo et al. (2020) restricted the norm of updated weights to improve generalization. Many of aforementioned methods can be combined with GSAM to further improve generalization.

Besides modified training schemes, there are other two types of techniques to improve generalization: data augmentation and model regularization. Data augmentation typically generates new data from training samples; besides standard data augmentation such as flipping or rotation of images, recent data augmentations include label smoothing (Müller et al., 2019) and mixup (Müller et al., 2019) which trains on convex combinations of both inputs and labels, automatically learned augmentation (Cubuk et al., 2018), and cutout (DeVries & Taylor, 2017) which randomly masks out parts of an image. Model regularization typically applies auxiliary losses besides the training loss such as weight decay (Loshchilov & Hutter, 2017), other methods randomly modify the model architecture during training, such as dropout (Srivastava et al., 2014) and shake-shake regularization (Gastaldi, 2017). Note that the data augmentation and model regularization literature mentioned here typically train with the standard back-propagation (Rumelhart et al., 1985) and first-order gradient optimizers, and both techniques can be combined with GSAM.

Besides SGD, Adam and AdaBelief, GSAM can be combined with other first-order gradient optimizers, such as AdaBound (Luo et al., 2019), RAdam (Liu et al., 2019), Yogi (Zaheer et al., 2018), AdaGrad (Duchi et al., 2011), AMSGrad (Reddi et al., 2019) and AdaDelta (Zeiler, 2012).