返回专辑
·Johan·3 分钟阅读

梯度裁剪:爆炸刹车

clip_grad_norm 能止住 NaN 连锁,但根因常是 LR 过大、loss scale 或 bad batch——裁剪是买时间,不是根治。

梯度裁剪:爆炸刹车

1. 裁剪能跑完,降学习率才真正稳定

视觉 Transformer 训到第八百步 loss 变 NaN,加 clip_grad_norm_(..., max_norm=1.0) 后能跑完 epoch——但同一 batch 上把峰值学习率从 3e-4 降到 1e-4 才真正稳定。裁剪只是让更新步长有上限,不替学习率调参。根因常是学习率过大、损失缩放异常或脏 batch;裁剪是买时间排查,不是根治。

2. 机制:全局范数与逐元素截断

PyTorch clip_grad_norm_ 算所有参数梯度的全局 L2 范数 ;若超过 max_norm,整体缩放 g \leftarrow g \cdot \frac{\text{max_norm}}{|g|_2+\epsilon}方向不变。另有 clip_grad_value_ 逐元素截断,会改方向,Transformer 里我很少用。

后归一化 Transformer 默认 max_norm=1.0;前归一化有时 0.5 更稳。循环神经网络长反向传播 0.5 也常见。阈值太大等于没裁;有裁剪也不代表可以把学习率随便加大。

3. 混合精度下顺序不能错

混合精度路径顺序必须是:scaler.unscale_(optimizer) → 裁剪 → scaler.step()scaler.update()。对缩放后的梯度做裁剪没有意义——范数被放大或缩小,阈值失效。

python
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()

梯度累积时裁剪在最后一个累积步做,不是每个微 batch。NaN 前五十步常能看到 grad_norm 缓慢爬升——应在此阶段介入,而不是等 NaN 后再加裁剪。

4. 根因排查

NaN 后依次查:学习率与 warmup、GradScaler 缩放、标签是否含 inf、自定义 loss 是否除零、优化器状态是否损坏(NaN 后常要重载 checkpoint)。反向 hook 按层记 L2 范数,谁先爆先查该块的初始化、学习率和数据。

批归一化在 batch_size=1model.train() 时统计量噪声也会诱发 spike,裁剪挡不住这种 NaN——应换组归一化或增大 batch。若 grad_norm 长期大于十仍被裁到一,有效更新方向被压扁——应降学习率而非只裁剪。

5. 何时需要裁剪

循环神经网络与 Transformer 训练初期、混合精度下损失缩放偏小导致下溢再 spike、强化学习策略梯度(方差本来就大)——这些场景默认开裁剪。卷积分类小模型有时不需要;加了也不 harm,但别把它当万能药。

6. 案例:裁剪掩盖学习率过大

某项目加裁剪后训练「稳定」,但收敛极慢——查日志发现每步梯度范数都被裁到阈值,有效步长恒定偏小。降学习率后裁剪几乎不再触发,收敛反而更快。裁剪是安全阀,不是替代学习率搜索。

7. 验收

  • 混合精度路径 unscale → clip → step 顺序正确。
  • 日志含裁剪前 grad_norm;spike 时能定位层。
  • 裁剪后仍 NaN → 查学习率、缩放、数据,不继续加大阈值。
  • 累积时在最后一步裁剪。
  • 长期梯度范数被压扁 → 降学习率。

8. 落地与衔接

在训练框架里把裁剪前的梯度范数作为一级指标落库,与损失、学习率同频写入可视化面板。混合精度路径在配置里固化先反缩放再裁剪再更新参数的顺序,并在持续集成里用单步反向传播测试顺序是否正确。出现非数报警时自动保存当前批次与优化器状态,方便复现尖峰。推理路径不需要裁剪,但训练日志应保留梯度范数曲线供回归对比。新成员接手时应能从配置直接看出阈值来源,而不是口头约定。

梯度裁剪是训练稳定性的安全阀。方向保留的全局范数裁剪是默认选择;根因仍在学习率、缩放和数据。裁剪阈值应随模型规模与任务记录在实验元数据里,避免后人只开裁剪却不查学习率;长期被压扁的梯度范数是降低学习率的信号,不是继续加大阈值的理由。发版前对照基线运行,确认裁剪触发频率在合理区间。

← 全部文章

johan's blog