梯度裁剪:爆炸刹车
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()。对缩放后的梯度做裁剪没有意义——范数被放大或缩小,阈值失效。
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=1 且 model.train() 时统计量噪声也会诱发 spike,裁剪挡不住这种 NaN——应换组归一化或增大 batch。若 grad_norm 长期大于十仍被裁到一,有效更新方向被压扁——应降学习率而非只裁剪。
5. 何时需要裁剪
循环神经网络与 Transformer 训练初期、混合精度下损失缩放偏小导致下溢再 spike、强化学习策略梯度(方差本来就大)——这些场景默认开裁剪。卷积分类小模型有时不需要;加了也不 harm,但别把它当万能药。
6. 案例:裁剪掩盖学习率过大
某项目加裁剪后训练「稳定」,但收敛极慢——查日志发现每步梯度范数都被裁到阈值,有效步长恒定偏小。降学习率后裁剪几乎不再触发,收敛反而更快。裁剪是安全阀,不是替代学习率搜索。
7. 验收
- 混合精度路径 unscale → clip → step 顺序正确。
- 日志含裁剪前
grad_norm;spike 时能定位层。 - 裁剪后仍 NaN → 查学习率、缩放、数据,不继续加大阈值。
- 累积时在最后一步裁剪。
- 长期梯度范数被压扁 → 降学习率。
8. 落地与衔接
在训练框架里把裁剪前的梯度范数作为一级指标落库,与损失、学习率同频写入可视化面板。混合精度路径在配置里固化先反缩放再裁剪再更新参数的顺序,并在持续集成里用单步反向传播测试顺序是否正确。出现非数报警时自动保存当前批次与优化器状态,方便复现尖峰。推理路径不需要裁剪,但训练日志应保留梯度范数曲线供回归对比。新成员接手时应能从配置直接看出阈值来源,而不是口头约定。
梯度裁剪是训练稳定性的安全阀。方向保留的全局范数裁剪是默认选择;根因仍在学习率、缩放和数据。裁剪阈值应随模型规模与任务记录在实验元数据里,避免后人只开裁剪却不查学习率;长期被压扁的梯度范数是降低学习率的信号,不是继续加大阈值的理由。发版前对照基线运行,确认裁剪触发频率在合理区间。
相关
也可以看看
- ·3 分钟阅读
grad_norm spike 与 misclass grid
只看 train loss 不够——要 log lr、grad norm、val slice、GPU 利用率和样例可视化,异常通常是 loss 曲线先以外的信号。
- ·5 分钟阅读
beta=0.999 与 eval 权重
训练时对参数做指数滑动平均,验证和部署常用 EMA 权重——曲线更顺,有时泛化更好,但要与 optimizer 步和 checkpoint 策略一致。
- ·5 分钟阅读
use_reentrant=False 粒度
Gradient checkpointing 用重算 forward 换显存——训练大模型或高分辨率输入时 OOM 的常用解法,wall-clock 涨 20–40% 要算进预算。
johan's blog