🏠 总目录📚 本教程 13 · RNN 与序列建模
📑 本页目录(点开跳转)

13 · RNN 与序列建模

26 分钟 | 🎁 可快速过,但要理解它「为什么被淘汰」


🎯 一句话

RNN 用一个不断更新的记忆状态来处理序列。 它的思路是对的,但串行长距离遗忘两个硬伤让它输给了 Transformer——理解这两个硬伤,你才真正懂 Transformer 赢在哪。


🤔 零、先问:序列数据到底特殊在哪

   普通表格数据:每一行是独立的
   [年龄 收入 学历] → 买不买

   序列数据:顺序【本身携带信息】
   "这部电影不好看"  vs  "这部电影好看不"
    同样的字,顺序一变,意思就变了 ⭐

前面学过的模型为什么处理不了序列

模型 为什么不行
MLP 输入维度固定,句子长度可变;而且打乱词序结果不变(它看不到顺序)
CNN 能看局部顺序(卷积窗口内),但看不到远距离依赖
树模型 每个特征独立看,完全没有顺序概念

🔑 所以需要一个新机制:能吃变长输入、能记住之前看过什么。 RNN 就是这个问题的第一个像样的答案。


🔄 一、RNN:把上一步的输出喂回来

h₀ h₁ h₂ h₃ x₁ [RNN] y₁ x₂ [RNN] y₂ x₃ [RNN] y₃ 核心公式: hₜ = tanh( W·hₜ₋₁ + U·xₜ + b) 上一步的记忆 当前输入
图上把时间轴展开画成了三个格子,实际上只有一套 W、U 被重复用了三次(这就是 RNN 版的「权重共享」)。⭐ 要盯的是中间那条横线:每一步的 h 既接了当前输入 xₜ,又接了上一步的 hₜ₋₁ —— 「记忆」就是这么传下去的。

💡 人话h 是一个「记忆本」,每读一个新词就更新一次。 同一套权重 W、U 在每个时间步复用——这是 RNN 版的「权重共享」(对比第 12 章 CNN)。

💡 权重共享带来两个好处(和 CNN 完全同构的道理):

   ① 参数量【和序列长度无关】
      → 处理 10 个词和 1000 个词,用的是同一套 W、U

   ② 学到的规律能【跨位置迁移】
      → 在第 3 个词学到的模式,第 300 个词也能用

🔨 五行看懂 RNN

import numpy as np
h = np.zeros(hidden)
for x_t in sequence:                       # ⭐ 注意这个 for 循环
    h = np.tanh(W @ h + U @ x_t + b)       # 更新记忆
    y_t = V @ h                            # 输出

⚠️ 就是这个 for 循环毁了 RNN。记住它,第二节会回来。


💀 二、两个致命硬伤

硬伤 1:梯度消失 → 记不住长距离

   反向传播要穿过 T 个时间步,每步乘一次 W 和 tanh 的导数
   → 梯度 ∝ Wᵀ × (tanh')ᵀ
   → T=100 时,0.9¹⁰⁰ ≈ 0.000027   💀 前面的信息完全学不到

后果:「我在法国长大……所以我说流利的 __」——隔了 50 个词,RNN 已经忘了「法国」。

🔗 这和第 8 章的梯度消失是同一件事,但更严重: 深度网络是"层数"次连乘,RNN 是"序列长度"次连乘—— 而且乘的是同一个 W(权重共享的代价),所以要么一起消失,要么一起爆炸。

⚠️ 梯度爆炸在 RNN 里比在 CNN 里常见得多,标准解法是梯度裁剪

import torch
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)   # ⭐ RNN 几乎必加

LSTM / GRU 的补救:加一条门控的记忆通道

   LSTM 的核心是「细胞状态」C,它像一条传送带:
   ├─ 遗忘门 f:决定丢掉多少旧记忆
   ├─ 输入门 i:决定写入多少新信息
   └─ 输出门 o:决定输出多少

   关键的一行:  Cₜ = f ⊙ Cₜ₋₁ + i ⊙ C̃ₜ
                        ↑ 加法!不是反复相乘 ⭐

   → 当遗忘门 f ≈ 1 时,梯度沿 C 【几乎原样】传回去
   → 和 ResNet 残差连接同一个思想
📐 为什么"加法更新"能救梯度(想看再点)

普通 RNN:$\frac{\partial h_t}{\partial h_{t-1}} = W^\top\text{diag}(\tanh')$ —— 纯乘法,连乘 T 次必然指数衰减。

LSTM:$\frac{\partial C_t}{\partial C_{t-1}} = f_t$ —— 如果遗忘门 $f_t \approx 1$, 连乘 T 次还是 ≈ 1,梯度原样通过

🔗 这和 ResNet 的 ∂a^l/∂a^{l-1} = f'W + I 里那个 +I 是同一个技巧: 给梯度留一条不经过乘法的直通路径。

💡 记住这个模式,它在深度学习里出现了至少四次: LSTM 的细胞状态、ResNet 的残差、Transformer 的残差、Highway Network 的门控。

GRU 是 LSTM 的简化版(两个门),参数少约 25%、更快,效果通常相当。

LSTM GRU
门的数量 3(遗忘/输入/输出) 2(重置/更新)
状态 h 和 C 两条 只有 h
参数量 多 ~33%
效果 长序列上略好 中短序列相当,训练更快 ⭐

LSTM/GRU 缓解了梯度消失,但没根治——几百步以上仍然吃力。

硬伤 2:串行,无法并行 ⭐ 这才是致命的

   还记得那个 for 循环吗?

   RNN:必须算完 h₁ 才能算 h₂,算完 h₂ 才能算 h₃ ...
        → 序列长度 = 计算步数,GPU 的并行能力完全浪费

   Transformer:所有位置【同时】计算注意力
        → 一次矩阵乘法搞定,GPU 吃满 ✅

一个具体的对比

   序列长度 1000,GPU 有 10000 个核:

   RNN:      1000 个串行步,每步只用得上几十个核
              → 大部分核在等
   Transformer:1 次 1000×1000 的矩阵乘法
              → 所有核同时干活 ⭐

   → 同样的硬件,Transformer 能快几十倍

🔑 这是 Transformer 取代 RNN 的根本原因。 不是因为 Transformer 更"聪明",而是因为它能并行,所以能训得更大、喂更多数据可扩展性(scalability)胜过了精巧的结构设计——这是深度学习史上反复出现的教训。

💡 同样的教训还出现在:CNN 打败手工特征、Transformer 打败 CNN 在视觉上(ViT)、 大模型的"scaling law"。一个能吃下更多算力的笨办法,往往赢过吃不下算力的聪明办法。

⚖️ 但 Transformer 也付出了代价

   RNN:         时间 O(T),  内存 O(1)   ← 可以流式处理无限长序列
   Transformer: 时间 O(T²), 内存 O(T²)  ← 序列一长就爆 💀

💡 这就是为什么"长上下文"是大模型的核心难题—— 也是为什么 Mamba 这类"类 RNN"架构最近又被重新拾起。 RNN 输的是并行,赢的是复杂度。


🌉 三、通往注意力:Seq2Seq 的瓶颈

   早期机器翻译(Seq2Seq):

   [编码器 RNN] 读完整个中文句子 → 压成【一个固定长度的向量】
                                          ↓ 💀 瓶颈!
                                   [解码器 RNN] 生成英文

   ❌ 不管句子多长,都压成同一个向量 → 长句信息严重丢失

💡 有多严重:实验显示 Seq2Seq 的翻译质量在句子超过 30 个词后急剧下降—— 因为那个固定向量装不下了。

注意力机制的诞生(2014-2015)

   解码器生成每个词时,不再只看那一个向量,
   而是【回头看编码器的所有隐状态】,动态决定该关注哪几个

   生成 "France" 时 → 重点关注输入里的「法国」那个位置
                        ↑ 这就是注意力权重

💡 注意力的三个立刻可见的好处

好处 说明
打破信息瓶颈 不再压成一个向量,所有位置的信息都还在
梯度捷径 解码器到编码器任意位置只隔 1 步,不用穿过 T 个时间步
可解释 注意力权重能画成热力图,看到"翻译这个词时在看哪里"

💡 然后 2017 年有人问了一个关键问题

既然注意力这么好用,我们还需要 RNN 吗?

答案是那篇论文的标题:《Attention Is All You Need》—— 去掉 RNN,只留注意力,就是 Transformer。

🔗 完整的 Transformer 讲解在全景导论第 2 章, 而下一章会把这条演化线补完。 现在你知道它是怎么来的了。


🗺️ 四、序列任务的四种形态

形态 典型任务 关键点
一对多 图像描述 一张图 → 一句话
多对一 文本分类、情感分析 只用最后一个 h
多对多 机器翻译 长度可不同,需要 encoder-decoder
同步多对多 词性标注、命名实体识别 输入输出一一对应

⚠️ 一个实践细节:多对一任务里,用最后一个 h 不一定最好—— 更常见的做法是对所有时间步的 h 做池化(mean/max),或者加一层注意力。 因为最后一个 h 天然偏向记住结尾的内容。

🔗 Kaggle 第 7 章命名实体识别讲的 CRF / GlobalPointer 就是最后一类的实战方案。


🤔 五、RNN 今天还有用吗

场景 用什么
NLP 几乎所有任务 Transformer(RNN 已基本退场)
超长序列 + 极低延迟 + 小设备 RNN/GRU 仍有一席之地(O(1) 内存,流式处理天然友好
时间序列预测 树模型常更强;LSTM 可作为对比基线
在线/流式场景(如实时推荐) GRU 的增量更新很自然
状态空间模型(Mamba 等) 新一代"类RNN"架构,试图兼得并行与线性复杂度 ⭐

💡 学 RNN 的价值不在于用它,而在于: ① 理解 Transformer 解决了什么问题 ② 理解「门控 + 加法更新」这个防梯度消失的通用思想 (LSTM 的 C、ResNet 的残差、Transformer 的残差,全是它) ③ 理解「并行性是架构设计的一等公民」这个教训


🔗 和站内其他章的关系

相关的地方 这里的对应
第 8 章梯度消失 RNN 版更严重(连乘次数 = 序列长度,且乘同一个 W)
第 12 章 CNN 的权重共享 RNN 跨时间步共享,CNN 跨空间位置共享
第 12 章 ResNet 残差 和 LSTM 的加法式记忆更新是同一思想
推荐算法第 9 章 GRU4Rec 就是本章的 GRU 用在序列推荐上
推荐算法第 9 章 SASRec 用 Transformer 取代 RNN 本章讲的正是"为什么要取代"
全景导论第 2 章注意力 本章第三节讲的是它的起源

✅ 检查点

  1. 为什么 MLP 和 CNN 处理不了序列?
  2. RNN 的核心公式和直觉是什么?权重共享带来哪两个好处?
  3. RNN 的梯度消失为什么比普通深度网络更严重?梯度爆炸怎么处理?
  4. LSTM 靠什么缓解梯度消失?关键的那一行公式是什么?它和 ResNet 有什么共同点?
  5. RNN 的两个硬伤哪个更致命?为什么?
  6. Transformer 相比 RNN 赢在哪、输在哪?
  7. 注意力机制最初是为了解决什么问题?它带来了哪三个好处?
  8. 多对一任务里,为什么"用最后一个 h"不一定最好?
  9. 今天还有哪些场景 RNN 有优势?
👀 答案
  1. MLP 输入维度固定、打乱词序结果不变(看不到顺序);CNN 只能看卷积窗口内的局部顺序,看不到远距离依赖
  2. hₜ = tanh(W·hₜ₋₁ + U·xₜ + b),一个不断更新的记忆本。权重共享带来:①参数量和序列长度无关 ②学到的规律能跨位置迁移
  3. 因为连乘次数是序列长度(可达几百上千)而不是层数(几十),而且每次乘的是同一个 W——要么一起消失要么一起爆炸。爆炸用梯度裁剪 clip_grad_norm_(..., max_norm=1.0)
  4. 门控的细胞状态 C。关键行:Cₜ = f ⊙ Cₜ₋₁ + i ⊙ C̃ₜ,是加法不是反复相乘;f≈1 时 ∂Cₜ/∂Cₜ₋₁ = f ≈ 1,梯度原样通过。和 ResNet 的 +I 是同一技巧:给梯度留一条不经过乘法的直通路径
  5. 串行更致命。梯度消失还能靠 LSTM 缓解,但串行限制了可扩展性——GPU 并行能力完全浪费,训不大。
  6. 赢在并行(时间 O(T) 串行 → 一次矩阵乘法),输在复杂度(O(T²) 时间和内存 vs RNN 的 O(1) 内存)。这就是长上下文难、Mamba 被重拾的原因。
  7. 解决 Seq2Seq 把整句压成一个固定向量的信息瓶颈(超过 30 词质量急剧下降)。三个好处:①打破瓶颈 ②梯度捷径(解码器到编码器任意位置只隔 1 步)③注意力权重可解释。
  8. 最后一个 h 天然偏向记住结尾内容。更好的做法是对所有时间步的 h 做 mean/max 池化,或加一层注意力。
  9. 超长序列+低延迟+小设备(O(1) 内存,流式友好)、在线增量更新场景;时间序列上树模型常更强,LSTM 只作基线。

🛑 可以停在这里

走神救援

序列的特殊在于顺序本身携带信息("不好看"vs"好看不"),MLP打乱词序结果不变、CNN看不到远距离。RNN=不断更新的记忆本 hₜ=tanh(Whₜ₋₁+Uxₜ),跨时间步共享权重(参数量与序列长无关+规律可跨位置迁移)。两硬伤:①梯度消失——比普通网络严重,因为连乘次数=序列长度且乘的是同一个W(爆炸用 clip_grad_norm_)→ LSTM 用Cₜ = f⊙Cₜ₋₁ + i⊙C̃ₜ 的加法更新,f≈1 时梯度原样通过(和ResNet的+I同一技巧,这个模式在深度学习里出现至少四次)②串行不能并行(更致命)——就是那个 for 循环。Transformer赢的根本原因是能并行→能scale("能吃算力的笨办法赢过吃不下算力的聪明办法"),但输在 O(T²) 复杂度(RNN 是 O(1) 内存)——这就是长上下文难、Mamba 被重拾的原因。注意力起源:Seq2Seq把整句压成一个向量的瓶颈(超30词就崩)→ 三好处:打破瓶颈/梯度捷径(只隔1步)/可解释 → 2017年"只要注意力就够了"。多对一别只用最后一个h,要池化。

下一节 👉 14-通往Transformer.md

打卡记录保存在你的浏览器里,首页能看到总进度