上期简要回顾: 上篇把 CNN 从 LeNet 一路推到 ResNet,残差连接一举解决深层退化;但卷积天生擅长「图」,面对「一句话、一段视频」这种序列信号,还差一个专门结构。
一、序列数据与 RNN 基本原理
语言、语音、时间序列都是序列数据 ------当前时刻的信息依赖历史。全连接网络假设输入独立,天然不适合这类任务。循环神经网络(RNN)的核心思想是:让网络带"记忆",用隐藏状态 \mathbf{h}_t 在时间上传递信息。
朴素 RNN 的递推式:
:时刻
t的输入(如一个词向量);
:时刻
t的隐藏状态,携带截至当前的全部历史压缩信息;
- 权重
,,跨时刻共享,参数量不随序列长度增长。
PyTorch 中一个最简 RNN 层只需一行:
python
import torch.nn as nn
rnn = nn.RNN(input_size=128, hidden_size=256, num_layers=1, batch_first=True)
二、梯度问题与设计动机
RNN 把同一套权重 W*{hh} 沿时间连乘,反向传播时的梯度包含因子 (W* {hh})\^\\top 的连乘,带来两个经典灾难(与第 2 篇全连接层的梯度消失同源,只是发生在时间维):
-
梯度消失:长序列下早期时刻的梯度指数衰减,模型学不到远距离依赖(比如句首的主语影响句尾的动词);
-
梯度爆炸:连乘结果过大,loss 震荡 NaN。
工程应对
- 梯度裁剪(Gradient Clipping):限制梯度范数上限,专治爆炸:
python
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0)
-
截断反向传播(Truncated BPTT):只沿有限时间步回传,控制计算与梯度路径长度;
-
门控机制(治本):LSTM / GRU 用"门"显式控制信息的保留与遗忘,从结构上缓解长期依赖难题。
三、LSTM 门控机制拆解
长短期记忆网络(LSTM,1997 Hochreiter & Schmidhuber)引入一条**细胞状态(Cell State)**作为"信息高速公路",并用三个门精细调控:

关键设计 :细胞状态 \\mathbf{c}*t 的更新里有一个加性 项(* ),相加而非连乘,使得梯度可以沿
几乎无损地回传很远,这是 LSTM 能捕捉长依赖的根本原因。
四、GRU 结构与对比
门控循环单元(GRU,2014 Cho)是 LSTM 的"轻量版",把遗忘门与输入门合并为更新门 ,并引入重置门
:

三者对比
| 特性 | 朴素 RNN | LSTM | GRU |
|---|---|---|---|
| 门数量 | 0 | 3(遗忘/输入/输出) | 2(更新/重置) |
| 额外状态 | 无 | 细胞状态 \\mathbf{c}_t | 无(合并到 \\mathbf{h}) |
| 参数量 | 最少 | 最多(约 RNN 的 4 倍) | 介于两者之间 |
| 长依赖能力 | 弱 | 强 | 强(接近 LSTM) |
| 训练速度 | 快 | 慢 | 中等 |
经验法则:多数任务 GRU 与 LSTM 表现接近,优先用 GRU 省算力;极端长序列或大模型再上 LSTM。
五、IMDB 情感分类三模型对比实验
我们用 IMDB 影评二分类数据集,在同一预处理与超参下训练 RNN / LSTM / GRU,公平对比。
5.1 数据预处理(共用)
python
import torch
from torchtext.datasets import IMDB
from torchtext.data.utils import get_tokenizer
from collections import Counter
from torch.nn.utils.rnn import pad_sequence
tokenizer = get_tokenizer("basic_english")
counter = Counter()
train_iter = IMDB(split="train")
for label, line in train_iter:
counter.update(tokenizer(line))
vocab = torchtext.vocab.vocab(counter, min_freq=5)
# 截断到最长 256 词,批量 padding
MAX_LEN = 256
5.2 三种模型定义(仅 Cell 不同)
python
import torch.nn as nn
class SentimentRNN(nn.Module):
def __init__(self, vocab_size, embed=128, hidden=256, cell="LSTM"):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed, padding_idx=0)
cells = {"RNN": nn.RNN, "LSTM": nn.LSTM, "GRU": nn.GRU}
self.rnn = cells[cell](embed, hidden, batch_first=True)
self.fc = nn.Linear(hidden, 2)
def forward(self, x):
x = self.embedding(x)
out, _ = self.rnn(x) # out: (B, L, hidden)
return self.fc(out[:, -1, :]) # 取最后时刻隐状态分类
5.3 训练主循环(含梯度裁剪)
python
def train(model, loader, epochs=5, lr=1e-3):
opt = torch.optim.Adam(model.parameters(), lr=lr)
crit = nn.CrossEntropyLoss()
for ep in range(epochs):
model.train(); total, correct = 0, 0
for xb, yb in loader:
opt.zero_grad()
logits = model(xb)
loss = crit(logits, yb)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) # 防爆炸
opt.step()
total += yb.size(0)
correct += (logits.argmax(1) == yb).sum().item()
print(f"epoch {ep+1}: acc={correct/total:.3f}")
5.4 典型复现结果
在默认超参(embed=128, hidden=256, epoch=5, batch=64, Adam, 梯度裁剪 5.0)下,多次运行取典型区间:
| 模型 | 测试准确率(典型区间) | 单 epoch 耗时 | 备注 |
|---|---|---|---|
| 朴素 RNN | 78% ~ 81% | 基准 | 长依赖弱,易早熟饱和 |
| LSTM | 85% ~ 88% | 最慢 | 长程特征捕获最好 |
| GRU | 84% ~ 87% | 中等 | 与 LSTM 接近、更轻 |
说明:IMDB 句子普遍较短,RNN 的劣势不算极端;在真正长文本(如文档级分类)上,LSTM/GRU 对 RNN 的领先会更明显。精确数值随随机种子、vocab 截断、padding 策略浮动 1~2 个百分点,建议读者跑通脚本后贴出自己的结果。
观察结论:
-
LSTM/GRU 凭借门控,明显优于朴素 RNN,印证了"长依赖"是核心痛点;
-
GRU 以更少参数达到接近 LSTM 的效果,性价比高;
-
三条 loss 曲线都随 epoch 平滑下降,但 RNN 后期趋于平缓(记忆上限),LSTM/GRU 仍能继续优化。
六、局限性与演进方向
RNN 系(含 LSTM/GRU)有个结构性硬伤:时间步必须串行计算,无法像 CNN/Transformer 那样并行,训练和推理都慢;且超长序列下记忆仍会衰减。
2017 年 Transformer 用自注意力彻底抛弃循环,实现全并行 + 任意长依赖直连,成为 NLP 乃至多模态的新 backbone。下一篇我们就正式拆解它。
下一篇《注意力机制与 Transformer 核心原理剖析》我们直击现代 AI 的基石。
参考资料
-
Rumelhart, D., et al. (1986). Learning Representations by Back-propagating Errors. Nature.
-
Hochreiter, S., & Schmidhuber, J. (1997). Long Short-Term Memory. Neural Computation.
-
Cho, K., et al. (2014). Learning Phrase Representations using GRU. EMNLP.
-
Pascanu, R., et al. (2013). On the Difficulty of Training Recurrent Neural Networks. ICML.
往期回顾
📚 本篇出自专栏《深学之路:从感知机到智能体》 ------ 从感知机到大模型智能体的完整深度学习连载。每篇附可运行代码与原创实验,收藏专栏不迷路。