循环神经网络(RNN)#
学习目标与记号#
准确解释隐藏状态、参数共享与 BPTT;
根据公式和张量维度分析梯度消失,并识别常见实现错误;
用
Python实现或验证梯度爆炸;在其他条件相同的情况下,只改变一个需要研究的因素进行比较,并根据实验结果分析该因素可能带来的影响。
本节沿用统一记号:普通小写字母表示标量,粗体小写字母表示向量,粗体大写字母表示矩阵或高阶张量;层编号写作上标 \([l]\),样本或时间编号写作下标;转置写作 \(\trans\)。数学公式中的单个向量默认都是列向量;需要表示行向量时,必须显式写出转置符号,例如 \(\bx_t\trans\)。批量程序仍采用常见的 \(m\times T\times d_x\) 张量布局;固定时间 \(t\) 后,批量输入矩阵写为 \(\bX_t=[\bx_{1,t}\trans;\ldots;\bx_{m,t}\trans]\in\mathbb{R}^{m\times d_x}\)。正文与练习中的程序都应同时检查数值结果和数组维度。本节练习的参考答案见 循环神经网络与随时间后向传播答案。
RNN 常用于 自回归模型 与 语言模型;若普通 RNN 难以保留长期信息,可继续阅读 LSTM 和 GRU。
循环神经网络用固定维度的隐藏状态递归汇总历史信息。它不要求数据真的满足有限阶马尔可夫性质;更准确地说,模型假设“与任务相关的历史”可以被列向量 \(\bh_t\in\mathbb{R}^{d_h}\) 近似压缩,并用同一更新函数跨时间复用。
Vanilla RNN#
对输入列向量 \(\bx_t\in\mathbb{R}^{d_x}\),最基本的 RNN 单元为
其中 \(\ba_t,\bh_t\in\mathbb{R}^{d_h}\) 和 \(\bz_t,\widehat{\by}_t\in\mathbb{R}^{d_y}\) 均为列向量,并且
这些维度保证 \(\bW_{xh}\cdot\bx_t\)、\(\bW_{hh}\cdot\bh_{t-1}\) 和 \(\bb_h\) 都是 \(d_h\) 维列向量,可以直接相加。输出函数 \(g\) 由任务决定:回归可用恒等映射,多分类训练通常把线性运算结果 \(\bz_t\) 直接交给交叉熵函数。初始状态 \(\bh_0\) 可以设为列向量 \(\bzero\in\mathbb{R}^{d_h}\),也可以学习或由另一个编码器给出。
图 24 Vanilla RNN 在时间上的展开#
展开图中每个时间步看似有一套单元,但 \(\bW_{xh}\)、\(\bW_{hh}\) 和偏置在所有时间步共享。因此参数量不随序列长度增长,计算深度却随长度增长。隐藏状态是有损摘要,不能声称它完整保存了全部历史。
一个可运行的 RNN 单元#
下面用 PyTorch 手写一个最简单的循环神经网络,用于展示单个时间步的隐藏状态更新如何扩展到整个序列。SimpleRNN 的主要输入 x 形状为 (batch_size, seq_len, input_size),可选的初始隐藏状态 h 形状为 (batch_size, hidden_size);若不提供 h,程序会用全 0 状态开始。返回值 logits 包含每个时间步的输出,形状为 (batch_size, seq_len, output_size);第二个返回值是最后一个隐藏状态。这段代码只定义模型,不会自动训练或打印结果。
1 import torch
2 import torch.nn as nn
3
4 class SimpleRNNCell(nn.Module):
5 def __init__(self, input_size, hidden_size):
6 super().__init__()
7 self.hidden_size = hidden_size
8 self.i2h = nn.Linear(input_size + hidden_size, hidden_size)
9
10 def forward(self, x_t, h_prev=None):
11 # x_t: (batch_size, input_size),第 i 行表示 x_{i,t}^T
12 if h_prev is None:
13 h_prev = x_t.new_zeros(x_t.size(0), self.hidden_size)
14 combined = torch.cat((x_t, h_prev), dim=-1)
15 return torch.tanh(self.i2h(combined))
16
17 class SimpleRNN(nn.Module):
18 def __init__(self, input_size, hidden_size, output_size):
19 super().__init__()
20 self.cell = SimpleRNNCell(input_size, hidden_size)
21 self.readout = nn.Linear(hidden_size, output_size)
22
23 def forward(self, x, h=None):
24 # x: (batch_size, seq_len, input_size)
25 states = []
26 for t in range(x.size(1)):
27 h = self.cell(x[:, t, :], h)
28 states.append(h)
29 states = torch.stack(states, dim=1)
30 logits = self.readout(states)
31 return logits, h
SimpleRNNCell 先把当前输入与上一时刻的隐藏状态沿最后一轴拼接,再通过线性变换和 tanh 得到新状态。SimpleRNN 按时间顺序重复这一步,将所有状态堆叠后交给 readout,因而 logits[:, t, :] 就是第 \(t\) 个时间步的未归一化预测,h 则可用作下一段序列的初始状态。
这个实现显式展示循环,且正确继承输入的设备与数据类型,但它假定序列长度大于 0,也没有在内部处理变长序列的掩码。实际训练应优先使用 nn.RNN、nn.GRU 或 nn.LSTM,它们能调用更高效的融合算子,并处理多层、双向和 Dropout 等选项。
备注
数学公式把单个输入 \(\bx_t\) 和隐藏状态 \(\bh_t\) 写成列向量。固定时刻 \(t\) 后,程序为了同时计算 \(m\) 个样本,把它们按行组成 \(\bX_t=[\bx_{1,t}\trans;\ldots;\bx_{m,t}\trans]\) 和 \(\bH_{t-1}=[\bh_{1,t-1}\trans;\ldots;\bh_{m,t-1}\trans]\);因此 x_t[i] 对应 \(\bx_{i,t}\trans\),h_prev[i] 对应 \(\bh_{i,t-1}\trans\)。nn.Linear 对每一行完成的运算与数学公式中的矩阵乘列向量等价。
输入—输出组织方式#
序列到单值(many-to-one):使用最后一个有效状态或所有状态的池化结果完成分类。批量中有填充时,不能直接取统一数组的最后一列。
等长序列到序列(many-to-many):每个时间步都有监督信号,例如序列标注。损失应屏蔽填充位置。
编码器—解码器:编码器读取输入,解码器自回归地产生长度可能不同的输出。只用最后状态作为固定瓶颈时,长序列性能可能受限,注意力机制可直接访问全部编码器状态。
堆叠 RNN:第 \(l\) 层在时间 \(t\) 的输出成为第 \(l+1\) 层同一时间的输入。层下标与时间下标必须区分。
双向 RNN:分别从左到右和从右到左编码,再拼接两个方向的状态。它适合能够看到完整输入的标注或编码任务,不适合要求实时、因果预测的场景,因为从右到左的状态使用了未来输入。
图 25 双向 RNN 同时利用两个方向的上下文#
备注
变长序列批处理通常需要长度、填充掩码,或框架提供的 packed sequence。分类时应提取每个样本最后一个 有效 时间步,而不是填充后的最后位置。
随时间后向传播(BPTT)#
设总损失为 \(\mathcal{J}=\sum_{t=1}^{T}\ell_t\)。本节把标量损失关于向量的梯度写成与该向量维度相同的列向量,并记
从最后一个时间步开始进行后向传播,令列向量 \(\bdelta_t=\nabla_{\ba_t}\mathcal{J}\in\mathbb{R}^{d_h}\),则
第一项是当前输出损失的贡献,第二项是未来时间步沿循环边传回的贡献。共享参数的梯度必须对所有时间步累加:
上式中的 \(\bg_t\cdot\bh_t\trans\)、\(\bdelta_t\cdot\bh_{t-1}\trans\) 和 \(\bdelta_t\cdot\bx_t\trans\) 都是“列向量乘行向量”的外积,所得矩阵分别与 \(\bW_{hy}\)、\(\bW_{hh}\) 和 \(\bW_{xh}\) 的维度一致。
这就是随时间后向传播:把循环在计算图中展开,再按照普通后向传播的计算顺序,从 \(T\) 逐步计算到 1。
梯度消失与爆炸#
在后向传播过程中,若要计算更早时间步的梯度,就需要逐步计算相邻两个时间步之间隐藏状态的变化。对于使用 tanh 激活函数的 RNN,上一时刻的隐藏状态 \(\bh_{t-1}\) 发生微小变化时,当前隐藏状态 \(\bh_t\) 会如何变化,可以用下面的雅可比矩阵表示:
跨越 \(k\) 步时会出现 \(\bJ_t\cdot\bJ_{t-1}\cdots\bJ_{t-k+1}\)。其范数可能随连乘快速缩小或增大,具体取决于权重、状态、激活饱和与矩阵方向,不能仅凭 \(\lVert\bW_{hh}\rVert\) 大于或小于 1 就作充要判断。
梯度裁剪(gradient clipping)可限制爆炸梯度,例如把整体梯度范数裁到阈值以内;它不能恢复已经消失的长期梯度。合理初始化、归一化、残差/跳跃连接以及 LSTM、GRU 的加性状态路径可以改善传播,但都不保证无限长依赖。截断 BPTT 只在固定窗口内进行后向传播,降低内存和计算代价,也主动舍弃了更远的梯度。
Shiny 交互演示:RNN 的时间展开
交互页面会按照正文的 Vanilla RNN 公式显式计算每个时间步的隐藏状态,并展示 \(\bW_{xh}\)、\(\bW_{hh}\) 和 \(\bb_h\) 在所有时间步共享的事实。对于较短序列,补齐位置不会再次更新隐藏状态。
核心推导与实现核验#
核心关系
其中,\(\phi\) 表示激活函数,本节正文以 \(\tanh\) 为例;\(\bW_h,\bW_x,\bb\) 是每个时刻复用的权重和偏置。
推导路径。 把循环沿时间展开后,每个时刻复用同一参数;BPTT 从后向前累积来自当前输出和未来状态两条路径的梯度。
关键条件
跨 \(k\) 步的梯度包含雅可比矩阵乘积;若其谱范数长期小于 1 则范数上界指数衰减,长期大于 1 则可能指数放大。
数据规模
若每个时刻的列向量 \(\bx_t\in\mathbb{R}^{d_x}\) 含 \(d_x\) 个数,列向量 \(\bh_t\in\mathbb{R}^{d_h}\) 含 \(d_h\) 个数,则输入权重 \(\bW_x\) 的大小为 \(d_h\times d_x\),循环权重 \(\bW_h\) 的大小为 \(d_h\times d_h\),计算结果仍是含 \(d_h\) 个数的列向量。
常见误区
RNN 在所有时间步都应使用同一组参数。如果不同时间步使用不同参数,模型就不再具有参数共享的特点。当相邻批次包含彼此无关的序列时,每个批次开始前还应重新设置隐藏状态,否则模型会把前一个批次的信息错误地传入下一个批次。
动手检查
对长度 2 或 3 的序列把手写 BPTT 与 自动微分 对照;再改变序列长度,绘制初始时刻梯度范数以验证长期传播趋势。
数值稳定性与规模
训练长序列时,应检查所有参数的梯度合在一起后是否过大。如果梯度的整体大小超过预设上限,可以在确认后向传播公式和程序实现正确后,按相同比例缩小所有参数的梯度,这称为 全局范数裁剪。采用截断 BPTT 时,可以把长序列分成若干较短片段:后一片段继续使用前一片段最后的隐藏状态,但应切断两段之间的梯度联系,使梯度不再传回更早的片段。在 PyTorch 中可以使用 h = h.detach() 完成这一步。
本节小结#
RNN 在时间上共享同一状态更新参数。
BPTT 的核心是跨时刻雅可比连乘。
状态边界、截断长度和 梯度裁剪 必须显式管理。
综合练习#
程序题应固定随机种子、写出维度断言并报告运行环境。全部参考答案见 循环神经网络与随时间后向传播答案。
前向传播计算。 考虑标量 RNN:\(h_t=\tanh(x_t+0.5h_{t-1})\),且 \(h_0=0\)。对序列 \((x_1,x_2,x_3)=(1,0,-1)\),依次计算三个线性运算结果和隐藏状态;若输出为 \(z=2h_3\),继续计算 \(z\)。
BPTT 推导。 从正文的 RNN 公式和 \(\mathcal{J}=\sum_{t=1}^{T}\ell_t\) 出发,推导 \(\bdelta_t\) 的后向传播公式,以及 \(\partial\mathcal{J}/\partial\bW_{hh}\) 和 \(\partial\mathcal{J}/\partial\bW_{xh}\)。解释共享参数的梯度为何必须对时间求和。
梯度范数证明。 设从时刻 \(s\) 到 \(t\) 的状态梯度包含 \(\bJ_t\cdot\bJ_{t-1}\cdots\bJ_{s+1}\)。利用矩阵范数的次乘性证明其范数上界,并分别给出该上界指数衰减和可能放大的充分条件;说明这些条件为何不是梯度消失或爆炸的充要条件。
参数量与计算量。 对 \(d_x=3\)、\(d_h=4\)、\(d_y=2\) 的单层 RNN,计算 \(\bW_{xh}\)、\(\bW_{hh}\)、\(\bb_h\)、\(\bW_{hy}\) 和 \(\bb_y\) 的维度及参数总数。再写出长度为 \(T\) 的单样本前向传播主要乘加运算量的数量级。
变长序列前向实现。 从零实现批量 RNN 前向传播,输入张量的维度为 \(m\times T\times d_x\),并接收每个样本的真实长度。数学公式中的 \(\bx_{i,t}\) 是列向量;程序固定时间 \(t\) 后得到 \(\bX_t=[\bx_{1,t}\trans;\ldots;\bx_{m,t}\trans]\in\mathbb{R}^{m\times d_x}\)。程序应返回全部隐藏状态和每个样本最后一个有效状态;在填充位置保持旧状态不变,并用不同长度的小批次验证结果。
手写 BPTT 核验。 为长度为 2 或 3 的小序列实现手写 BPTT,返回每个 \(\bdelta_t\) 和所有参数梯度。把结果分别与 自动微分 和 中心差分 比较,给出相对误差标准,并加入会暴露“只保留最后一步参数梯度”错误的测试。
状态与截断测试。 为独立序列之间的隐藏状态重置、连续片段之间的状态传递和截断 BPTT 的计算图分离编写测试。验证批量结果与逐条序列结果一致,并检查改变 padding 值不会改变有效位置输出与损失。
逻辑回归与 RNN。 构造一个类别取决于词元顺序的二分类任务,比较基于词袋特征的逻辑回归与单层 RNN。固定训练/验证/测试划分、随机种子集合、输入词表、训练样本和各自的训练更新预算;报告测试准确率或 F1、训练时间、固定批量预测时间、峰值内存和模型参数量。
完整与截断 BPTT。 在同一长序列任务上比较完整 BPTT 与截断长度 \(K\in\{16,32,64\}\)。固定数据划分、随机种子集合、RNN 参数初值、优化器、总处理词元数和参数更新次数;报告预测质量、训练时间、固定批量预测时间、峰值内存、循环单元调用次数和梯度范数,并讨论截断造成的长期信用分配偏差。
隐藏维度与学习率。 比较隐藏维度 \(d_h\in\{16,32,64\}\) 和若干学习率的组合。固定数据划分、随机种子集合、最大训练更新次数、早停规则和评价代码;报告每个组合的预测质量、训练时间、固定批量预测时间、峰值内存与参数量,并区分“容量变化”和“优化设置变化”对结果的影响。