CASAM: Consistency-Anchored Sharpness-Aware Minimization for Improved Model Generalization
Abstract
Many studies argue that the success of sharpness-aware minimization (SAM) is due to the implicit regularization of the gradient norm or the eigenvalues of the Hessian. However, they cannot fully explain why the flatness-based measures do not always correlate well with generalization. For instance, we can easily construct a counterexample by only optimizing with hard examples. In this paper, we propose to resolve such a contradiction from the perspective of anti-overfitting. First, we show that the sharpness at the mini-batch level approximately equals the sum of the gradient norm and the Kullback–Leibler (KL) divergence between the predictive distributions evaluated at the current point and its adversary. Second, we show that the \emph{gradient consistency}, which approximately quantifies the inner product between the mini-batch gradient and the per-example gradient, decreases consistently during training. This result explains why SAM is less likely to suffer from overfitting. At last, building on these insights, we further introduce CASAM, an algorithm that reweights each example according to gradient consistency to enhance the generalization performance. Extensive experiments on popular benchmarks such as ImageNet-1K and Clothing1M corroborate its efficacy.