Optimization2021intermediate10 min read

Sharpness-Aware Minimization for Efficiently Improving Generalization

تصغير مُدرك للحِدّة لتحسين التعميم بكفاءة

Foret, P. · Kleiner, A. · Mobahi, H. · Neyshabur, B. — ICLR

The problem

In overparameterized deep networks, minimizing alone provides few guarantees about . A model can reach zero training error yet perform poorly on unseen data because it converged to a sharp minimum — a region of space where the loss changes rapidly with small perturbations. The geometry of the , not just the loss value, determines how well a model generalizes.

The contribution

SAM introduces a optimization that simultaneously minimizes loss value and loss . At each step, it first computes the worst-case weight within a bounded neighborhood using a single ascent step, then takes a step from that perturbed point. This encourages to where the loss is uniformly low across the entire neighborhood. The paper proves a PAC-Bayesian generalization bound connecting sharpness to generalization gap, and shows empirically that SAM achieves state-of-the-art results on CIFAR-10/100, ImageNet, and finetuning tasks while natively providing robustness to .

The impact

SAM became one of the most widely adopted training techniques in deep learning after its publication. It spawned a family of variants — ASAM, LookSAM, Efficient SAM, Fisher SAM — and shifted the optimization community's focus from loss-value-only training toward geometry-aware optimization. Its insight that flat minima generalize better, made actionable through a practical two-step algorithm, influenced the design of modern optimizers including Lion. SAM is now standard in training Vision Transformers and large-scale models where generalization is critical.

Imagine two hikers each searching for the lowest point in a mountain range. The first hiker uses GPS coordinates alone — she walks straight downhill and stops at the first valley floor, even if it's a narrow crack between two cliffs. One strong gust of wind and she's climbing a wall.

The second hiker is smarter. Before planting her flag, she walks a full circle around herself. If the ground rises steeply in any direction, she moves on. She only settles in a wide, flat basin where she can stumble in any direction and still be on low ground.

SAM is the second hiker. Standard gradient descent finds a minimum. SAM finds a stable minimum — one that stays low even when the ground shakes.

The problem: sharp valleys generalize poorly

When we train a , we minimize a over the training data. The searches through a high-dimensional parameter space looking for points where the loss is low. But not all low-loss points are equal.

Some minima sit at the bottom of narrow, sharp valleys. The loss is low exactly at that point, but rises steeply in every direction. Such a minimum is fragile: even a tiny shift in the data distribution — the difference between training data and test data — moves the parameters slightly, and the loss explodes. This is in geometric language.

Other minima sit in wide, flat basins. The loss is low not just at the exact parameter values found during training, but across an entire neighborhood. These flat minima are robust: small shifts in data or parameters barely change the loss. Empirically, flat minima correlate strongly with good generalization.

The question is: how do we make the optimizer prefer flat minima over sharp ones?

Open in Lab
Drag the parameter slider to see how the loss changes around sharp vs flat minima. Notice that the flat minimum stays low over a wider range.
The demo wakes as you arrive…

The SAM objective: minimizing the worst neighbor

Standard training solves a simple optimization: find parameters ww that minimize the training loss LS(w)L_S(w). SAM changes the question. Instead of asking "what is the loss at ww?", it asks "what is the worst loss in a small ball around ww?"

The intuition is like stress-testing a bridge. A standard optimizer checks whether the bridge holds under normal load. SAM checks whether it holds under the worst load within a reasonable range. If the bridge survives the stress test, it will certainly survive normal conditions.

Formally, SAM replaces the standard objective with:

min⁡w  max⁡∥ϵ∥p≤ρ  LS(w+ϵ)\min_{w} \; \max_{\|\epsilon\|_p \leq \rho} \; L_S(w + \epsilon)
SAM objective — minimize the worst-case loss in a ρ-neighborhood — The inner maximization finds the worst perturbation ϵ\epsilon of the weights within a ball of radius ρ\rho. The outer minimization then adjusts the weights to make even that worst case as small as possible. The parameter ρ\rho controls how wide the neighborhood is — how far the "stress test" reaches.

This is a minimax problem: minimize over ww, maximize over ϵ\epsilon. The two objectives pull in opposite directions — ϵ\epsilon tries to find a weakness, and ww tries to eliminate it. The result is parameters that have no sharp weakness in any direction.

But solving the inner maximization exactly is computationally expensive. SAM's key practical insight is that a single-step gradient ascent approximation works remarkably well.

The algorithm: a two-step dance

SAM's elegance lies in its simplicity. Each training step has two phases — think of it as "probe, then correct."

Step 1 — Probe (gradient ascent): Compute the gradient of the loss at the current weights wtw_t. Use this gradient to find the worst-case perturbation — the direction that increases the loss most within the ρ\rho-ball:

ϵ^(wt)=ρ  ∇wLS(wt)∥∇wLS(wt)∥\hat{\epsilon}(w_t) = \rho \; \frac{\nabla_w L_S(w_t)}{\|\nabla_w L_S(w_t)\|}
Worst-case perturbation via first-order approximation — The gradient points in the direction of steepest ascent. Normalizing it and scaling by ρ\rho gives us the point on the ρ\rho-ball boundary where the loss is approximately highest. This is a first-order Taylor approximation — cheap but effective.

Step 2 — Correct (gradient descent): Now compute the gradient at the perturbed point wt+ϵ^w_t + \hat{\epsilon}, and use it to update the actual weights:

wt+1=wt−η  ∇wLS(w)∣w=wt+ϵ^(wt)w_{t+1} = w_t - \eta \; \nabla_w L_S(w)\big|_{w = w_t + \hat{\epsilon}(w_t)}
SAM weight update — gradient at the perturbed point — Instead of descending from the current position (like standard SGD), SAM descends from the worst-case neighbor. This means the update accounts for the *curvature* around the current point, not just the slope at the point itself. The learning rate η\eta controls the step size as usual.
Open in Lab
Watch SAM's two-step process: first perturb to the worst neighbor (red arrow), then compute the gradient there and update the weights (blue arrow). Compare with standard SGD which only uses the gradient at the current point.
The demo wakes as you arrive…

The total computational cost is roughly 2× that of standard — one forward and to compute ϵ^\hat{\epsilon}, and another to compute the gradient at w+ϵ^w + \hat{\epsilon}. This doubles the per-step cost, but the authors show that the improved generalization more than compensates for the overhead, and SAM often converges in fewer epochs.

SAM pseudocode — the full training steppython

Simplified to show the idea — not the real implementation.

# SAM: Sharpness-Aware Minimization (one training step)
# Wraps any base optimizer (SGD, Adam, etc.)

def sam_step(model, loss_fn, batch, rho=0.05, base_optimizer):
    # Step 1: Compute gradient at current weights
    loss = loss_fn(model(batch.x), batch.y)
    loss.backward()

    # Compute worst-case perturbation epsilon_hat
    with torch.no_grad():
        grad_norm = torch.stack([
            p.grad.norm() for p in model.parameters()
        ]).norm()
        for p in model.parameters():
            # Scale gradient to lie on the rho-ball
            epsilon = rho * p.grad / grad_norm
            p.add_(epsilon)          # w -> w + epsilon

    # Step 2: Compute gradient at perturbed point
    model.zero_grad()
    loss_perturbed = loss_fn(model(batch.x), batch.y)
    loss_perturbed.backward()

    # Remove perturbation and update with SAM gradient
    with torch.no_grad():
        for p in model.parameters():
            epsilon = rho * p.grad / grad_norm
            p.sub_(epsilon)          # w + epsilon -> w
        base_optimizer.step()        # w -> w - eta * grad_perturbed

The theory: why flat minima generalize

SAM is not just a heuristic — it is grounded in a formal generalization bound derived from PAC-Bayesian theory. The bound provides a mathematical guarantee: if the training loss is low and the loss landscape is flat around the solution, then the gap between training loss and test loss is bounded.

PAC-Bayesian theory works by considering a distribution over possible parameter values rather than a single point. Think of it as asking: "if I wobble the weights randomly in a small region, does the model still work?" If yes, the model is robust to perturbation, and the theory guarantees it will generalize.

The key insight is that the term inside the bound that captures loss sharpness is precisely the SAM objective — the worst-case loss in a neighborhood:

LD(w)≤max⁡∥ϵ∥≤ρLS(w+ϵ)+h ⁣(∥w∥22/ρ2)L_\mathcal{D}(w) \leq \max_{\|\epsilon\| \leq \rho} L_S(w + \epsilon) + h\!\left(\|w\|_2^2 / \rho^2\right)
PAC-Bayesian generalization bound connecting sharpness to generalization — The population loss LDL_\mathcal{D} (test performance) is bounded by the worst-case training loss in the neighborhood (the SAM objective) plus a complexity term hh that grows with the ratio of weight norm to perturbation radius. Minimizing both terms — keeping the landscape flat and the weights small — bounds the generalization gap.

Measuring sharpness: the m-sharpness variant

The paper introduces a refined notion called m-sharpness. Instead of computing one worst-case perturbation for the entire training set, m-sharpness computes it per batch. Each gets its own worst-case perturbation, and the sharpness is the average over batches.

Why does this matter? Because per-batch perturbations are more targeted. A single global perturbation might average out across different data regions. Per-batch perturbations probe the landscape more precisely in each local neighborhood, catching fine-grained sharpness that a global perturbation would miss.

The authors find that m-sharpness correlates more strongly with generalization than the global sharpness measure, and that using mini-batch perturbations during training (which SAM naturally does, since it operates on batches) is key to SAM's effectiveness.

Open in Lab
Compare global sharpness (one perturbation for all data) vs m-sharpness (per-batch perturbation). m-sharpness captures local geometry more precisely.
The demo wakes as you arrive…

Results: SAM across benchmarks

SAM delivers consistent improvements across a wide range of tasks and architectures. On CIFAR-10, SAM achieves approximately 96.5% accuracy with ResNet, compared to ~95.4% for standard training — a gap of over 1%. On CIFAR-100, the improvement is even larger, jumping from ~78% to ~80%.

On ImageNet, SAM with EfficientNet achieves state-of-the-art results, pushing top-1 accuracy beyond what the same architecture achieves with standard optimizers.

Perhaps most impressively, SAM shines on finetuning tasks — FGVC (Fine-Grained Visual Categorization) benchmarks like Oxford Pets, Stanford Cars, and Flowers. These scenarios test how well a pretrained model adapts to new domains, and SAM consistently outperforms standard finetuning.

Open in Lab
SAM vs standard training across benchmarks. Toggle between datasets to see the improvement.
The demo wakes as you arrive…

Bonus: native robustness to label noise

A surprising finding: SAM is naturally robust to label noise, without being designed for it. When a fraction of training labels are randomly corrupted, standard SGD quickly overfits to the noise — it memorizes the wrong labels. SAM resists this memorization.

The explanation connects back to sharpness. Noisy labels create sharp, isolated minima in the loss landscape — the model has to contort itself into a narrow region of parameter space to simultaneously fit both clean and corrupted labels. SAM's sharpness penalty discourages the optimizer from entering these sharp regions.

At 80% label corruption on CIFAR-10, standard training collapses while SAM retains substantial accuracy, performing on par with methods specifically designed for learning with noisy labels, like MentorMix.

Open in Lab
Increase the label noise fraction and compare SAM vs SGD. SAM maintains accuracy where SGD collapses.
The demo wakes as you arrive…

The big picture: geometry-aware optimization

SAM represents a paradigm shift in how we think about training neural networks. Before SAM, optimizers focused exclusively on the loss value — find the lowest point. SAM showed that the loss landscape matters just as much — not just where you are, but what the terrain looks like around you.

This insight spawned a rich research direction. ASAM (Adaptive SAM) made the perturbation scale-invariant by normalizing per parameter. LookSAM reduced the computational cost by periodically skipping the perturbation step. Fisher SAM used the Fisher information matrix to guide perturbations more intelligently.

The lesson extends beyond SAM itself: geometry-aware optimization is now recognized as a fundamental ingredient for training large models that generalize well. Modern optimizers like Lion incorporate ideas from the same family of insights about loss landscape geometry that SAM helped popularize.

Timeline: from sharpness theory to SAM and beyond

  1. 1997

    Flat minima hypothesis (Hochreiter & Schmidhuber)

    First formal argument that flat minima of the training loss generalize better than sharp ones. Introduced a minimum description length approach to quantify flatness.

  2. 2017

    Large-batch training reveals sharpness problem (Keskar et al.)

    Showed that large-batch SGD converges to sharp minima, explaining the generalization gap between small-batch and large-batch training. Established sharpness as a measurable quantity correlated with generalization.

  3. 2018

    Loss landscape visualization (Li et al.)

    Introduced filter-normalized visualization of loss landscapes, revealing that architectures with skip connections (ResNet) have dramatically flatter landscapes than those without.

  4. 2021

    SAM (Foret et al.) — this paper

    Proposed a practical minimax optimizer that directly targets flat minima. Proved a PAC-Bayesian bound connecting sharpness to generalization, and demonstrated state-of-the-art results across benchmarks with native label noise robustness.

  5. 2021

    ASAM — scale-invariant sharpness (Kwon et al.)

    Identified that SAM's sharpness measure is scale-dependent. Proposed Adaptive SAM with element-wise normalization, improving results on tasks where parameter scales vary.

  6. 2022

    SAM for Vision Transformers (Chen et al.)

    Showed that SAM is particularly effective for Vision Transformers, which are more prone to sharp minima than convolutional networks. SAM became near-essential for ViT training.

  7. 2023

    Lion optimizer and geometry-aware training

    Google's Lion optimizer, discovered via program search, shares SAM's philosophy of considering the loss landscape geometry. The era of loss-value-only optimization gave way to geometry-aware training.

CitationForet, Kleiner, Mobahi, Neyshabur. Sharpness-Aware Minimization for Efficiently Improving Generalization. ICLR, 2021.

Terms in this paper