Towards Understanding Sharpness-Aware Minimization

Maksym AndriushchenkoNicolas Flammarion

article2022ICML187 citations

Explains the generalization benefits of Sharpness-Aware Minimization by analyzing its implicit bias and proving its convergence with stochastic gradients, revealing why mini-batch sharpness yields superior performance over standard gradient descent.

Listen

Modern deep neural networks often contain far more parameters than training examples. While standard optimization algorithms like stochastic gradient descent achieve zero training error, they can converge to solutions that generalize poorly to unseen data or overfit severely when data contains label noise. Sharpness-Aware Minimization has emerged as an effective training method to improve generalization by calculating worst-case parameter perturbations during optimization. However, the theoretical mechanisms explaining why this method works remain incomplete, and standard justifications fail to capture why calculating perturbations over small data subsets is essential to its real-world performance.

The main objective of the article is to establish a rigorous theoretical and empirical foundation for the success of Sharpness-Aware Minimization. It evaluates the limitations of existing explanations, demonstrates the implicit optimization bias that drives its generalization benefits, and provides formal convergence guarantees for non-convex models trained with stochastic gradients.

To investigate these questions, the authors combine mathematical proofs with controlled computer vision experiments. The theoretical analysis characterizes algorithm behavior on overparametrized diagonal linear networks and general non-convex optimization objectives. The experimental work evaluates deep neural networks, including ResNet-18 and ResNet-34 architectures trained on standard vision benchmarks such as CIFAR-10 and CIFAR-100, across varying batch sizes, perturbation radii, network widths, and label noise conditions.

The investigation yields four central findings. First, existing explanations based on flat minima and theoretical generalization bounds are incomplete; models converging to flatter loss regions can actually generalize worse than sharper models trained with standard gradient descent, and random perturbations fail to reproduce the performance gains of worst-case perturbations. Second, the success of the method strictly requires computing perturbations over small data batches rather than the full training set, and this benefit stems from a superior implicit optimization bias that favors sparse, lower-complexity model solutions. Third, the authors mathematically prove that Sharpness-Aware Minimization converges to stationary points at rates matching standard gradient descent, provided the perturbation step size decreases appropriately relative to the learning rate. Fourth, empirical trials show that switching from standard training to Sharpness-Aware Minimization for only the final 10% of training epochs allows the model to escape suboptimal solutions and achieve the full generalization benefit.

These findings provide clear practical implications for training costs and workflow efficiency. Because the benefits of Sharpness-Aware Minimization are concentrated toward the end of optimization, practitioners do not need to incur the double computational cost of the algorithm throughout the entire training lifecycle. Instead, models can be pre-trained rapidly using standard empirical risk minimization and fine-tuned cheaply with Sharpness-Aware Minimization. Furthermore, because the algorithm provably converges to zero training loss, it will eventually fit incorrect data, demonstrating that explicit early stopping is still required when datasets contain significant label noise.

Based on these results, machine learning practitioners should adopt small-batch implementations of Sharpness-Aware Minimization and consider fine-tuning pre-trained models rather than training from scratch to reduce compute overhead. Perturbation radii should remain constant rather than decaying rapidly alongside standard learning rates. When deploying on noisy real-world data, teams must pair the method with validation-based early stopping to prevent over-fitting to corrupted labels.

The primary limitation of the study is that full hyperparameter grid searches were constrained to single-GPU workloads on benchmark datasets like CIFAR rather than web-scale datasets like ImageNet. However, the theoretical derivations are mathematically rigorous and closely align with empirical observations across multiple network architectures and normalization schemes, providing high confidence in the core findings.

Cover for Towards Understanding Sharpness-Aware Minimization

Abstract

Sharpness-Aware Minimization (SAM) is a recent training method that relies on worst-case weight perturbations which significantly improves generalization in various settings. We argue that the existing justifications for the success of SAM which are based on a PAC-Bayes generalization bound and the idea of convergence to flat minima are incomplete. Moreover, there are no explanations for the success of using m-sharpness in SAM which has been shown as essential for generalization. To better understand this aspect of SAM, we theoretically analyze its implicit bias for diagonal linear networks. We prove that SAM always chooses a solution that enjoys better generalization properties than standard gradient descent for a certain class of problems, and this effect is amplified by using m-sharpness. We further study the properties of the implicit bias on non-linear networks empirically, where we show that fine-tuning a standard model with SAM can lead to significant generalization improvements. Finally, we provide convergence results of SAM for non-convex objectives when used with stochastic gradients. We illustrate these results empirically for deep networks and discuss their relation to the generalization behavior of SAM. The code of our experiments is available at https://github.com/tml-epfl/understanding-sam.

Table of Contents

  • 1. Introduction
  • 2. Background on SAM
  • 3. Challenging the Existing Understanding of SAM
  • 4. Understanding the Generalization Benefits of SAM
  • 4.1. Testing Two Natural Hypotheses for Why Low ρ in ρ-SAM Could be Beneficial
  • 4.2. Provable Benefit of SAM for Diagonal Linear Networks
  • 4.3. Empirical Study of the Implicit Bias in Non-Linear Networks
  • 5. Understanding the Optimization Aspects of SAM
  • 5.1. Theoretical Analysis of Convergence of SAM
  • 5.2. Convergence of SAM for Deep Networks
  • 6. Conclusions
  • References
  • Organization of the Appendix
  • A. Implementations of the SAM Algorithm in the Full-Batch Setting
  • B. Theoretical Analysis of the Implicit Bias for Diagonal Linear Networks
  • Appendix
  • B.1. Implicit Bias of the ρ-SAM Algorithm
  • B.2. Implicit Bias of the 1-SAM Algorithm
  • B.3. Comparison between 1-SAM and ρ-SAM
  • C. Convergence of the SAM Algorithm
  • C.1. Convergence of Full-Batch ρ-SAM
  • C.2. Convergence of Stochastic SAM
  • C.2.1. Convergence of ρ-SAM
  • C.2.2. Convergence of m-SAM
  • D. Experimental Details
  • E. Additional Deep Learning Experiments
  • E.1. The Effect of ρ in ρ-SAM
  • E.2. The Effect of the Batch Size on SAM
  • E.3. The Effect of the Model Width on SAM
  • E.4. Sharpness for Models with Batch Normalization
  • E.5. Training Loss for ERM vs. SAM Models
  • E.6. SAM with a Decreasing Perturbation Radius ρ
  • E.7. Experiments with Noisy Labels

Knowls

  1. Knowl 1 — Sharpness-aware objectives and stochastic SAM update

    model/method

    Let the training set be Sexttrain={(xi,yi)}i=1nS_{ ext{train}}=\{(x_i,y_i)\}_{i=1}^n, let w∈Rdw\in\mathbb{R}^d be the model parameters, and let ℓi(w)\ell_i(w) be the loss on example (xi,yi)(x_i,y_i). For a subset S⊆StrainS\subseteq S_{\text{train}} and perturbation radius ρ>0\rho>0, the sharpness on SS is

    s(w,S)=max⁡∥δ∥2≤ρ1∣S∣∑(xi,yi)∈S[ℓi(w+δ)−ℓi(w)].s(w,S)=\max_{\|\delta\|_2\le \rho}\frac{1}{|S|}\sum_{(x_i,y_i)\in S}\bigl[\ell_i(w+\delta)-\ell_i(w)\bigr].

    The nn-SAM objective maximizes the total loss over all nn training examples using one shared perturbation, whereas mm-SAM averages objectives over subsets of mm examples, with a separate worst-case perturbation for each subset:

    Fm(w)=1∣Bm∣∑S∈Bmmax⁡∥δ∥2≤ρ∑i∈Sℓi(w+δ),F_m(w)=\frac{1}{|\mathcal{B}_m|}\sum_{S\in\mathcal{B}_m}\max_{\|\delta\|_2\le\rho}\sum_{i\in S}\ell_i(w+\delta),

    where Bm\mathcal{B}_m is the collection of training subsets of size mm; nn-SAM is the special case m=nm=n.

    A practical stochastic implementation samples a batch ItI_t of size mm, uses the same batch for the ascent and descent steps, and updates

    wt+1=wt−γtm∑i∈It∇ℓi ⁣(wt+ρtm∑j∈It∇ℓj(wt)),w_{t+1}=w_t-\frac{\gamma_t}{m}\sum_{i\in I_t}\nabla\ell_i\!\left(w_t+\frac{\rho_t}{m}\sum_{j\in I_t}\nabla\ell_j(w_t)\right),

    where γt>0\gamma_t>0 is the outer descent step size and ρt≥0\rho_t\ge 0 is the inner perturbation step size. The commonly used gradient-normalized variant chooses ρt\rho_t inversely proportional to the norm of the batch gradient, but the paper finds that this normalization is not essential for the generalization benefit.

  2. Knowl 2 — Worst-case perturbations and low m are essential for generalization

    empirical result

    Experiments with pre-activation ResNet-18 on CIFAR-10 and ResNet-34 on CIFAR-100, using standard data augmentation and batch size 128128, compare ERM, random weight perturbations, nn-SAM, and 128128-SAM. Random perturbations produce only marginal changes in test error, and nn-SAM does not substantially improve generalization. In contrast, low-mm SAM with worst-case perturbations, represented by m=128m=128, substantially reduces test error.

    Additional experiments with group normalization and fixed batch size 256256 show that the generalization improvement grows as mm decreases from 256256 to 6464, 1616, and 44, approximately continuously in mm. Thus, the benefit is attributed to the combination of worst-case perturbations and the use of a relatively small mm, rather than to random weight noise or simply perturbing the full training objective. Very small mm is less computationally efficient, so values such as m=128m=128 can trade off generalization and hardware utilization.

  3. Knowl 3 — Evidence that low m changes the implicit bias rather than merely solving the inner maximization better

    empirical result

    Two explanations for the advantage of low mm were tested. First, projected gradient ascent with 100100 inner steps was compared with the usual one-step approximation for computing mm-sharpness at ρ=0.1\rho=0.1. The ratio between the resulting sharpness values increases with mm and reaches roughly 10×10\times for an ERM model on CIFAR-10 at m=1024m=1024, showing that one-step SAM can solve the inner maximization less accurately for large mm. However, using 1010 ascent steps instead of one does not improve final generalization; on CIFAR-10 it mainly shifts the best perturbation radius from about 0.10.1 to 0.050.05.

    Second, the possibility that low mm helps only through batch-normalization regularization was tested using group normalization, which has no batch-size-dependent statistics. Low-mm SAM still substantially improves generalization, including with m=4m=4 and outer batch size 256256, on both CIFAR-10 and CIFAR-100. These results support the hypothesis that low-mm SAM induces a more favorable optimization implicit bias.

  4. Knowl 4 — Implicit-bias theorem for diagonal linear networks

    theoretical result

    Consider an overparameterized diagonal linear network with parameters w=(w+,w−)w=(w_+,w_-), w+,w−∈Rdw_+,w_-\in\mathbb{R}^d, predictor β(w)=w+⊙2−w−⊙2\beta(w)=w_+^{\odot 2}-w_-^{\odot 2}, data matrix X∈Rn×dX\in\mathbb{R}^{n\times d}, labels y∈Rny\in\mathbb{R}^n, and squared training loss

    L(w)=14n∥Xβ(w)−y∥22.L(w)=\frac{1}{4n}\|X\beta(w)-y\|_2^2.

    Assume n<dn<d, initialize w+=w−=αw_+=w_-=\alpha coordinatewise for α∈R>0d\alpha\in\mathbb{R}_{>0}^d, and assume the relevant gradient flow converges to an interpolating predictor β∞\beta_\infty satisfying Xβ∞=yX\beta_\infty=y. Define

    ϕα(β)=∑k=1dαk2q ⁣(βkαk2),q(z)=2−4+z2+z arcsinh⁡(z/2).\phi_\alpha(\beta)=\sum_{k=1}^d \alpha_k^2 q\!\left(\frac{\beta_k}{\alpha_k^2}\right), \qquad q(z)=2-\sqrt{4+z^2}+z\,\operatorname{arcsinh}(z/2).

    The limiting predictor selected by gradient flow, full-batch 11-SAM, or full-batch nn-SAM is the constrained minimizer of this potential, but with an algorithm-dependent effective initialization scale:

    β∞=arg⁡min⁡β∈Rd: Xβ=yϕαalg(β).\beta_\infty=\arg\min_{\beta\in\mathbb{R}^d:\,X\beta=y}\phi_{\alpha_{\mathrm{alg}}}(\beta).

    For sufficiently small inner step size ρ\rho, the effective scales for the two SAM variants satisfy, coordinatewise,

    αn-SAM=α⊙exp⁡ ⁣[−2ρn2∫0∞(X⊤r(s))⊙2 ds+O(ρ2)],\alpha_{n\text{-SAM}} =\alpha\odot\exp\!\left[-\frac{2\rho}{n^2}\int_0^\infty (X^\top r(s))^{\odot 2}\,ds+O(\rho^2)\right],

    and

    α1-SAM=α⊙exp⁡ ⁣[−8ρn∫0∞∑i=1nxi⊙2(xi⊤β(s)−yi)2 ds+O(ρ2)],\alpha_{1\text{-SAM}} =\alpha\odot\exp\!\left[-\frac{8\rho}{n}\int_0^\infty\sum_{i=1}^n x_i^{\odot 2}\bigl(x_i^\top\beta(s)-y_i\bigr)^2\,ds+O(\rho^2)\right],

    where xi⊤x_i^\top is row ii of XX, β(s)\beta(s) is the predictor along the corresponding flow, and r(s)=Xβ(s)−yr(s)=X\beta(s)-y. The potential ϕα\phi_\alpha approaches an ℓ1\ell_1-type regularizer for small initialization scales and an ℓ2\ell_2-type regularizer for large scales. SAM reduces the effective scale, and 11-SAM generally reduces it more strongly than nn-SAM, producing a stronger sparsity-inducing bias.

  5. Knowl 5 — The stronger 1-SAM bias scales with the loss accumulated along training

    theoretical result

    For the diagonal linear-network setting, the leading biasing terms in the effective initialization scales are governed by

    In-SAM(t)=1n2∥X⊤r(t)∥22,I1-SAM(t)=1n∑i=1n∥xi∥22ri(t)2,I_{n\text{-SAM}}(t)=\frac{1}{n^2}\left\|X^\top r(t)\right\|_2^2, \qquad I_{1\text{-SAM}}(t)=\frac{1}{n}\sum_{i=1}^n\|x_i\|_2^2 r_i(t)^2,

    where r(t)=Xβ(t)−yr(t)=X\beta(t)-y and ri(t)r_i(t) is its iith component. For approximately isotropic Gaussian inputs xi∼N(0,Id)x_i\sim\mathcal{N}(0,I_d) in the overparameterized regime d≫nd\gg n, the paper estimates

    ∥In-SAM(t)∥1≈dnL(w(t)),∥I1-SAM(t)∥1≈dL(w(t)).\|I_{n\text{-SAM}}(t)\|_1\approx \frac{d}{n}L(w(t)), \qquad \|I_{1\text{-SAM}}(t)\|_1\approx dL(w(t)).

    Consequently, the cumulative biasing effect of 11-SAM is typically about a factor of nn larger than that of nn-SAM. The effect is also larger when training loss decreases slowly, because the effective scale depends on an integral of the loss along the optimization trajectory. This explains why nn-SAM often behaves similarly to ERM while 11-SAM more strongly favors sparse, better-generalizing interpolating solutions.

  6. Knowl 6 — Existing PAC-Bayes sharpness bounds do not explain low-m SAM

    limitation

    The usual PAC-Bayes justification for SAM controls an expected loss under random Gaussian parameter perturbations, with a leading term of the form

    Eδ∼N(0,σ2I)[∑i=1nℓi(w+δ)].\mathbb{E}_{\delta\sim\mathcal{N}(0,\sigma^2 I)}\left[\sum_{i=1}^n\ell_i(w+\delta)\right].

    Replacing this expectation by a worst-case maximum over ∥δ∥2≤ρ\|\delta\|_2\le\rho is a post-hoc upper bound and makes the bound looser. More importantly, an analogous bound for mm-SAM can be obtained by upper-bounding the contribution of every mini-batch by a maximum over batches. That bound would predict that low-mm SAM, such as 128128-SAM, has worse generalization than random perturbations and nn-SAM.

    The experiments show the opposite ordering: random perturbations and nn-SAM provide little improvement, whereas low-mm SAM substantially improves test error. Therefore, the existing PAC-Bayes sharpness bound does not distinguish the worst-case, batch-local objective that makes low-mm SAM effective and cannot explain its observed generalization improvement.

  7. Knowl 7 — Lower sharpness does not universally identify better-generalizing solutions

    empirical result

    The paper finds counterexamples in which a model with lower mm-sharpness generalizes worse. For group-normalized networks achieving zero training error, a large-batch SAM model is flatter than a small-batch ERM model according to m=128m=128 sharpness across the tested perturbation radii, yet the small-batch ERM model has lower test error. On CIFAR-10, the respective test errors are 6.17%6.17\% for small-batch ERM and 6.80%6.80\% for large-batch SAM; on CIFAR-100 they are 25.06%25.06\% and 28.31%28.31\%.

    A simple analytic example gives the same limitation. For a linear predictor fw(x)=⟨w,x⟩f_w(x)=\langle w,x\rangle with labels yi∈{−1,1}y_i\in\{-1,1\} and a decreasing margin loss ℓ\ell, the 11-sharpness contribution is

    ∑i=1nmax⁡∥δ∥2≤ρ[ℓ(yi⟨w+δ,xi⟩)−ℓ(yi⟨w,xi⟩)]=∑i=1n[ℓ(yi⟨w,xi⟩−ρ∥xi∥2)−ℓ(yi⟨w,xi⟩)].\sum_{i=1}^n\max_{\|\delta\|_2\le\rho}\left[\ell\bigl(y_i\langle w+\delta,x_i\rangle\bigr)-\ell\bigl(y_i\langle w,x_i\rangle\bigr)\right] = \sum_{i=1}^n\left[\ell\bigl(y_i\langle w,x_i\rangle-\rho\|x_i\|_2\bigr)-\ell\bigl(y_i\langle w,x_i\rangle\bigr)\right].

    The perturbation effect is determined by the feature norms and the margins, so 11-sharpness cannot universally distinguish among equally fitting solutions. Thus, convergence to a flatter minimum is not by itself a sufficient explanation for SAM's generalization behavior.

  8. Knowl 8 — SAM induces simpler functions and can improve an already-trained model

    empirical result

    In a one-hidden-layer ReLU network with 100100 hidden units trained by full-batch gradient descent on a one-dimensional quadratic-regression task, SAM produces simpler interpolations of the observed data than ERM and is substantially more stable across five random initializations. The learned functions under SAM are consistent with a bias toward using a sparse combination of ReLU units.

    For deep networks, experiments that switch between ERM and SAM at different training times show that the method used only at the beginning has little influence on the final test error. In contrast, switching from ERM to SAM near the end of training can substantially improve generalization, while switching from SAM back to ERM after SAM has already trained the model does not substantially degrade it. The ERM-to-SAM endpoint remains in the same connected basin as the ERM endpoint under weight interpolation. These results indicate that SAM can gradually move an ERM solution toward a better-generalizing point without requiring optimization from a special initialization, making SAM-based fine-tuning of standard pretrained models effective.

  9. Knowl 9 — Convergence guarantees for stochastic m-SAM

    theoretical result

    Let L(w)=1n∑i=1nℓi(w)L(w)=\frac{1}{n}\sum_{i=1}^n\ell_i(w), let L∗=inf⁡wL(w)L^*=\inf_w L(w), and assume: (A1) the per-example gradients have bounded variance, Ei∥∇ℓi(w)−∇L(w)∥22≤σ2\mathbb{E}_i\|\nabla\ell_i(w)-\nabla L(w)\|_2^2\le\sigma^2; (A2) every per-example loss is β\beta-smooth, ∥∇ℓi(w)−∇ℓi(v)∥2≤β∥w−v∥2\|\nabla\ell_i(w)-\nabla\ell_i(v)\|_2\le\beta\|w-v\|_2; and, when a function-value guarantee is desired, (A3) LL satisfies the Polyak--Łojasiewicz condition 12∥∇L(w)∥22≥μ(L(w)−L∗)\frac12\|\nabla L(w)\|_2^2\ge\mu(L(w)-L^*) for some μ>0\mu>0.

    For stochastic mm-SAM with batch size bb, the same batch used in the ascent and descent steps, and update

    wt+1=wt−γt1b∑i∈It∇ℓi ⁣(wt+ρt1b∑j∈It∇ℓj(wt)),w_{t+1}=w_t-\gamma_t\frac{1}{b}\sum_{i\in I_t}\nabla\ell_i\!\left(w_t+\rho_t\frac{1}{b}\sum_{j\in I_t}\nabla\ell_j(w_t)\right),

    choose γt=1/(βT)\gamma_t=1/(\beta\sqrt{T}) and ρt=1/(βT1/4)\rho_t=1/(\beta T^{1/4}) for TT iterations. Then

    1TE[∑t=0T−1∥∇L(wt)∥22]≤4βT(L(w0)−L∗)+8σ2bT.\frac{1}{T}\mathbb{E}\left[\sum_{t=0}^{T-1}\|\nabla L(w_t)\|_2^2\right] \le \frac{4\beta}{\sqrt{T}}\bigl(L(w_0)-L^*\bigr)+\frac{8\sigma^2}{b\sqrt{T}}.

    Under the PL condition, using γt=min⁡{(8t+4)/(3μ(t+1)2),1/(2β)}\gamma_t=\min\{(8t+4)/(3\mu(t+1)^2),1/(2\beta)\} and ρt=γt/β\rho_t=\sqrt{\gamma_t/\beta} gives

    E[L(wT)]−L∗≤3β2μ2T2(L(w0)−L∗)+22βσ2μ2bT.\mathbb{E}[L(w_T)]-L^* \le \frac{3\beta^2}{\mu^2T^2}\bigl(L(w_0)-L^*\bigr)+\frac{22\beta\sigma^2}{\mu^2bT}.

    Thus, stochastic SAM reaches stationary points at the usual stochastic-gradient order and achieves a PL function-value rate, while the inner ascent step may decrease more slowly than the outer descent step: ρt=O(γt)\rho_t=O(\sqrt{\gamma_t}) is sufficient. An analogous result holds for stochastic nn-SAM when independent batches are used for ascent and descent.

  10. Knowl 10 — Deep-network training converges for SAM, with or without gradient normalization

    empirical result

    On CIFAR-10 and CIFAR-100, long training runs of ResNet models show that both ERM and SAM fit the training data and approach nearly zero training loss. On CIFAR-10, ERM reaches training loss 0.0013±0.000020.0013\pm0.00002 and test error 4.75%±0.14%4.75\%\pm0.14\%, whereas standard SAM reaches training loss 0.0034±0.00040.0034\pm0.0004 and test error 3.94%±0.09%3.94\%\pm0.09\%. The improvement in test error therefore does not require SAM to stop before fitting clean training data.

    Using a constant, unnormalized perturbation radius gives similar behavior. On CIFAR-10, standard gradient-normalized SAM obtains 3.94%±0.09%3.94\%\pm0.09\% test error compared with 4.15%±0.16%4.15\%\pm0.16\% for constant-radius SAM. On CIFAR-100, the corresponding errors are 19.22%±0.38%19.22\%\pm0.38\% and 19.30%±0.38%19.30\%\pm0.38\%. The best radii differ between the two implementations, so removing normalization without retuning the radius can be suboptimal. These results support the theoretical conclusion that SAM can converge even when its inner perturbation radius is not decreased in lockstep with the outer learning rate.

  11. Knowl 11 — Label noise exposes a limitation: SAM can eventually overfit

    limitation

    When 60%60\% of CIFAR-10 or CIFAR-100 training labels are replaced by fixed random labels, SAM still improves test error over ERM during part of training. However, continued optimization eventually causes SAM to fit the noisy examples, after which test error increases. The beneficial behavior is therefore not equivalent to preventing convergence to zero training loss.

    For noisy-label data, SAM requires early stopping, either explicitly through a validation set or implicitly through a predetermined training budget. This experiment also shows that SAM's beneficial effect can occur along the optimization trajectory and is not restricted to neighborhoods of a final sharp or flat minimum. Consequently, convergence guarantees for the training objective do not imply uniformly improved generalization when the training labels contain substantial noise.

Coverage note — Proof-only lemmas, detailed derivations, and secondary ablations on model width, batch-normalization measurement discrepancies, and decreasing perturbation schedules were omitted because they support or qualify the main contributions without adding separate load-bearing results.

References

  1. 1.Guy Blanc, Neha Gupta, Gregory Valiant, and Paul Valiant. Implicit regularization for deep neural networks driven by an Ornstein-Uhlenbeck like process. In COLT, 2020.
  2. 2.Lénaïc Chizat and Francis Bach. Implicit bias of gradient descent for wide two-layer neural networks trained with the logistic loss. In COLT, 2020.
  3. 3.Laurent Dinh, Razvan Pascanu, Samy Bengio, and Yoshua Bengio. Sharp minima can generalize for deep nets. In ICML, pp. 1019–1028. PMLR, 2017.
  4. 4.Gintare Karolina Dziugaite and Daniel Roy. Entropy-sgd optimizes the prior of a pac-bayes bound: Generalization properties of entropy-sgd and data-dependent priors. In ICML, 2018.
  5. 5.Pierre Foret, Ariel Kleiner, Hossein Mobahi, and Behnam Neyshabur. Sharpness-aware minimization for efficiently improving generalization. In ICLR, 2021.
  6. 6.Saeed Ghadimi and Guanghui Lan. Stochastic first- and zeroth-order methods for nonconvex stochastic programming. SIAM Journal on Optimization, 23(4):2341–2368, 2013.
  7. 7.Robert Mansel Gower, Nicolas Loizou, Xun Qian, Alibek Sailanbayev, Egor Shulgin, and Peter Richtarik. SGD: General analysis and improved rates. In ICML, 2019.
  8. 8.Priya Goyal, Piotr Dollar, Ross Girshick, Pieter Noordhuis, Lukasz Wesolowski, Aapo Kyrola, Andrew Tulloch, Yangqing Jia, and Kaiming He. Accurate, large mini-batch sgd: Training imagenet in 1 hour. arXiv preprint arXiv:1706.02677, 2017.
  9. 9.Alex Graves, Abdel-rahman Mohamed, and Geoffrey Hinton. Speech recognition with deep recurrent neural networks. In 2013 IEEE ICASSP, 2013.
  10. 10.Kosuke Haruki, Taiji Suzuki, Yohei Hamakawa, Takeshi Toda, Ryuji Sakai, Masahiro Ozawa, and Mitsuhiro Kimura. Gradient noise convolution (GNC): Smoothing loss function for distributed large-batch sgd. arXiv preprint arXiv:1906.10822, 2019.
  11. 11.Haowei He, Gao Huang, and Yang Yuan. Asymmetric valleys: Beyond sharp and flat local minima. In NeurIPS, 2019.
  12. 12.Kaiming He, Xiangyu Zhang, Shaoqing Ren, and Jian Sun. Identity mappings in deep residual networks. In ECCV, 2016.
  13. 13.Sepp Hochreiter and Jürgen Schmidhuber. Simplifying neural nets by discovering flat minima. In NeurIPS, 1995.
  14. 14.Elad Hoffer, Itay Hubara, and Daniel Soudry. Train longer, generalize better: closing the generalization gap in large batch training of neural networks. In NeurIPS, 2017.
  15. 15.Sergey Ioffe and Christian Szegedy. Batch normalization: Accelerating deep network training by reducing internal covariate shift. In ICML, 2015.
  16. 16.Yiding Jiang, Behnam Neyshabur, Hossein Mobahi, Dilip Krishnan, and Samy Bengio. Fantastic generalization measures and where to find them. In ICLR, 2019.
  17. 17.Kam-Chuen Jim, C Lee Giles, and Bill G Horne. An analysis of noise in recurrent neural networks: convergence and generalization. In IEEE Transactions on Neural Networks, 1996.
  18. 18.Hamed Karimi, Julie Nutini, and Mark Schmidt. Linear convergence of gradient and proximal-gradient methods under the Polyak-Lojasiewicz condition. In Joint European Conference on Machine Learning and Knowledge Discovery in Databases, 2016.
  19. 19.Nitish Shirish Keskar, Dheevatsa Mudigere, Jorge Nocedal, Mikhail Smelyanskiy, and Ping Tak Peter Tang. On large-batch training for deep learning: Generalization gap and sharp minima. In ICLR, 2016.
  20. 20.Galina Korpelevich. Extragradient method for finding saddle points and other problems. In Matekon, 1977.
  21. 21.Alex Krizhevsky and Geoffrey Hinton. Learning multiple layers of features from tiny images. Technical Report, 2009.
  22. 22.Jungmin Kwon, Jeongseop Kim, Hyunseo Park, and In Kwon Choi. ASAM: Adaptive sharpness-aware minimization for scale-invariant learning of deep neural networks. In ICML, 2021.
  23. 23.Yann A LeCun, Leon Bottou, Genevieve B Orr, and Klaus-Robert Müller. Efficient backprop. In Neural networks: Tricks of the trade, pp. 9–48. Springer, 2012.
  24. 24.Tao Lin, Lingjing Kong, Sebastian Stich, and Martin Jaggi. Extrapolation for large-batch training in deep learning. In ICML, 2020.
  25. 25.Shengchao Liu, Dimitris Papailiopoulos, and Dimitris Achlioptas. Bad global minima exist and SGD can reach them. In NeurIPS, 2019.
  26. 26.Rotem Mulayoff and Tomer Michaeli. Unique properties of flat minima in deep networks. In ICML, 2020.
  27. 27.Alan F Murray and Peter J Edwards. Synaptic weight noise during MLP learning enhances fault-tolerance, generalization and learning trajectory. In NeurIPS, 1993.
  28. 28.Vaishnavh Nagarajan and J Zico Kolter. Uniform convergence may be unable to explain generalization in deep learning. In NeurIPS, 2019.
  29. 29.Preetum Nakkiran, Gal Kaplun, Yamini Bansal, Tristan Yang, Boaz Barak, and Ilya Sutskever. Deep double descent: Where bigger models and more data hurt. In ICLR, 2020.
  30. 30.Yurii Nesterov. Introductory Lectures on Convex Optimization. Kluwer Academic, 2004.
  31. 31.Gergely Neu. Information-theoretic generalization bounds for stochastic gradient descent. In COLT, 2021.
  32. 32.Behnam Neyshabur, Ryota Tomioka, and Nathan Srebro. In search of the real inductive bias: On the role of implicit regularization in deep learning. In ICLR workshops, 2015.
  33. 33.Behnam Neyshabur, Srinadh Bhojanapalli, David McAllester, and Nati Srebro. Exploring generalization in deep learning. In NeurIPS, 2017.
  34. 34.Scott Pesme, Loucas Pillaud-Vivien, and Nicolas Flammarion. Implicit bias of sgd for diagonal linear networks: a provable benefit of stochasticity. In NeurIPS, 2021.
  35. 35.Leslie Rice, Eric Wong, and J Zico Kolter. Overfitting in adversarially robust deep learning. In ICML, 2020.
  36. 36.Nitish Srivastava, Geoffrey Hinton, Alex Krizhevsky, Ilya Sutskever, and Ruslan Salakhutdinov. Dropout: a simple way to prevent neural networks from overfitting. JMLR, 2014.
  37. 37.Ziqiao Wang and Yongyi Mao. On the generalization of models trained with SGD: Information-theoretic bounds and implications. In ICLR, 2022.
  38. 38.Michael L. Waskom. Seaborn: statistical data visualization. Journal of Open Source Software, 6(60):3021, 2021. doi: 10.21105/joss.03021. URL https://doi.org/10.21105/joss.03021.
  39. 39.Wei Wen, Yandan Wang, Feng Yan, Cong Xu, Chunpeng Wu, Yiran Chen, and Hai Li. Smoothout: Smoothing out sharp minima to improve generalization in deep learning. arXiv preprint arXiv:1805.07898, 2018.
  40. 40.Blake Woodworth, Suriya Gunasekar, Jason D. Lee, Edward Moroshko, Pedro Savarese, Itay Golan, Daniel Soudry, and Nathan Srebro. Kernel and rich regimes in overparametrized models. In COLT, 2020.
  41. 41.Dongxian Wu, Shu-tao Xia, and Yisen Wang. Adversarial weight perturbation helps robust generalization. In NeurIPS, 2020.
  42. 42.Yuxin Wu and Kaiming He. Group normalization. In ECCV, 2018.
  43. 43.Chiyuan Zhang, Samy Bengio, Moritz Hardt, Benjamin Recht, and Oriol Vinyals. Understanding deep learning requires rethinking generalization. In ICLR, 2017.
  44. 44.Yaowei Zheng, Richong Zhang, and Yongyi Mao. Regularizing neural networks via adversarial model perturbation. In CVPR, 2021.

Citation

MLA
Andriushchenko, M., and N. Flammarion. “Towards Understanding Sharpness-Aware Minimization”. International Conference on Machine Learning, vol. 162, 2022, pp. 639–68, https://proceedings.mlr.press/v162/andriushchenko22a.html.
APA
Andriushchenko, M., & Flammarion, N. (2022). Towards Understanding Sharpness-Aware Minimization. International Conference on Machine Learning, 162, 639–668. https://proceedings.mlr.press/v162/andriushchenko22a.html
Chicago
Andriushchenko, M., and N. Flammarion. 2022. “Towards Understanding Sharpness-Aware Minimization”. International Conference on Machine Learning 162: 639–68. https://proceedings.mlr.press/v162/andriushchenko22a.html.
Harvard
Andriushchenko, M. and Flammarion, N. (2022) “Towards Understanding Sharpness-Aware Minimization”, International Conference on Machine Learning. PMLR, pp. 639–668. Available at: https://proceedings.mlr.press/v162/andriushchenko22a.html.
Vancouver
1. Andriushchenko M, Flammarion N (2022) Towards Understanding Sharpness-Aware Minimization. In: International Conference on Machine Learning. PMLR, pp 639–668

BibTeX

@InProceedings{pmlr-v162-andriushchenko22a,
  title = 	 {Towards Understanding Sharpness-Aware Minimization},
  author =       {Andriushchenko, Maksym and Flammarion, Nicolas},
  booktitle = 	 {Proceedings of the 39th International Conference on Machine Learning},
  pages = 	 {639--668},
  year = 	 {2022},
  editor = 	 {Chaudhuri, Kamalika and Jegelka, Stefanie and Song, Le and Szepesvari, Csaba and Niu, Gang and Sabato, Sivan},
  volume = 	 {162},
  series = 	 {Proceedings of Machine Learning Research},
  month = 	 {17--23 Jul},
  publisher =    {PMLR},
  pdf = 	 {https://proceedings.mlr.press/v162/andriushchenko22a/andriushchenko22a.pdf},
  url = 	 {https://proceedings.mlr.press/v162/andriushchenko22a.html},
  abstract = 	 {Sharpness-Aware Minimization (SAM) is a recent training method that relies on worst-case weight perturbations which significantly improves generalization in various settings. We argue that the existing justifications for the success of SAM which are based on a PAC-Bayes generalization bound and the idea of convergence to flat minima are incomplete. Moreover, there are no explanations for the success of using m-sharpness in SAM which has been shown as essential for generalization. To better understand this aspect of SAM, we theoretically analyze its implicit bias for diagonal linear networks. We prove that SAM always chooses a solution that enjoys better generalization properties than standard gradient descent for a certain class of problems, and this effect is amplified by using m-sharpness. We further study the properties of the implicit bias on non-linear networks empirically, where we show that fine-tuning a standard model with SAM can lead to significant generalization improvements. Finally, we provide convergence results of SAM for non-convex objectives when used with stochastic gradients. We illustrate these results empirically for deep networks and discuss their relation to the generalization behavior of SAM. The code of our experiments is available at https://github.com/tml-epfl/understanding-sam.}
}
Metadata:DOI registry

Source Code

This paper has an official code repository available. Click below to access the source code.

View Repository

Access the Paper

This paper is available from its original source. Click below to access the PDF.

Open PDF
License: https://creativecommons.org/licenses/by/4.0/