chapter 05 / sequence-models · 预计学习时间 120 分钟 · Transformer 前最后一站

序列模型
注意力的诞生

AUDIO // 本章语音导读
本章目录
  1. 序列为什么是新问题
  2. RNN:把网络沿时间折叠
  3. BPTT 与梯度的指数命运
  4. LSTM:给记忆装上阀门
  5. Seq2Seq 与固定向量瓶颈
  6. 注意力:让解码器学会回头看
  7. 交互实验室:注意力对齐热图
  8. 教师强制与曝光偏差
  9. 代码实战:手写 LSTM 单元
  10. 章节测验

序列为什么是新问题

前四章的模型有个隐含假设:输入是定长的,且各维之间没有顺序(图像有空间结构,但尺寸固定)。语言、语音、时间序列打破了这两点:

CNN 的答案是空间上的权重共享,序列模型的答案如出一辙:时间上的权重共享——同一组参数在每个时间步重复使用。这就是循环神经网络。

RNN:把网络沿时间折叠

RNN 维护一个隐藏状态 $\mathbf{h}_t$(网络的「工作记忆」),每读入一个词就更新一次:

$$\mathbf{h}_t = \tanh\big(W_h \mathbf{h}_{t-1} + W_x \mathbf{x}_t + \mathbf{b}\big)$$

同一组 $W_h, W_x$ 用于所有时间步——处理 3 个词和 300 个词用的都是这点参数(变长问题解决),且「位置 7 学到的规律」自动适用于位置 70(时间平移共享)。把循环按时间展开,RNN 就是一个深度等于序列长度、且每层共享权重的特殊深网络——这个视角是理解下一节灾难的钥匙。

VIDEO 01
Recurrent Neural Networks (RNNs), Clearly Explained!!!
StatQuest with Josh Starmer 16:37
观看指南
  • 04:00 沿时间展开(unroll)的画法——把循环看成共享权重的深网络。
  • 10:30 同一个权重在展开图里出现 N 次意味着什么——为 BPTT 的连乘埋伏笔。
  • 13:30 梯度爆炸/消失的直观演示——下一节我们给出严格推导。

BPTT 与梯度的指数命运

训练 RNN 用的还是反向传播,只是沿展开的时间轴回传,得名 BPTT(Backpropagation Through Time)。问题藏在「远距离信用分配」里:第 $t$ 步的损失对 $k$ 步之前隐藏状态的梯度,由链式法则:

$$\frac{\partial \mathbf{h}_t}{\partial \mathbf{h}_k} = \prod_{i=k+1}^{t} \frac{\partial \mathbf{h}_i}{\partial \mathbf{h}_{i-1}} = \prod_{i=k+1}^{t} W_h^\top\, \text{diag}\big(\tanh'(\mathbf{z}_i)\big)$$

同一个矩阵 $W_h$ 连乘 $t-k$ 次。设其最大奇异值为 $\sigma_{\max}$:$\sigma_{\max} < 1$ 时梯度随距离指数消失——50 步外的「小明」对当前梯度的影响约等于零,长程依赖学不到;$\sigma_{\max} > 1$ 时指数爆炸——损失变 NaN。和第 3 章深度方向的梯度消失同源,但更恶劣:深网络各层是不同矩阵,运气可以互相抵消;RNN 是同一个矩阵自乘,命运完全由 $W_h$ 的谱决定,没有侥幸。

爆炸有个简单粗暴的工程解:梯度裁剪——梯度范数超过阈值 $c$ 就整体缩放 $\mathbf{g} \leftarrow c\,\mathbf{g}/\|\mathbf{g}\|$(方向不变、模长封顶)。至今训练 LLM 仍然标配(典型 $c{=}1.0$,你会在第 7 章的训练配置里再见到它)。但消失没有这种贴膏药的解法——信号没了就是没了,需要动手术改架构。

LSTM:给记忆装上阀门

LSTM(1997)的手术方案:在 $\mathbf{h}_t$ 之外增设一条细胞状态 $\mathbf{c}_t$——专用的长期记忆传送带,并用三个可学习的「门」控制读写。门就是一个 sigmoid 输出的 0~1 向量,逐元素相乘实现「软开关」:

$$\mathbf{f}_t = \sigma(W_f [\mathbf{h}_{t-1}, \mathbf{x}_t] + \mathbf{b}_f) \qquad \text{遗忘门:旧记忆保留多少}$$ $$\mathbf{i}_t = \sigma(W_i [\mathbf{h}_{t-1}, \mathbf{x}_t] + \mathbf{b}_i), \qquad \tilde{\mathbf{c}}_t = \tanh(W_c [\mathbf{h}_{t-1}, \mathbf{x}_t] + \mathbf{b}_c) \qquad \text{输入门:新信息写入多少}$$ $$\boxed{\ \mathbf{c}_t = \mathbf{f}_t \odot \mathbf{c}_{t-1} + \mathbf{i}_t \odot \tilde{\mathbf{c}}_t\ } \qquad \text{细胞状态更新}$$ $$\mathbf{o}_t = \sigma(W_o [\mathbf{h}_{t-1}, \mathbf{x}_t] + \mathbf{b}_o), \qquad \mathbf{h}_t = \mathbf{o}_t \odot \tanh(\mathbf{c}_t) \qquad \text{输出门:暴露多少给本步输出}$$

盯住方框里的更新式:$\mathbf{c}_t$ 的主干是加法。求梯度 $\frac{\partial \mathbf{c}_t}{\partial \mathbf{c}_{t-1}} = \text{diag}(\mathbf{f}_t)$——不再有 $W_h$ 连乘!只要遗忘门学会保持开启($f \approx 1$),梯度就能沿细胞状态无衰减地流过几百步。

认出来了吗?$\mathbf{c}_t = \mathbf{f}\odot\mathbf{c}_{t-1} + \cdots$ 与 ResNet 的 $y = x + F(x)$ 是同一个思想:用加法主干替代乘法链,给梯度修高速公路。LSTM(1997)比 ResNet(2015)早了 18 年发现它——深度学习史上最重要的设计模式(恒等捷径)在时间维和深度维各被独立发明一次。GRU(2014)是它的精简版:两个门、无独立细胞状态、参数少 1/4,效果通常打平——工程上「更简单且不差」就是赢。
VIDEO 02
Long Short-Term Memory (LSTM), Clearly Explained
StatQuest with Josh Starmer 20:44
观看指南
  • 05:00 三个门的逐步拆解——配合本章公式逐门对照。
  • 12:00 用具体数字跑一遍记忆的写入与遗忘。
  • 17:30 为什么细胞状态那条「直线」是长程记忆的关键——加法主干的可视化。

Seq2Seq 与固定向量瓶颈

2014 年的 Seq2Seq 把两个 LSTM 接力,解决「输入输出都是变长序列」的任务(翻译、摘要):编码器读完源句,把全部理解压进最后一个隐藏状态 $\mathbf{h}_{enc}$;解码器以它为初始状态,逐词生成目标句。优雅,但有个先天残疾:

固定向量瓶颈:无论源句 3 个词还是 100 个词,全部信息都要挤进同一个(如 512 维的)向量。实测翻译质量随句长明显衰减——长句的开头在编码结束时已被「冲淡」。这不是容量调大能根治的:问题在于「读完才动笔,且只许带一张小抄」的流程本身。人类译者不是这样工作的——他们随时回头看原文。

注意力:让解码器学会回头看

Bahdanau 注意力(2015)就是给解码器这个权利。编码器保留每个位置的隐藏状态 $\mathbf{h}_1,\dots,\mathbf{h}_n$(不再只留最后一个);解码器生成第 $t$ 个词之前,做三步:

① 打分——拿当前解码状态 $\mathbf{s}_{t-1}$ 去问每个源位置「你和我现在要生成的东西相关吗」:

$$e_{tj} = \mathbf{v}^\top \tanh\big(W_s \mathbf{s}_{t-1} + W_h \mathbf{h}_j\big) \qquad j = 1,\dots,n$$

② 归一化——softmax 把分数变成权重(和为 1):

$$\alpha_{tj} = \frac{\exp(e_{tj})}{\sum_{k} \exp(e_{tk})}$$

③ 加权汇总——按权重混合所有源位置的信息,得到本步专属的上下文向量:

$$\mathbf{c}_t = \sum_{j=1}^{n} \alpha_{tj}\, \mathbf{h}_j$$

瓶颈消失了:每生成一个词,解码器都现场定制一份源句摘要。生成「yesterday」时 $\alpha$ 集中在「昨天」,生成「park」时集中在「公园」——对齐是从翻译数据里自己学出来的,没有任何人工词典。

用检索的语言重述这三步:解码状态是一个查询(query),源位置状态既当被匹配的键(key)又当被取用的值(value)——注意力本质上是一次可微分的软检索。把 query/key/value 拆成三个独立投影、把打分简化成点积、再让序列「自己查自己」,就是第 6 章的 self-attention。你已经站在 Transformer 的门口了。

交互实验室:注意力对齐热图

下面是一个模拟训练好的注意力模型(中→英)。必做实验:① 悬停「yesterday」——它的注意力越过整个句子对齐到第 2 个词「昨天」(中英语序不同,这正是注意力优于「按位置对齐」的地方);② 悬停「the」——虚词没有明确对应,注意力弥散——模型表达「不确定」的方式;③ 把温度拉到 0.1——软对齐退化成硬指针;拉到 5——退化成均匀平均(≈ 没有注意力的 Seq2Seq)。温度就是第 2 章 softmax 那个 $T$,在第 9 章 LLM 采样里你会第三次遇到它。

attention.align(zh→en)

悬停/点击英文词 · 中文词的高亮深浅 = 注意力权重 α · 柱状图显示精确数值

注意力机制把解码每个词的计算量从 O(1) 变成了 O(n)(要对全部 n 个源位置打分)。这个代价在 Transformer 里会变成什么?值得吗?
变成著名的 $O(n^2)$:self-attention 中 n 个位置每个都要对其余所有位置打分。这是 Transformer 上下文长度昂贵的根源(也催生了第 7 章会提的线性注意力/混合架构这条研究线)。但历史的回答是:值得——因为换来的是①任意两位置一步直达(最长梯度路径从 O(n) 降到 O(1),长程依赖彻底解决);②所有位置可并行计算(RNN 必须串行等 h_{t-1},GPU 利用率天差地别)。「多花算力买并行和直达」正是 scaling 时代最划算的交易——序章《苦涩的教训》再次应验。

教师强制与曝光偏差

训练 Seq2Seq 时有个微妙选择:解码器第 $t$ 步的输入用什么?用模型上一步自己的输出——错一个词,后面全在垃圾上训练,且无法并行;所以实践用教师强制(teacher forcing):训练时永远喂真实的上一个词。代价是曝光偏差(exposure bias):模型在训练中从未见过自己的错误,推理时一旦走错就进入完全陌生的状态分布,错误滚雪球。

这个 2015 年的老问题在 LLM 时代依然活着:GPT 预训练就是大规模教师强制(每个位置都以真实前文为条件),幻觉的雪球效应与之相关;而第 8 章的 RLHF/RLVR 之所以有效,部分原因恰恰是强化学习让模型在自己生成的轨迹上接受训练——可以视为对曝光偏差的系统性修复。一根线索,牵了十年。

代码实战:手写 LSTM 单元

python · lstm_cell_from_scratch.py
import torch, torch.nn as nn

class LSTMCell(nn.Module):
    """四组门计算合并成一个大矩阵乘(工程标准做法),逐行对应 §4 公式"""
    def __init__(self, d_in, d_h):
        super().__init__()
        self.W = nn.Linear(d_in + d_h, 4 * d_h)   # f, i, c̃, o 一次算完
        self.d_h = d_h

    def forward(self, x, h, c):
        z = self.W(torch.cat([x, h], dim=-1))
        f, i, c_tilde, o = z.chunk(4, dim=-1)
        f = torch.sigmoid(f)          # 遗忘门
        i = torch.sigmoid(i)          # 输入门
        c_tilde = torch.tanh(c_tilde) # 候选记忆
        o = torch.sigmoid(o)          # 输出门
        c = f * c + i * c_tilde       # ★ 加法主干:时间维的残差连接
        h = o * torch.tanh(c)
        return h, c

# 跑一个变长序列,验证状态形状
cell = LSTMCell(d_in=32, d_h=64)
h = torch.zeros(1, 64); c = torch.zeros(1, 64)
for t in range(100):                  # 100 步后梯度依然能流回 t=0(拜加法主干所赐)
    h, c = cell(torch.randn(1, 32), h, c)
print(h.shape, c.shape)               # torch.Size([1, 64]) ×2

# 生产中直接用 nn.LSTM(内部 cuDNN 融合实现,快一个数量级):
# rnn = nn.LSTM(input_size=32, hidden_size=64, num_layers=2, batch_first=True)

章节测验