梯度消失(Vanishing Gradient)

深层网络训练时梯度在反向传播过程中逐层衰减趋近于零的现象。导致底层参数几乎无法更新,网络"学不动"。残差连接和归一化是主要解决方案。

梯度消失(Vanishing Gradient)

直觉理解

想象一条传话链——100 个人依次传话,到最后一个人时原始信息已经面目全非。梯度消失就是类似的问题:反向传播的梯度信号经过多层传递后衰减到接近零,导致底层网络"学不动"。

为什么会消失

反向传播时,梯度通过链式法则逐层相乘:

LW1=Lhnhnhn1h2h1h1W1\frac{\partial L}{\partial W_1} = \frac{\partial L}{\partial h_n} \cdot \frac{\partial h_n}{\partial h_{n-1}} \cdot \cdots \cdot \frac{\partial h_2}{\partial h_1} \cdot \frac{\partial h_1}{\partial W_1}W1L=hnLhn1hnh1h2W1h1

如果每一步的偏导 hi+1hi\frac{\partial h_{i+1}}{\partial h_i}hihi+1 的绝对值小于 1(使用 sigmoid/tanh 激活函数时很常见),连乘 nnn 次后趋近于零:

0.5100103000.5^{100} \approx 10^{-30} \approx 00.510010300

动手试试

下面的动画展示梯度信号从输出层(右)向输入层(左)传播时逐层衰减的过程。拖动衰减率滑块观察变化——打开"残差连接"开关,看梯度如何被保住:

梯度消失 vs 残差连接

在 RNN 中的表现

RNN 沿时间步展开后等价于一个很深的网络。100 个时间步 = 100 层。梯度要从第 100 步传回第 1 步,经过 99 次连乘,早期时间步的梯度几乎为零——模型"记不住"长距离信息。

这就是 RNN 难以处理长文本的根本原因。

解决方案

方案原理代表
LSTM 门控通过遗忘门控制信息保留,提供梯度旁路LSTM (1997)
残差连接output = f(x) + x,梯度可以直接跳过子层ResNet / Transformer
归一化控制每层输出的分布范围LayerNorm / BatchNorm
ReLU 激活正区间导数恒为 1,不衰减几乎所有现代网络

Transformer 如何解决

Transformer 通过三重保障抵御梯度消失:

  1. Self-Attention:任意 token 之间直接连接,梯度路径长度 O(1)
  2. 残差连接:每个子层都有跳线,梯度可以绕过变换直接回传
  3. LayerNorm:稳定每层的数值范围

这就是为什么 Transformer 能堆到几十上百层而不退化。

引用本术语的文章