指南
基于 cs224n 讲义 + Understanding LSTM Networks(Chris Olah)+ PyTorch RNN 文档编写,对照当前版本行为
速查
- LSTM 核心:细胞状态
C_t像传送带,三门(遗忘/输入/输出)控制信息的丢弃、写入、读出,让梯度长距离流动 - 三门公式:遗忘门
f_t=σ(W_f·[h_{t-1},x_t]);输入门i_t=σ(W_i·[...]),候选C̃_t=tanh(...);输出门o_t=σ(W_o·[...]) - 细胞更新:
C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t(遗忘旧 + 写入新) - 隐藏态:
h_t = o_t ⊙ tanh(C_t)(输出门控制读出多少) - GRU 简化:更新门
z_t(合并遗忘+输入)、重置门r_t,无独立细胞状态,参数少 25% - 双向 RNN:正反向各跑一遍,每位置拼接
[h_forward, h_backward],需完整序列(不能实时) - Seq2Seq:编码器 LSTM 把输入序列压成上下文向量 → 解码器 LSTM 从该向量展开生成输出
- Attention 雏形:解码每步动态对编码器所有隐藏态算权重,加权求和——突破单一上下文向量瓶颈
- BPTT 截断:长序列按固定步长(如 35)截断反向传播,权衡梯度质量与计算成本
- 为何被 Transformer 取代:①串行无法并行 ②长距离仍需逐步传递 ③单一上下文向量瓶颈——Attention 直接「直连」任意两位置解决全部
LSTM 门控机制深析
LSTM(Long Short-Term Memory)通过精心设计的门控解决朴素 RNN 的梯度消失。Chris Olah 的经典解读把它比喻成一条「传送带 + 三道阀门」。
细胞状态:传送带
C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t
↑ ↑
遗忘旧信息 写入新信息细胞状态 C_t 是 LSTM 的核心:它沿时间步线性流动,信息加上去或移除都由门控制。关键是这个加法更新的梯度近似为 1——梯度可以沿传送带长距离回流而不衰减,这是 LSTM 能学长距离依赖的数学根基。
三道门
| 门 | 公式 | 作用 |
|---|---|---|
| 遗忘门(forget gate) | f_t = σ(W_f·[h_{t-1}, x_t] + b_f) | 输出 [0,1],决定细胞状态每个维度保留多少旧信息(0 全弃,1 全留) |
| 输入门(input gate) | i_t = σ(W_i·[h_{t-1}, x_t] + b_i) | 决定每个维度写入多少新信息 |
| 候选值 | C̃_t = tanh(W_C·[h_{t-1}, x_t] + b_C) | 生成 [-1,1] 的候选新信息 |
| 输出门(output gate) | o_t = σ(W_o·[h_{t-1}, x_t] + b_o) | 决定细胞状态有多少输出到隐藏状态 |
完整前向流程
1. 遗忘:f_t = σ(W_f·[h_{t-1}, x_t]) ← 该忘掉什么旧记忆
2. 候选:C̃_t = tanh(W_C·[h_{t-1}, x_t]) ← 新的候选信息
3. 输入门:i_t = σ(W_i·[h_{t-1}, x_t]) ← 该写入多少新信息
4. 更新细胞:C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t ← 传送带更新
5. 输出门:o_t = σ(W_o·[h_{t-1}, x_t]) ← 该输出多少
6. 隐藏态:h_t = o_t ⊙ tanh(C_t) ← 当前输出import torch
import torch.nn as nn
# PyTorch 已封装好,无需手写门控
lstm = nn.LSTM(input_size=64, hidden_size=128, num_layers=1, batch_first=True)
# 内部就是上述 6 步计算,输出 (output, (h_n, c_n))直觉:细胞状态是「长期记忆」(可以跨越数百步),隐藏状态是「短期输出」(每步都更新)。遗忘门决定「清空哪些旧记忆」,输入门决定「写入哪些新记忆」,输出门决定「此刻读出哪些记忆」。这种设计让 LSTM 像一个有擦写控制的笔记本,远比朴素 RNN 的「全部覆盖式更新」更适合长序列。
GRU:LSTM 的轻量化
GRU(Gated Recurrent Unit, Cho et al. 2014)把 LSTM 的三门简化为两门,合并细胞状态与隐藏状态,参数更少、训练更快,多数任务效果与 LSTM 相当。
| LSTM | GRU | 简化点 |
|---|---|---|
| 遗忘门 f + 输入门 i | 更新门 z | 合并为一个门:z_t = σ(...) 直接控制新旧比例 |
| 独立细胞状态 C | 隐藏状态 h 兼任 | 不再维护单独的传送带 |
| 输出门 o | 重置门 r | 控制计算候选隐藏态时用多少旧记忆 |
GRU 公式:
z_t = σ(W_z·[h_{t-1}, x_t]) # 更新门:决定保留多少旧状态
r_t = σ(W_r·[h_{t-1}, x_t]) # 重置门:决定计算候选时用多少旧状态
h̃_t = tanh(W·[r_t ⊙ h_{t-1}, x_t]) # 候选隐藏态
h_t = (1 - z_t) ⊙ h_{t-1} + z_t ⊙ h̃_t # 更新(注意是 1-z 与 z 互补)选择建议:新项目优先 GRU——参数少 25%、训练快、效果通常与 LSTM 持平。若任务对长距离依赖极敏感(如长文档摘要),可对比 LSTM 看是否有提升。
双向 RNN
单向 RNN 的 h_t 只看到 [x_1, ..., x_t](左侧上下文),但语言理解常需右侧上下文(如「我喜欢吃___」要填「苹果」得看后面)。双向 RNN 同时跑正向与逆向两个 RNN,把两者的隐藏态拼接。
lstm = nn.LSTM(64, 128, batch_first=True, bidirectional=True)
# 输出 shape: [B, seq_len, 256] ← 正向 128 + 逆向 128 拼接- 优点:每个位置都有完整左右上下文,表示质量显著提升
- 限制:必须拿到完整序列才能跑,不能用于实时流式(语音识别实时转写、在线翻译)
- 用途:离线 NLP 任务(命名实体识别、句法分析、文本分类)几乎都用双向
注意:双向 RNN 在生成任务(如机器翻译解码器、语言模型)中不能随意用——生成时未来词还未产生,无法看右侧。它主要用在「理解型」任务上。
Seq2Seq 与编码器-解码器
Seq2Seq(Sequence-to-Sequence)把变长输入序列映射到变长输出序列,是机器翻译、摘要、对话的开山范式(Sutskever et al. 2014)。
编码器(Encoder LSTM) 解码器(Decoder LSTM)
┌─────────────────────┐ ┌─────────────────────┐
│ 我 → 爱 → 编 → 程 │ │ I → love → coding │
└──────────┬──────────┘ └──────────▲──────────┘
│ │
└──→ 上下文向量 c ──────────────────┘
(编码器最后隐藏态)工作机制:
- 编码:编码器 LSTM 读完整个输入序列,把最终隐藏状态(与细胞状态)作为「上下文向量 c」
- 解码:解码器 LSTM 以 c 为初始状态,逐步生成输出序列,每步用上一步生成的词作为下一步输入
致命瓶颈:基础 Seq2Seq 把整个输入压成固定长度的单一向量 c。输入越长,信息损失越严重——这是 Seq2Seq 在长句翻译上崩盘的根本原因。
Attention 雏形
为破解单一上下文向量的瓶颈,Bahdanau et al. 2014 提出 Attention:解码每一步不再依赖固定 c,而是动态计算对编码器每个隐藏态的关注权重,加权求和得到该步专属的上下文。
解码第 t 步:
对每个编码器隐藏态 h_i 计算注意力分数 e_ti = score(decoder_state, h_i)
softmax 归一化:α_ti = softmax(e_ti)
加权求和上下文:c_t = Σ α_ti · h_i
解码器用 c_t 生成第 t 个输出词Attention 的革命性:让解码器「直连」编码器任意位置,无需信息逐步传递,彻底打破长距离瓶颈。这一思想被 Vaswani et al. 2017 推到极致——干脆完全抛弃 RNN,只用 Attention,这就是 Transformer。
本叶只讲 RNN 系列及其如何孕育出 Attention。Transformer 的自注意力、多头、位置编码等机制在独立的「Transformer」叶展开。
为何 RNN 被 Transformer 取代
2017 年「Attention Is All You Need」发表后,RNN 在 NLP 主力任务上迅速被 Transformer 取代。三个工程动因:
| RNN 的缺陷 | Transformer 的解法 |
|---|---|
串行计算:h_t 依赖 h_{t-1},无法在时间步上并行,训练慢 | 全并行:Self-Attention 同时计算所有位置对,GPU 充分利用 |
| 长距离建模仍弱:即便有门控,信息仍需逐步传递,远距离衰减 | 直接直连:任意两位置一步 Attention 即可交互,距离无衰减 |
| 单一上下文瓶颈:Seq2Seq 把输入压成固定向量 | 动态聚合:每步对全部编码器位置动态加权,无信息瓶颈 |
结果:
- 机器翻译:Transformer(2017+)全面取代 LSTM Seq2Seq
- 语言模型:GPT/BERT 系列全部基于 Transformer,RNN 语言模型退场
- 文本分类/NER:BERT 等预训练 Transformer 主导,BiLSTM 退为轻量备选
RNN 仍保留的场景:
- 资源受限设备:移动端、嵌入式,模型必须小且串行无妨
- 实时流式:在线语音识别、实时翻译,必须增量处理不能等完整序列
- 时序预测:传统时间序列(股票、传感器),数据量小,轻量 GRU 仍实用
- 教学:理解序列建模与门控思想的最佳起点
cs224n 现行课程也以 Transformer 为主线,RNN 作为历史背景与对比对象讲授。掌握 RNN 仍是理解 Attention 为何被发明、为何有效的前提。
反模式(生产坑)
- 朴素 RNN 处理长序列:超过 20 步梯度消失,学不到长依赖。正确:长序列用 LSTM/GRU
- 不裁剪梯度:LSTM/GRU 仍可能爆炸,训练中 loss 突然变 NaN。正确:每步
clip_grad_norm_(model.parameters(), 5.0) - 双向 RNN 用于生成:解码器看右侧未来词,泄漏答案。正确:生成任务只用单向;理解任务才用双向
- batch_first 维度搞反:PyTorch 默认
[seq, batch, feat],搞反会维度对不上报错或静默算错。正确:明确设batch_first=True - 超长序列不截断 BPTT:显存爆炸且梯度质量差。正确:用 truncated BPTT,按固定步长(如 35)截断
- 新项目首选朴素 RNN:2017 年后默认应选 Transformer,或至少 LSTM/GRU;朴素 RNN 仅作教学
下一步
- 参考:RNN API 速查 + 超参默认值表 + 官方资源