梯度截断阈值

wen IT资讯 27

本文目录导读:

梯度截断阈值

  1. 核心原理
  2. 这张图很好地解释了过程
  3. 为什么需要这个参数?
  4. 梯度截断类型
  5. 常见设定值
  6. 代码示例(PyTorch)

梯度截断阈值(Gradient Clipping Threshold)是训练深度神经网络(尤其是RNN、LSTM或Transformer)时用于防止梯度爆炸(Gradient Explosion)的一个超参数。

它设定了一个“安全上限”,当梯度的范数(通常是L2范数)超过这个阈值时,就直接把梯度缩放到这个阈值的大小。

核心原理

  1. 计算梯度的全局范数:在一次反向传播后,计算所有参数梯度的L2范数(或L2范数的平方,取决于实现),公式大致为: ( \text{totalnorm} = \sqrt{\sum{i} |g_i|^2} ) (( g_i ) 是第i个参数的梯度向量)

  2. 比较与缩放

    • total_norm <= threshold:梯度保持不变。
    • total_norm > threshold:将所有梯度乘以一个缩放因子 scale = threshold / total_norm,这样,所有梯度的总范数被强行拉回到阈值大小,但梯度的方向保持不变。

这张图很好地解释了过程

  • 灰色轨道是梯度原本爆炸的方向(幅度很大)。
  • 红色圆圈是设定的“安全区域”(阈值)。
  • 执行梯度截断后,无论原始梯度多大,都会被拉回圆圈的边界上,保证更新步长不会过大。

为什么需要这个参数?

  • 防止训练崩溃:在训练早期或处理长序列时,梯度可能突然变得非常大(例如指数级增长),如果不加以限制,参数更新会剧烈震荡,导致损失变成NaN(非数值),模型彻底失效,梯度截断是防止这种情况最直接的方法。
  • 稳定训练:它保证每一步的参数更新都在一个可控的、较“温和”的范围内,使得训练过程更加平滑,尤其是对于那些容易发生梯度爆炸的模型(如RNN的BPTT)。
  • 允许使用更大的学习率:因为梯度被限制了上限,所以我们通常可以设置比不使用截断时稍大一点的学习率,而不用担心瞬间爆炸。

梯度截断类型

虽然最常用的是“按范数截断”(Norm Clipping),但还有另一种:

  1. 按范数截断(Norm Clipping):(最常用,如PyTorch的 torch.nn.utils.clip_grad_norm_
    • 对整个网络所有梯度的总范数进行截断。
    • 优点:保持不同层之间梯度的相对比例。
  2. 按值截断(Value Clipping):(例如PyTorch的 torch.nn.utils.clip_grad_value_
    • 对每个单独的梯度元素设定一个上下限([-1.0, 1.0])。
    • 优点:实现简单,但会完全破坏梯度的相对大小,不常用,除非梯度值极其不稳定。

常见设定值

这是一个需要调参的超参数,但有一些常见起点:

  • 常用范围:通常在 1 到 10.0 之间。
  • RNN/LSTM:建议从 0 或 5.0 开始尝试,如果训练经常崩溃,可以降低(如0.5)。
  • Transformer:现代Transformer(如BERT、GPT)通常使用 0clip by global norm = 1.0(OpenAI的常见做法),有时也会用到0.5或0.25。
  • GAN(生成对抗网络):常设为 0

经验法则:

  • 如果损失在训练中突然变成NaN,应降低梯度截断阈值(例如从1.0降到0.5)。
  • 如果训练非常缓慢,可以尝试调高阈值(例如从1.0升到5.0),或者干脆不截断(设为无穷大)。

代码示例(PyTorch)

这是最常见的使用方式(训练一个步):

import torch
import torch.nn as nn
from torch.nn.utils import clip_grad_norm_
...
model = YourModel()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
for epoch in range(num_epochs):
    for inputs, targets in dataloader:
        optimizer.zero_grad()
        outputs = model(inputs)
        loss = loss_fn(outputs, targets)
        loss.backward()
        # ----- 关键步骤:梯度截断 -----
        # 对模型的所有参数进行按L2范数截断,阈值为1.0
        clip_grad_norm_(model.parameters(), max_norm=1.0, norm_type=2)
        # ----------------------------
        optimizer.step()
方面 说明
现象 梯度爆炸
作用 防止更新步长过大导致震荡或NaN
常用值 0 (推荐初值)
过大后果 失去截断效果,仍然可能爆炸
过小后果 梯度被过于频繁地大幅缩放,抑制模型学习能力

梯度截断阈值 = 1.0 是一个被广泛验证的、很好的默认起点,如果你的模型老是爆掉,把它调小;如果你的模型训练得太慢且稳定,可以适当调大。

抱歉,评论功能暂时关闭!