RNN的反向传播算法为BPTT,全称为Back Propagation Through Time,即沿时间轴的反向传播。
RNN的正向计算链路包括两个方向:
和前馈网络一样,从输入层到输出层,即:
的计算链路。
沿着时间步从左到右,即:
的计算链路。
在利用RNN进行文本分类(图a)的例子中,RNN最后一个时间步隐藏层的输出会作为整个序列的表征 (即特征),
在经过输出层之后 (Linear + Softmax),最后输出:
其中Y为表示预测输出的随机变量,概率的预测公式为:
再计算交叉熵损失:
y为标注的真值。
我们来看一下这个网络反向传播的过程:
(1) 计算的梯度:
同时计算的梯度:
(2) 计算通过传播过来的
和
的梯度:
同时计算通过传播过来的
和
的梯度:
(3) 计算通过传播过来的
和
的梯度,并与之前计算的梯度累加:
同时计算通过传播过来的
和
的梯度,并进行梯度累加:
(4) 计算通过传播过来的
和
的梯度,并与之前计算的梯度累加:
同时计算通过传播过来的
和
的梯度,并进行梯度累加:
(5) 计算通过传播过来的
和
的梯度,并与之前计算的梯度累加:
同时计算通过传播过来的
梯度:
当我们解决的问题为序列标注时 (图b),每个时间步的输出都会用于计算损失函数,因此梯度会有更多来源,在此就不进行详细推导了。
CrossEntropyLoss — PyTorch 2.14 documentation
在基于交叉熵的损失计算中,可通过标签屏蔽(label masking)机制控制样本对损失的贡献。具体做法是:将待排除样本的目标标签设置为哨兵值
-100(即 PyTorchCrossEntropyLoss默认的ignore_index),损失函数在反向传播时自动跳过这些位置——其损失项不参与汇总,梯度贡献为零,等价于该样本未参与本批次(batch)的参数更新。(严格来说,PyTorch 内部仍会计算这些位置的损失,只是在汇总阶段将其置零并屏蔽梯度)
对模型前向计算无影响:模型仍会照常对该位置产生预测输出,只是该输出不参与目标值与预测值的比对;
对训练过程的影响:被屏蔽样本不会影响模型参数的学习,也不干扰同批次中其他有效样本的梯度;
典型应用场景:序列标注/文本生成中的 padding 位置屏蔽、负采样样本的排除、以及样本级降权(如标注噪声数据)。
对应实现:
import torch.nn as nn # 设置 ignore_index=-100,标签为 -100 的样本自动被忽略 loss_fn = nn.CrossEntropyLoss(ignore_index=-100)
RNN的缺陷
RNN存在沿时间步传播的梯度,因此,序列长度在某种程度上可以看作是网络的“深度”。梯度传播的步数(即偏导相乘的次数)越大,梯度消失或梯度爆炸出现的概率越大。在RNN反向传播的过程中,如果发现梯度的幅度超过一定的阈值,则需要进行截断 (最大取该阈值),防止梯度爆炸的发生。此外,RNNCell的激活函数一般采用Tanh或ReLU,也是为了缓解梯度消失的程度。
同时,随着时间步的增加,序列中靠前位置的输入特征,它们的影响也会随着时间步逐渐衰减,造成信息的遗忘,影响模型序列建模的效果。