深度学习进阶文档深度页
梯度裁剪
在 backward 后按总范数或元素值限制梯度,缓解梯度爆炸。
学习目标
- 正确放置裁剪
- 记录总范数
- 区分根因修复
为什么需要这个概念
RNN、长序列和不稳定训练可能产生巨大梯度。
- 防止单步破坏参数
- AMP/RNN 常用护栏
学习前需要掌握
梯度爆炸内容待补充
背景与问题
RNN、长序列和不稳定训练可能产生巨大梯度。
概念定义
范数裁剪保持方向并统一缩小过大梯度;值裁剪逐元素截断。
直观理解
方向不变,只把过长箭头缩到安全半径。
核心原理
顺序 backward→AMP unscale→clip→step;裁剪不能修复 NaN 损失。
clip_grad_norm_ 原地修改梯度并返回裁剪前范数。
数学表达
核心公式
M 是最大范数。
变量含义
g所有参数梯度M最大范数
计算步骤
- 1缩放系数 1/5
代码对应位置
示例 1梯度裁剪 · PyTorch 示例:验证梯度裁剪的输入、计算和输出。
变量解释
| 变量 | 含义 |
|---|---|
g | 所有参数梯度 |
M | 最大范数 |
完整数值示例
范数裁剪
已知条件
- ||g||=5
- M=1
- 1
缩放系数 1/5
新范数 1
处理前后对比
计算与实现步骤
- 1
zero_grad
- 2
forward/loss
- 3
backward
- 4
unscale AMP
- 5
clip
- 6
step
代码实现
示例 1
梯度裁剪 · PyTorch 示例
梯度裁剪 · PyTorch 示例
PyTorch
example_01.py用途:验证梯度裁剪的输入、计算和输出。
loss.backward()
norm=torch.nn.utils.clip_grad_norm_(model.parameters(),1.0,error_if_nonfinite=True)
optimizer.step()代码解析
解析始终位于完整代码下方,并按实际代码段逐项对应。
输入数据与任务
验证梯度裁剪的输入、计算和输出。
model- 组合预处理和估计器的模型对象。
loss- 当前预测与目标之间的损失值。
Step 1 · 1–3 行
计算损失对参数的梯度,对应数学表达中的 ∇J;梯度只给出局部变化方向。
loss.backward()
norm=torch.nn.utils.clip_grad_norm_(model.parameters(),1.0,error_if_nonfinite=True)
optimizer.step()代码与数学原理
参数更新由当前梯度与学习率共同决定。
平方误差代码对应均方损失。
预期输出或运行结果
norm 可大于 1,但实际梯度已限制到 1。
常见错误 · 3 条
- backward 前裁剪
- 掩盖学习率过大
- NaN 后才裁剪
实际应用
- RNN
- Transformer
- AMP
常见错误
backward 前裁剪
掩盖学习率过大
NaN 后才裁剪
文档来源
13_梯度问题与权重初始化原始文档
学习资料/deeplearning/13_梯度问题与权重初始化.md- §3 梯度爆炸
- §4 监控
- §5 梯度裁剪
常见问题
裁剪解决梯度消失吗?
不能。
AMP 时先做什么?
先 scaler.unscale_。
推荐学习资料
暂未收录相关资料
参考库不会生成虚假资源或无效外部链接。