SAM 的临时扰动该不该留下:用一个二次损失检查两次梯度
把SAM接进训练代码时,一个容易藏住的错误是:先把权重移到扰动位置,再直接在那里执行优化器更新。两次反向传播都成功,损失也可能下降,结果却偏离算法。拿一个能手算的目标检查“在哪里求梯度、从哪里出发”,比只看训练曲线更容易发现问题。
AI生成的概念示意图:虚线表示临时探测,实线表示实际更新;地形、箭头和距离不对应精确损失曲面。
先固定一批数据对应的目标
教学损失为L(x,y)=(x²+4y²)/2,当前参数w=(3,1),初始损失6.5。它在两个坐标上的弯曲程度不同,梯度为g=(x,4y),所以当前位置得到(3,4),长度为5。设扰动半径ρ=0.5、学习率η=0.1,不加权重衰减。
SAM希望降低参数邻域内的高损失。采用二范数与常用一阶近似时,先取ε=ρg/||g||,得到(0.3,0.4)。这是沿损失上升方向寻找探测点,不能误用梯度下降的负号。也不能把两个坐标各加0.5,否则扰动总长度会超过设定半径。
第二次求导完成后回到原来的起点
探测点为w+ε=(3.3,1.4),损失9.365。这里重新求梯度,得到g′=(3.3,5.6)。常用SAM近似随后用这个梯度更新原参数:w新=w-ηg′=(2.67,0.44),代回原损失得到3.95165。
若错误地从探测点更新,会得到(2.97,0.84),损失5.82165。它仍低于初始6.5,因此“损失下降了”抓不住错误。验收应直接断言最终参数等于原始快照减去第二次梯度,而不是只比较损失有没有改善。
普通梯度下降则得到(2.7,0.6),损失4.365。本例SAM一步更低,只是这个目标和步长下的计算结果。更复杂的训练里,两种方法的原始训练损失、邻域损失和验证集表现应分开比较,不能据此预测泛化收益。
一阶探测没有精确求出最坏位置
把半径0.5的圆周写成ε=(0.5cosθ,0.5sinθ),可以另做一维数值最大化。本例最坏扰动约为(0.24610,0.43524),损失约9.38842,稍高于刚才的9.365。原因是选择扰动时用了当前位置的一阶信息,没有完全追踪邻域内的曲率。
因此报告应写“使用归一化梯度近似内层最大化”。常见实现还忽略扰动随参数变化所带来的二阶导数项。如果让自动微分直接穿过扰动生成过程,就需要明确这是另一种梯度计算,不能声称和上述两遍计算完全一致。
给AI的检查任务要带状态快照
可以要求AI按顺序列出原参数、第一遍梯度、扰动范数、探测参数、第二遍梯度、恢复后的参数以及最终值。对本例,把梯度也用中心有限差分单独复核;再令ρ=0,确认结果退回相同基础优化器的一次普通更新。
两遍计算使用同一批样本,并说明随机增强与随机层怎样处理
临时扰动只改变探测位置,不应被误计作一次优化器历史更新
梯度范数为零时明确处理规则,避免除零
调试真实网络还应检查两遍前向传播是否重复改变批归一化统计等状态。先在没有随机层的小目标上通过精确参数断言,再扩大到网络层和训练批次。若第二遍计算失败,也应恢复原始参数后再退出,避免错误处理把半完成的探测权重留给下一个批次。这个恢复分支同样可以用小例触发并验收。这个练习交付的是一套定位实现偏差的方法,二维地形本身不代表真实模型的全部优化困难。
资料核对日期:2026年10月3日。本文数值为原创教学设定,已用本地程序复算,未进行真实网络训练或生产效果测试。
参考资料
Foret等:Sharpness-Aware Minimization,第2节与算法1


