Flat loss landscapes have long been linked to better generalization in neural networks. However, its role as a causal mechanism for generalization is less established. Grokking provides an unique testbed to understand this distinction: models are prone to fit observed data using non-generalizing structure and remain in that regime for prolonged periods, transitioning to generalization only under particular training conditions. In this work, we study whether flat loss landscapes can act as a driving mechanism in this transition. While recent work has argued for flatness as a necessary geometric condition for this transition, we find that biasing training toward flatter solutions using sharpness-aware minimization (SAM) is insufficient to reliably induce this transition, despite producing flatter solutions. However, when SAM is paired with mechanisms that drive generalization such as weight decay, an interesting property emerges: SAM can accelerate the transition to generalizing solutions by up to 4x at the epoch-level. We theoretically untangle this relationship between SAM and weight decay using a minimal interpolating two-layer ReLU model with both memorizing and generalizing solutions. We show that even in this simple setup, flatness alone cannot distinguish a memorizing solution from a generalizing one, while weight decay favors generalizing solutions. However, under a local stability analysis, there exists a window where a memorizing interpolant is locally stable under gradient descent but unstable under SAM in the low-norm regime, which can explain SAM's ability to accelerate this transition. Overall, our results provide a more interpretable account of the role of flatness in driving generalization, especially in settings where models are vulnerable to minimizing loss through learning non-generalizing structure.
Figures & tables
Figure 1 : Flatness regularization accelerates generalization in grokking where grokking already occurs. It does not, however, trigger generalization alone. Lighter colors correspond to stronger regularization. Stronger regularization is linked with faster generalization.
Figure 2 : Flatness alone does not explain the generalization transition. Red bars (left axis) show the percentage of modular arithmetic runs reaching the generalization transition; blue bars (right axis, log scale) show the mean final-checkpoint sharpness λmax of the training-loss Hessian. SAM-only finds by far the flattest solutions yet no run generalizes, while weight-decay-based runs generalize in 99% of cases despite substantially sharper final solutions. No regularization pools 20 runs (4 settings × 5 seeds); SAM-only pools 180 runs across several values of ρ ; Weight Decay methods pool WD and SAM+WD runs across all values of ρ (200 runs).
Figure 3 : SAM accelerates generalization. Acceleration of the grokking transition under SAM+WD relative to WD alone, as a function of the SAM radius ρ , across Transformer architectures on modular addition. Acceleration is the median per-seed ratio of the WD grokking epoch to the SAM+WD grokking epoch; the dashed line marks the WD baseline ( 1× ). Shaded regions show ±1 standard deviation of the SAM+WD grokking epoch across 5 seeds. Blue stars mark increases in the scale-normalized sharpness λmax/∥θ∥22 , indicating where first-order Taylor estimates from SAM begin to degrade and can no longer enforce flatness.
Figure 4 : Relationship between SAM radius ρ and grokking speedup across models trained on CIFAR-10 (left) and SST-5 (right). Larger ρ generally induces stronger flatness regularization and leads to a faster grokking transition when combined with weight decay until ρ is too large. Shaded regions indicate one standard deviation above and below the median across 3 seeds.
Figure 5 : Relationship between SAM radius ρ and grokking speedup across models trained on MNIST and IMDB as defined by Omnigrok. Larger ρ generally induces stronger flatness regularization and leads to faster grokking transition when combined with weight decay until ρ is too large. Shaded regions indicate one standard deviation above and below the median across 3 seeds.
Appendix figures & tables3 assets
Supplementary material from the paper’s appendix.
Appendix
Condition
ρ
Grok
Grok epoch
Val. acc.
Speedup
Han RF ↓
λmax50↓
λmax50/∥θ∥22↓
No reg
–
0/5
–
0.307 ± 0.0242
–
48.02 ± 37.47
0.00213 ± 0.00157
4.54e-07 ± 3.37e-07
WD
–
5/5
5538 ± 2976.6
1 ± 4.06e-04
1x
9.247 ± 10.17
0.0547 ± 0.0553
1.25e-05 ± 7.95e-06
SAM-only
0.01
0/5
–
0.307 ± 0.0144
–
8.757 ± 3.058
6.67e-04 ± 2.51e-04
1.56e-07 ± 6.39e-08
SAM-only
0.05
0/5
–
0.351 ± 0.0711
–
1.026 ± 0.483
1.58e-04 ± 9.46e-05
4.08e-08 ± 2.42e-08
SAM-only
0.2
0/5
–
0.433 ± 0.144
–
5.152 ± 6.666
8.43e-04 ± 0.00126
2.39e-07 ± 3.61e-07
SAM-only
0.3
0/5
–
0.413 ± 0.172
–
10.99 ± 15.36
0.00177 ± 0.00263
5.02e-07 ± 7.56e-07
Appendix
Table 2 : Modular-addition results for 1-layer Transformers. Grokking and validation summaries use matched five-seed runs. Speedup is the median paired-seed SAM+WD speedup over WD. Han RF is the official Han et al. relative-flatness metric. Hessian columns use the previously computed three-seed Hessian-50 posthoc table matched by task, condition, and ρ .
Condition
ρ
Grok
Grok epoch
Val. acc.
Speedup
Han RF ↓
λmax50↓
λmax50/∥θ∥22↓
No reg
–
0/5
–
0.311 ± 0.032
–
1.23e+04 ± 2.12e+04
7.04 ± 12.14
0.00107 ± 0.00184
WD
–
5/5
11266 ± 9705.0
0.876 ± 0.174
1x
6.44e+06 ± 1.99e+06
5.91e+04 ± 2.06e+04
83.45 ± 63.48
SAM-only
0.01
0/5
–
0.41 ± 0.171
–
3.877 ± 2.619
0.00403 ± 0.00223
7.42e-07 ± 4.63e-07
SAM-only
0.05
0/5
–
0.44 ± 0.177
–
900.7 ± 1558.1
0.48 ± 0.826
7.57e-05 ± 1.30e-04
SAM-only
0.2
0/5
–
0.39 ± 0.119
–
2.052 ± 1.191
0.00105 ± 6.43e-04
1.75e-07 ± 1.05e-07
SAM-only
0.3
0/5
–
0.383 ± 0.101
–
13.66 ± 20.79
0.00301 ± 0.00365
4.87e-07 ± 5.87e-07
Appendix
Table 3 : Modular-addition results for 2-layer Transformers. Grokking and validation summaries use matched five-seed runs. Speedup is the median paired-seed SAM+WD speedup over WD. Han RF is the official Han et al. relative-flatness metric. Hessian columns use the previously computed three-seed Hessian-50 posthoc table matched by task, condition, and ρ .
Experiment
Runs
Seeds
Wall time / run
Total GPU-hours
Modular Arithmetic
400
5
0.2–41 min
73
CIFAR-10
96
3
0.9–2.0 h
113
SST-5
168
3
0.35–3.82 h
290
MNIST
110
3
–
2
IMDb
261
3
0.5–1.2 h
113
Appendix
Table 4 : We report the number of total runs, seeds, wall-time per run, and total GPU-hours to recreate all results on each task independently.
A widely held intuition in deep learning is that stochastic gradient descent (SGD) implicitly favors flat minima and that flat minima generalize better, but standard Euclidean measures of flatness such as the trace or maximum eigenvalue of the loss Hessian are not invariant under reparametrizations that preserve the network function, which undermines the theoretical foundations of this narrative. In this study we resolve this issue by grounding flatness in the Riemannian geometry of the statistical manifold induced by the Fisher Information Matrix (FIM). We define Riemannian sharpness mathematically and prove that it is invariant under smooth, function-preserving reparametrizations, which directly addresses the critique of Dinh et al. in the paper ``Sharp minima can generalize for deep nets''.We note that this invariance is a property of the true FIM; the diagonal empirical estimator used in practice (and in all experiments below) inherits invariance only approximately, and exact invariance under arbitrary reparametrizations would require structured estimators such as K-FAC. We formalize the gradient noise of mini-batch SGD as having a covariance structure proportional to the FIM, derive the stationary distribution of the resulting stochastic differential equation, and then show that the probability mass is exponentially concentrated at Riemannian-flat minima. A PAC-Bayes generalization bound controlled explicitly by SR formally links this geometric bias to test performance. Our experiments on MNIST and CIFAR-10 confirm that SR reliably tracks generalization in ways that Euclidean sharpness does not, and that its scaling with η/B matches the theoretical predictions. Together these results provide a rigorous, reparametrization-invariant account of why flat minima generalize.
Md Sakir Ahmed, Kumaresh Sarmah, Hemen Dutta
Department of Electronics and Communication Technology Gauhati University Guwahati, Assam, India · Department of Mathematics Gauhati University Guwahati, Assam, India
The sharpness-aware minimization (SAM) algorithm and its variants, including gap guided SAM (GSAM), have been successful at improving the generalization capability of deep neural network models by finding flat local minima of the empirical loss in training. Meanwhile, it has been shown theoretically and practically that increasing the batch size or decaying the learning rate avoids sharp local minima of the empirical loss. In this paper, we consider the GSAM algorithm with increasing batch sizes or decaying learning rates, such as cosine annealing or linear learning rate, and theoretically show its convergence. Moreover, we numerically compare SAM (GSAM) with and without an increasing batch size and conclude that using an increasing batch size { achieves a lower worst-case ℓ∞ adaptive sharpness} than compared with using a constant batch size and learning rate.
Sharpness-Aware Minimization (SAM) improves generalization by seeking parameters whose loss is robust to local adversarial perturbations, but the quantitative mechanism underlying its implicit bias toward flat minima remains unclear. In particular, the perturbation radius ρ is typically treated as an isolated tuning parameter, despite defining the neighborhood in which SAM measures sharpness. We analyze mini-batch SAM near an interpolating minimum through linear stability. Under local linearization and gradient-noise alignment assumptions, we prove that every linearly stable minimum satisfies λmax≤3bΓ/(2ρη2), where λmax is the largest Hessian eigenvalue, b is the batch size, η is the learning rate, and Γ bounds the gradient norm. The bound quantitatively characterizes SAM's implicit flatness bias: holding the other quantities fixed, a smaller batch size, a larger learning rate, or a larger radius restricts linearly stable SAM to flatter minima. It also exposes a necessary trade-off: ρ should be large enough to promote flatness, yet remain local enough to preserve the approximation and stable training. We validate this prediction in a controlled study of 900 models on CIFAR-100 with ResNet-18 and VGG-19, where increasing ρ is consistently associated with a smaller largest Hessian eigenvalue across batch-size and learning-rate settings. Finally, we instantiate the analysis in Taylor-Locality Controlled SAM (TLC-SAM), which adjusts ρ using the observed Taylor-approximation error and further reduces the top Hessian eigenvalue relative to fixed-radius SAM. Our results provide quantitative hyperparameter bounds and a stability--locality perspective for analyzing and designing SAM variants.