一句话定义
锐度感知最小化(Sharpness-Aware Minimization,SAM)是一种训练神经网络的方法:它不只要求"当前这个参数点上损失低",还要求"这个点附近一小圈内的损失都低"。换句话说,它想找的是一片平坦的谷底,而不是一根陡峭的针尖。
它是怎么做的
想象你在山里挑营地。普通训练(比如随机梯度下降)只关心脚下这块地是不是最低点:坡度指向哪,就往哪走一步。问题是,它可能停在一个很窄的裂缝底部——那里确实是个局部最低点,但两边都是陡壁。
SAM 多了一步"环顾四周":先找出附近哪个方向最糟(让损失最大),把参数朝那个方向推一小步,然后在这个"最糟的点"上算梯度,再用它来更新原来的参数。写成公式就是 min_w max_{‖ε‖≤ρ} L(w+ε):内层在半径 ρ 的小邻域里挑出最坏情况,外层再把最坏情况下的损失压下去。实际实现通常用两阶段近似,所以每一步的代价比普通训练高一些。
这个比方也能解释为什么平坦解更抗分布偏移。上线后的数据和训练数据不会完全一样,相当于地形变了或者水位涨了:裂缝底部的水位抬一点就被淹,宽阔的盆地抬一点还是干的。参数轻微漂移、样本轻微偏移,平坦谷底里的模型性能掉得慢。
和相邻概念的区别
| 方法 | 扰动什么 | 主要目标 |
|---|---|---|
| 普通 SGD / Adam | 不扰动,只看当前点 | 尽快压低训练损失 |
| SAM | 扰动模型参数,找邻域内的最坏点 | 让解落在平坦区域,泛化更稳 |
| 对抗训练 | 扰动输入样本 | 让模型对输入扰动不敏感 |
| L2 正则 / 权重衰减 | 惩罚参数的大小 | 抑制过拟合,间接且与数据无关 |
对抗训练和 SAM 形式上都是 min-max,只是"折腾"的对象不同:一个折腾数据,一个折腾参数。
对从业者和普通人的意义
对做模型的人来说,SAM 常被用在分布偏移、噪声标签、域泛化这类场景,因为它给模型留了"冗余度":参数抖一抖、数据偏一偏,性能不至于断崖。代价是训练更慢、更占显存,邻域半径 ρ 也要调——调小了几乎没效果,调大了可能学不动。各家框架的接口和默认行为不一样,具体以官方页面为准。
对不太碰模型的职场人,可以记住一句话:模型在测试集上分数高,不代表它上线后稳。它的高分是长在针尖上还是长在平地上,是两回事。SAM 就是在训练阶段主动去找那块平地。
