梯度裁剪解决什么问题?
简化版
梯度裁剪限制一次更新的梯度幅度,常用全局范数裁剪 g←g·min(1,c/||g||);它缓解爆炸和偶发尖峰,却不能修复错误损失、过高学习率或持续数值不稳定。AMP 下必须先 unscale 再裁剪。
详细版
-
按值裁剪逐元素截断,按范数裁剪保持整体方向,后者更常见。
-
全局范数应跨全部参数计算,分布式/分片训练需正确聚合。
-
阈值过小会让大多数步骤都被压缩,训练等价于长期缩小有效学习率。
-
应记录裁剪前范数和触发比例,而不只设置一个阈值。
-
NaN 梯度不是“特别大的有限值”,需要另外定位溢出来源。
完整版教学
一、裁剪把更新限制在可信半径内
RNN 长链或深网中梯度可因 Jacobian 连乘迅速变大。
一个异常 batch 产生的巨大更新足以破坏已学权重,范数裁剪保留方向并压缩长度。
裁剪是最后一道保险,不改变产生异常梯度的前向机制。
若 90% 步骤都触发,应检查学习率、数据异常和损失尺度。
二、底层机制与公式
if ||g||_2 > c:
g_clipped = c * g / ||g||_2
else:
g_clipped = g
三、带数字的推演
梯度向量 [3,4] 的 L2 范数是 5。
阈值 c=2 时乘 2/5,得到 [1.2,1.6],方向不变且新范数正好为 2。
四、方案对比
| 方案/对象 | 核心特点 | 代价或边界 |
|---|---|---|
| 按值裁剪 | 每元素限制到区间 | 会改变梯度方向 |
| 全局范数裁剪 | 整体等比例缩放 | 需正确汇总所有参数 |
| 自适应裁剪 | 相对参数尺度限制 | 实现和解释更复杂 |
五、执行流程
backward -> AMP unscale -> 计算全局 norm -> 记录是否触发
-> clip -> optimizer.step -> 监控 loss、norm 分布和溢出
六、边界条件与工程代价
梯度累积时应对累积完成后的梯度裁剪一次;每个 micro-batch 分别裁剪再求和会得到不同方向。
分布式参数分片下,局部 norm 不是全局 norm。
必须使用框架提供的 sharded clipping,不能每卡独立裁剪自己的碎片。
记忆钩子:裁剪的动作是“方向尽量不动,只把步子缩到半径 c 内”。
七、常见误区与追问
-
误区:裁剪能把 NaN 自动变正常。 NaN 参与范数仍是 NaN,应先定位溢出。
-
追问:范数裁剪为何较少改变方向? 所有梯度元素乘同一缩放系数。
-
误区:阈值越小越稳定。 过小会长期压制有效学习信号。
-
追问:AMP 下顺序是什么? 先用 scaler.unscale_ 恢复真实梯度,再算范数和裁剪。
-
追问:如何判断阈值合理? 看裁剪前范数分位数与触发率,并结合收敛和异常 batch。
八、加强记忆
裁剪的动作是“方向尽量不动,只把步子缩到半径 c 内”。
回答要补上 unscale、累积后裁剪和触发率监控,避免把保险丝当根治方案。