\[ \begin{align}\begin{aligned}\newcommand{\ba}{\boldsymbol{a}} \newcommand{\bb}{\boldsymbol{b}} \newcommand{\be}{\boldsymbol{e}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bw}{\boldsymbol{w}} \newcommand{\bx}{\boldsymbol{x}} \newcommand{\by}{\boldsymbol{y}} \newcommand{\bz}{\boldsymbol{z}} \newcommand{\bd}{\boldsymbol{d}} \newcommand{\bv}{\boldsymbol{v}} \newcommand{\bs}{\boldsymbol{s}}\\\newcommand{\btheta}{\boldsymbol{\theta}} \newcommand{\bbeta}{\boldsymbol{\beta}} \newcommand{\bgamma}{\boldsymbol{\gamma}} \newcommand{\bsigma}{\boldsymbol{\sigma}} \newcommand{\md}{\mbox{d}} \newcommand{\bmu}{\boldsymbol{\mu}} \newcommand{\bone}{\boldsymbol{1}} \newcommand{\bzero}{\boldsymbol{0}} \newcommand{\bepsilon}{\boldsymbol{\epsilon}} \newcommand{\bphi}{\boldsymbol{\phi}} \newcommand{\bh}{\boldsymbol{h}} \newcommand{\bc}{\boldsymbol{c}} \newcommand{\br}{\boldsymbol{r}} \newcommand{\bQ}{\boldsymbol{Q}} \newcommand{\bK}{\boldsymbol{K}} \newcommand{\bV}{\boldsymbol{V}} \newcommand{\bSigma}{\boldsymbol{\Sigma}} \newcommand{\bg}{\boldsymbol{g}} \newcommand{\bxi}{\boldsymbol{\xi}} \newcommand{\bvarepsilon}{\boldsymbol{\varepsilon}} \newcommand{\bdelta}{\boldsymbol{\delta}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bJ}{\boldsymbol{J}} \newcommand{\bp}{\boldsymbol{p}} \newcommand{\bi}{\boldsymbol{i}} \newcommand{\bo}{\boldsymbol{o}} \newcommand{\bE}{\boldsymbol{E}} \newcommand{\bH}{\boldsymbol{H}} \newcommand{\bL}{\boldsymbol{L}} \newcommand{\bu}{\boldsymbol{u}} \newcommand{\bLambda}{\boldsymbol{\Lambda}} \newcommand{\trans}{^{\rm\scriptsize T}} \newcommand{\var}{\mathrm{var}}\\\newcommand{\bA}{\boldsymbol{A}} \newcommand{\bB}{\boldsymbol{B}} \newcommand{\bC}{\boldsymbol{C}} \newcommand{\bD}{\boldsymbol{D}} \newcommand{\bG}{\boldsymbol{G}} \newcommand{\bI}{\boldsymbol{I}} \newcommand{\bM}{\boldsymbol{M}} \newcommand{\bP}{\boldsymbol{P}} \newcommand{\bS}{\boldsymbol{S}} \newcommand{\bU}{\boldsymbol{U}} \newcommand{\bW}{\boldsymbol{W}} \newcommand{\bX}{\boldsymbol{X}} \newcommand{\bY}{\boldsymbol{Y}} \newcommand{\bZ}{\boldsymbol{Z}} \newcommand{\cotp}{\textcolor[RGB]{48,209,88}{TP}} \newcommand{\cotn}{\textcolor[RGB]{100,210,255}{TN}} \newcommand{\cofp}{\textcolor[RGB]{94,92,230}{FP}} \newcommand{\cofn}{\textcolor[RGB]{191,90,242}{FN}}\\\newcommand{\numcotp}{\textcolor[RGB]{48,209,88}{50}} \newcommand{\numcotn}{\textcolor[RGB]{100,210,255}{30}} \newcommand{\numcofp}{\textcolor[RGB]{94,92,230}{10}} \newcommand{\numcofn}{\textcolor[RGB]{191,90,242}{10}} \DeclareMathOperator*{\argmin}{arg\,min}\end{aligned}\end{align} \]

循环神经网络(RNN)#

学习目标与记号#

  1. 准确解释隐藏状态、参数共享与 BPTT;

  2. 根据公式和张量维度分析梯度消失,并识别常见实现错误;

  3. 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 难以保留长期信息,可继续阅读 LSTMGRU⁠。

  循环神经网络用固定维度的隐藏状态递归汇总历史信息。它不要求数据真的满足有限阶马尔可夫性质;更准确地说,模型假设“与任务相关的历史”可以被列向量 \(\bh_t\in\mathbb{R}^{d_h}\) 近似压缩,并用同一更新函数跨时间复用。

Vanilla RNN#

  对输入列向量 \(\bx_t\in\mathbb{R}^{d_x}\)⁠,最基本的 RNN 单元为

(87)#\[\begin{split}\begin{aligned} \ba_t&=\bW_{xh}\cdot\bx_t+\bW_{hh}\cdot\bh_{t-1}+\bb_h,\\ \bh_t&=\tanh(\ba_t),\\ \bz_t&=\bW_{hy}\cdot\bh_t+\bb_y,\\ \widehat{\by}_t&=g(\bz_t), \end{aligned}\end{split}\]

其中 \(\ba_t,\bh_t\in\mathbb{R}^{d_h}\)\(\bz_t,\widehat{\by}_t\in\mathbb{R}^{d_y}\) 均为列向量,并且

\[\bW_{xh}\in\mathbb{R}^{d_h\times d_x},\quad \bW_{hh}\in\mathbb{R}^{d_h\times d_h},\quad \bW_{hy}\in\mathbb{R}^{d_y\times d_h},\quad \bb_h\in\mathbb{R}^{d_h},\quad \bb_y\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}\)⁠,也可以学习或由另一个编码器给出。

../_images/Figure_6_3_RNN.png

图 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.RNNnn.GRUnn.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 对每一行完成的运算与数学公式中的矩阵乘列向量等价。

输入—输出组织方式#

  1. 序列到单值(many-to-one)⁠:使用最后一个有效状态或所有状态的池化结果完成分类。批量中有填充时,不能直接取统一数组的最后一列。

  2. 等长序列到序列(many-to-many)⁠:每个时间步都有监督信号,例如序列标注。损失应屏蔽填充位置。

  3. 编码器—解码器⁠:编码器读取输入,解码器自回归地产生长度可能不同的输出。只用最后状态作为固定瓶颈时,长序列性能可能受限,注意力机制可直接访问全部编码器状态。

  4. 堆叠 RNN⁠:第 \(l\) 层在时间 \(t\) 的输出成为第 \(l+1\) 层同一时间的输入。层下标与时间下标必须区分。

  5. 双向 RNN⁠:分别从左到右和从右到左编码,再拼接两个方向的状态。它适合能够看到完整输入的标注或编码任务,不适合要求实时、因果预测的场景,因为从右到左的状态使用了未来输入。

../_images/Figure_6_5_Bidirectional_RNN.png

图 25 双向 RNN 同时利用两个方向的上下文#

备注

  变长序列批处理通常需要长度、填充掩码,或框架提供的 packed sequence。分类时应提取每个样本最后一个 有效 时间步,而不是填充后的最后位置。

随时间后向传播(BPTT)#

  设总损失为 \(\mathcal{J}=\sum_{t=1}^{T}\ell_t\)⁠。本节把标量损失关于向量的梯度写成与该向量维度相同的列向量,并记

\[\bg_t=\nabla_{\bz_t}\ell_t\in\mathbb{R}^{d_y}.\]

从最后一个时间步开始进行后向传播,令列向量 \(\bdelta_t=\nabla_{\ba_t}\mathcal{J}\in\mathbb{R}^{d_h}\)⁠,则

(88)#\[\begin{split}\begin{aligned} \bdelta_t &=\left( \bW_{hy}\trans\cdot\bg_t +\bW_{hh}\trans\cdot\bdelta_{t+1} \right) \odot\left(\bone-\bh_t\odot\bh_t\right),\\ \bdelta_{T+1}&=\bzero. \end{aligned}\end{split}\]

  第一项是当前输出损失的贡献,第二项是未来时间步沿循环边传回的贡献。共享参数的梯度必须对所有时间步累加:

\[\begin{split}\begin{aligned} \frac{\partial\mathcal{J}}{\partial\bW_{hy}} &=\sum_{t=1}^{T}\bg_t\cdot\bh_t\trans, &\frac{\partial\mathcal{J}}{\partial\bb_y} &=\sum_{t=1}^{T}\bg_t,\\ \frac{\partial\mathcal{J}}{\partial\bW_{hh}} &=\sum_{t=1}^{T}\bdelta_t\cdot\bh_{t-1}\trans, &\frac{\partial\mathcal{J}}{\partial\bW_{xh}} &=\sum_{t=1}^{T}\bdelta_t\cdot\bx_t\trans,\\ \frac{\partial\mathcal{J}}{\partial\bb_h} &=\sum_{t=1}^{T}\bdelta_t. \end{aligned}\end{split}\]

上式中的 \(\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\) 会如何变化,可以用下面的雅可比矩阵表示:

\[\bJ_t =\frac{\partial\bh_t}{\partial\bh_{t-1}} =\operatorname{diag}\!\left(\bone-\bh_t\odot\bh_t\right)\cdot\bW_{hh}.\]

  跨越 \(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\) 在所有时间步共享的事实。对于较短序列,补齐位置不会再次更新隐藏状态。

点击打开“序列张量与循环神经网络”交互演示

核心推导与实现核验#

核心关系

\[\bh_t=\phi(\bW_h\cdot\bh_{t-1}+\bW_x\cdot\bx_t+\bb).\]

其中,\(\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() 完成这一步。

本节小结#

  1. RNN 在时间上共享同一状态更新参数。

  2. BPTT 的核心是跨时刻雅可比连乘。

  3. 状态边界、截断长度和 梯度裁剪 必须显式管理。

综合练习#

  程序题应固定随机种子、写出维度断言并报告运行环境。全部参考答案见 循环神经网络与随时间后向传播答案⁠。

  1. 前向传播计算。 考虑标量 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\)⁠。

  2. BPTT 推导。 从正文的 RNN 公式和 \(\mathcal{J}=\sum_{t=1}^{T}\ell_t\) 出发,推导 \(\bdelta_t\) 的后向传播公式,以及 \(\partial\mathcal{J}/\partial\bW_{hh}\)\(\partial\mathcal{J}/\partial\bW_{xh}\)⁠。解释共享参数的梯度为何必须对时间求和。

  3. 梯度范数证明。 设从时刻 \(s\)\(t\) 的状态梯度包含 \(\bJ_t\cdot\bJ_{t-1}\cdots\bJ_{s+1}\)⁠。利用矩阵范数的次乘性证明其范数上界,并分别给出该上界指数衰减和可能放大的充分条件;说明这些条件为何不是梯度消失或爆炸的充要条件。

  4. 参数量与计算量。\(d_x=3\)⁠、\(d_h=4\)⁠、\(d_y=2\) 的单层 RNN,计算 \(\bW_{xh}\)⁠、\(\bW_{hh}\)⁠、\(\bb_h\)⁠、\(\bW_{hy}\)\(\bb_y\) 的维度及参数总数。再写出长度为 \(T\) 的单样本前向传播主要乘加运算量的数量级。

  5. 变长序列前向实现。 从零实现批量 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}\)⁠。程序应返回全部隐藏状态和每个样本最后一个有效状态;在填充位置保持旧状态不变,并用不同长度的小批次验证结果。

  6. 手写 BPTT 核验。 为长度为 2 或 3 的小序列实现手写 BPTT,返回每个 \(\bdelta_t\) 和所有参数梯度。把结果分别与 自动微分中心差分 比较,给出相对误差标准,并加入会暴露“只保留最后一步参数梯度”错误的测试。

  7. 状态与截断测试。 为独立序列之间的隐藏状态重置、连续片段之间的状态传递和截断 BPTT 的计算图分离编写测试。验证批量结果与逐条序列结果一致,并检查改变 padding 值不会改变有效位置输出与损失。

  8. 逻辑回归与 RNN。 构造一个类别取决于词元顺序的二分类任务,比较基于词袋特征的逻辑回归与单层 RNN。固定训练/验证/测试划分、随机种子集合、输入词表、训练样本和各自的训练更新预算;报告测试准确率或 F1、训练时间、固定批量预测时间、峰值内存和模型参数量。

  9. 完整与截断 BPTT。 在同一长序列任务上比较完整 BPTT 与截断长度 \(K\in\{16,32,64\}\)⁠。固定数据划分、随机种子集合、RNN 参数初值、优化器、总处理词元数和参数更新次数;报告预测质量、训练时间、固定批量预测时间、峰值内存、循环单元调用次数和梯度范数,并讨论截断造成的长期信用分配偏差。

  10. 隐藏维度与学习率。 比较隐藏维度 \(d_h\in\{16,32,64\}\) 和若干学习率的组合。固定数据划分、随机种子集合、最大训练更新次数、早停规则和评价代码;报告每个组合的预测质量、训练时间、固定批量预测时间、峰值内存与参数量,并区分“容量变化”和“优化设置变化”对结果的影响。