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?
The SAM objective: minimizing the worst neighbor
Standard training solves a simple optimization: find parameters that minimize the training loss . SAM changes the question. Instead of asking "what is the loss at ?", it asks "what is the worst loss in a small ball around ?"
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:
This is a minimax problem: minimize over , maximize over . The two objectives pull in opposite directions — tries to find a weakness, and 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 . Use this gradient to find the worst-case perturbation — the direction that increases the loss most within the -ball:
Step 2 — Correct (gradient descent): Now compute the gradient at the perturbed point , and use it to update the actual weights:
The total computational cost is roughly 2× that of standard — one forward and to compute , and another to compute the gradient at . 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.
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:
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.
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.
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.
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
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.
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.
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.
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.
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.
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.
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
- Sharpnessالحدّة
- Flat Minimaالحدود الدنيا المسطّحة
- Generalizationالتعميم
- Loss Landscapeسطح الخسارة
- Perturbationاضطراب
- Minimaxالأصغري-الأعظمي
- PAC-Bayesنظرية PAC-Bayes
- Regularizationالضبط الهيكلي
- Overfittingفرط التخصيص
- Label Noiseتشويش التسميات
- Gradient Descentالانحدار التدريجي
- Stochastic Gradient Descent (SGD)الانحدار التدريجي العشوائي
- Weight Decayاضمحلال الأوزان
- Optimizerالـمُحسِّن
- Hessianمصفوفة هيسي