【人工智能】【深度学习】 ① RNN核心算法介绍:从循环结构到LSTM门控机制(白话版)
📖目录
前言
💡 提示:本系列按《全解深度学习——九大核心算法》目录顺序展开。由于笔者此前已在 这篇博客 中系统讲解了 CNN 的结构、数学原理与多模态应用,故本篇直接切入 循环神经网络(RNN),开启序列建模之旅。
1. 为什么需要 RNN?——大白话理解“记忆”的价值
想象你在读一句话:“我昨天去了银行,取了____。”
你大概率会填“钱”,而不是“鱼”或“星星”。为什么?因为你记得前面的词“银行”和“取”,这些上下文信息帮助你预测下一个词。
传统神经网络(如全连接网络、CNN)是“无记忆”的:每个输入独立处理,无法利用历史信息。
而 RNN(Recurrent Neural Network) 的核心思想就是:让网络拥有“短期记忆”,把上一时刻的输出作为当前时刻的输入之一,从而捕捉时间序列中的依赖关系。
✅ 应用场景:机器翻译、语音识别、股票预测、文本生成、DNA序列分析……
2. RNN 基本结构与前向传播
2.1 网络架构图(文字描述)
t=1 t=2 t=3
x₁ ──► x₂ ──► x₃ ──► ...
│ │ │
▼ ▼ ▼
[RNN] → [RNN] → [RNN] → ...
│ │ │
▼ ▼ ▼
h₁ h₂ h₃
│ │ │
▼ ▼ ▼
y₁ y₂ y₃
- xₜ:t 时刻的输入(如一个词的 one-hot 向量)
- hₜ:t 时刻的隐藏状态(即“记忆”)
- yₜ:t 时刻的输出(如预测下一个词的概率分布)
2.2 核心公式(前向传播)+ 推导解释
RNN 的前向传播由两个简洁但深刻的公式组成:
h t = tanh ( W h x x t + W h h h t − 1 + b h ) y t = W y h h t + b y \begin{aligned} h_t &= \tanh(W_{hx} x_t + W_{hh} h_{t-1} + b_h) \\ y_t &= W_{yh} h_t + b_y \end{aligned} htyt=tanh(Whxxt+Whhht−1+bh)=Wyhht+by
乍看之下,这只是带了“上一时刻隐藏状态”的普通神经网络。但它的设计背后有清晰的推导逻辑和工程直觉。下面我们一步步拆解。
2.2.1 公式是怎么来的?——从全连接网络到循环结构
-
起点:普通神经网络
对于静态输入 x x x,单层网络输出为:
h = tanh ( W x + b ) h = \tanh(W x + b) h=tanh(Wx+b)
每个样本独立处理,无历史依赖。 -
需求:序列需要“记忆”
在处理“我吃了一个___”时,模型必须知道前面出现了“吃”,才能预测“苹果”。
→ 因此,当前状态 h t h_t ht 应同时依赖 当前输入 x t x_t xt 和 历史状态 h t − 1 h_{t-1} ht−1。 -
自然扩展:拼接输入
最直接的做法是把 x t x_t xt 和 h t − 1 h_{t-1} ht−1 拼成一个大向量:
h t = tanh ( W all [ x t h t − 1 ] + b h ) h_t = \tanh\left( W_{\text{all}} \begin{bmatrix} x_t \\ h_{t-1} \end{bmatrix} + b_h \right) ht=tanh(Wall[xtht−1]+bh) -
拆解权重矩阵
将 W all W_{\text{all}} Wall 按维度拆成两部分:- W h x W_{hx} Whx:作用于输入 x t x_t xt
- W h h W_{hh} Whh:作用于历史状态 h t − 1 h_{t-1} ht−1
于是:
W all [ x t h t − 1 ] = W h x x t + W h h h t − 1 W_{\text{all}} \begin{bmatrix} x_t \\ h_{t-1} \end{bmatrix} = W_{hx} x_t + W_{hh} h_{t-1} Wall[xtht−1]=Whxxt+Whhht−1 -
最终形式
代入后即得 RNN 隐藏状态更新公式:
h t = tanh ( W h x x t + W h h h t − 1 + b h ) h_t = \tanh(W_{hx} x_t + W_{hh} h_{t-1} + b_h) ht=tanh(Whxxt+Whhht−1+bh)
✅ 关键洞察:RNN 本质是一个带反馈连接的全连接网络,通过将上一时刻的输出重新作为输入,实现“短期记忆”。
输出层则更简单:只需基于当前“理解”(即
h
t
h_t
ht)做预测,因此用线性变换即可:
y
t
=
W
y
h
h
t
+
b
y
y_t = W_{yh} h_t + b_y
yt=Wyhht+by
2.2.2 如何快速记忆这个公式?
记住这个口诀:
“新输入 + 旧记忆 → 融合 → 激活 → 新记忆”
对应步骤:
- 新输入: W h x x t W_{hx} x_t Whxxt —— 当前时刻看到的内容
- 旧记忆: W h h h t − 1 W_{hh} h_{t-1} Whhht−1 —— 上一时刻留下的印象
- 融合 + 偏置:两者相加再加 b h b_h bh
- 激活: tanh ( ⋅ ) \tanh(\cdot) tanh(⋅) 压缩到 [-1, 1],引入非线性
- 新记忆:结果就是 h t h_t ht,用于下一步和输出
输出则遵循:“记忆 → 线性映射 → 预测”
2.2.3 举个栗子:字符级语言模型
假设我们在预测单词 “hello” 中第 3 个字母之后的下一个字符。
- 输入: x 3 = ’l’ x_3 = \text{'l'} x3=’l’(one-hot 向量)
- 上一隐藏状态: h 2 = [ 0.2 , − 0.1 , 0.5 ] T h_2 = [0.2, -0.1, 0.5]^T h2=[0.2,−0.1,0.5]T
- 权重(简化):
- W h x x 3 = [ 0.3 , 0.1 , 0 ] T W_{hx} x_3 = [0.3, 0.1, 0]^T Whxx3=[0.3,0.1,0]T
- W h h h 2 = [ 0.1 , − 0.05 , 0.25 ] T W_{hh} h_2 = [0.1, -0.05, 0.25]^T Whhh2=[0.1,−0.05,0.25]T
计算:
h
3
=
tanh
(
[
0.3
,
0.1
,
0
]
+
[
0.1
,
−
0.05
,
0.25
]
)
=
tanh
(
[
0.4
,
0.05
,
0.25
]
)
≈
[
0.38
,
0.05
,
0.24
]
h_3 = \tanh\left( [0.3, 0.1, 0] + [0.1, -0.05, 0.25] \right) = \tanh([0.4, 0.05, 0.25]) \approx [0.38, 0.05, 0.24]
h3=tanh([0.3,0.1,0]+[0.1,−0.05,0.25])=tanh([0.4,0.05,0.25])≈[0.38,0.05,0.24]
这个 h 3 h_3 h3 编码了 “hel” 的语义信息,模型据此更可能预测下一个字母是 “l” 而非 “z”。
2.2.4 为什么这样设计合理?
| 组件 | 作用 |
|---|---|
| W h x x t W_{hx} x_t Whxxt | 引入当前时刻的新信息 |
| W h h h t − 1 W_{hh} h_{t-1} Whhht−1 | 保留并转换历史记忆(“循环”的来源) |
| tanh \tanh tanh | 非线性激活,控制数值范围,缓解梯度爆炸 |
| 参数共享(所有 t 共用 W h h W_{hh} Whh) | 使模型能处理任意长度序列,且参数量不随序列增长 |
📌 一句话理解 RNN 的灵魂:
“用同一个大脑,在不同时间看不同的东西,但记得之前看到过什么。”
这正是 RNN 区别于 CNN、MLP 的核心——时间维度上的信息传递。
3. RNN 的训练:BPTT(Backpropagation Through Time)
当然可以!以下是按照你要求全面增强后的 “1. 损失函数(以分类为例)+ 解释” 小节,融合了大白话解释、记忆口诀、具体举例和设计动机,风格与你博客一致,逻辑清晰、通俗易懂又不失专业性:
3.1 损失函数(以分类为例)+ 解释
在 RNN 中,若每个时间步都有监督信号(如语言模型中每个词都要预测下一个词),总损失通常定义为所有时刻损失之和:
L = ∑ t = 1 T CrossEntropy ( y t , y t true ) \mathcal{L} = \sum_{t=1}^{T} \text{CrossEntropy}(y_t, y_t^{\text{true}}) L=t=1∑TCrossEntropy(yt,yttrue)
其中:
- T T T:序列总长度
- y t y_t yt:模型在 t 时刻的预测输出(通常是 softmax 后的概率分布)
- y t true y_t^{\text{true}} yttrue:t 时刻的真实标签(如 one-hot 向量)
3.1.1 大白话解释:为什么要把每个时刻的损失加起来?
想象你在教一个学生写作文。
他每写一个词,你就立刻判断“这个词用得对不对”。
- 如果他在“我今天吃了___”后面写了“饭”,你点头;
- 如果写了“火箭”,你摇头。
RNN 的训练就像这样——每一步都要负责。
所以不能只看最后一词对不对,而是从第一个词到最后一个词,每一步都算分,错得越多,总分越低。
✅ 这种训练方式叫 “Teacher Forcing”:训练时用真实历史输入(而非模型自己生成的)来预测下一步,加速收敛。
3.1.2 记忆口诀
“步步都要对,错一步就扣分,全对才满分。”
对应公式含义:
- “步步都要对” → 每个 t t t 都有损失项
- “错一步就扣分” → 单个时刻预测不准, CrossEntropy \text{CrossEntropy} CrossEntropy 变大
- “全对才满分” → 所有 y t = y t true y_t = y_t^{\text{true}} yt=yttrue 时, L → 0 \mathcal{L} \to 0 L→0
3.1.3 再举个栗子:情感分析中的序列标注(简化版)
假设我们有一个极简任务:判断句子中每个词是否“表达情绪”。
句子:“这个电影太棒了!”
分词后:[“这个”, “电影”, “太”, “棒”, “了”]
真实标签(1=情绪词,0=非情绪词):
y
true
=
[
0
,
0
,
1
,
1
,
0
]
y^{\text{true}} = [0, 0, 1, 1, 0]
ytrue=[0,0,1,1,0]
模型输出(经 softmax 后取 argmax):
y
=
[
0
,
0
,
1
,
0
,
0
]
y = [0, 0, 1, 0, 0]
y=[0,0,1,0,0]
→ 第4个词“棒”被错判为 0!
那么损失为:
L
=
CE
1
+
CE
2
+
CE
3
+
CE
4
+
CE
5
\mathcal{L} = \text{CE}_1 + \text{CE}_2 + \text{CE}_3 + \text{CE}_4 + \text{CE}_5
L=CE1+CE2+CE3+CE4+CE5
其中
CE
4
\text{CE}_4
CE4 会显著大于其他项,拉高总损失,促使模型重点修正“棒”这类词的判断能力。
3.1.4 为什么用交叉熵(CrossEntropy)?
- 它天然适合多分类问题(如词汇表大小为 10000,每步预测下一个词)
- 当预测概率接近真实标签时,损失趋近于 0;反之急剧上升,梯度信号强
- PyTorch 中
nn.CrossEntropyLoss内部已包含 softmax,直接输入 logits 即可
3.1.5 注意:不是所有 RNN 任务都用这种损失!
- 机器翻译 / 文本生成:通常只在解码阶段每个输出词计算损失
- 情感分类(整句分类):可能只用最后一个隐藏状态做预测,此时损失为单个 CrossEntropy
- 语音识别:可能用 CTC(Connectionist Temporal Classification)等特殊损失
但在标准序列建模(如语言模型)中,逐时间步求和是最经典、最直观的做法。
📌 一句话总结:
“RNN 的损失 = 所有时间步犯错的总代价” —— 时间越长,责任越大。
3.2 梯度计算(关键问题:梯度消失)+ 推导思路
在深度学习中,模型能否学会,取决于梯度是否能有效传递。对于 RNN 而言,训练的核心方法是 BPTT(Backpropagation Through Time,通过时间反向传播)——即将时间序列展开为一个深度网络,然后像普通神经网络一样进行反向传播。
然而,正是这种看似自然的展开方式,暴露出 RNN 最致命的缺陷:梯度消失(Vanishing Gradient Problem)。它使得 RNN 几乎无法捕捉跨越数十甚至上百个时间步的依赖关系,比如一句话开头的主语对结尾谓语的影响。
下面,我们将从直观理解 → 数学推导 → 数值实验 → 工程后果四个层面,层层深入剖析这一问题。
3.2.1 直观理解:为什么“记忆会断”?
想象你在玩一个“传话游戏”:
第一个人说:“明天下午三点在图书馆见面。”
第二个人听成:“明天下午三点在书店见面。”
第三个人:“明天三点在书店……”
……
到第十个人时,变成了:“今天去吃饭。”
每传一次,信息就丢失一点。RNN 的隐藏状态 h t h_t ht 就像这个“话”,在时间步之间传递。而反向传播的梯度,则是从“最后一人”往回追溯“谁传错了”的信号。
但问题在于:每次传递都会压缩信号。传得越远,追溯信号越弱,最终早期参与者根本收不到反馈——他们不知道自己错在哪,也就无法改正。
这就是梯度消失的本质:长期依赖的学习信号在反向传播中被指数级衰减,导致早期时间步几乎无法更新参数。
3.2.2 数学推导:梯度是如何一步步消失的?
我们以循环权重 W h h W_{hh} Whh 为例,分析其梯度 ∂ L ∂ W h h \frac{\partial \mathcal{L}}{\partial W_{hh}} ∂Whh∂L 的结构。
Step 1:总损失对 W h h W_{hh} Whh 的梯度
由于 W h h W_{hh} Whh 在每个时间步都被使用,根据链式法则:
∂ L ∂ W h h = ∑ t = 1 T ∂ L t ∂ W h h \frac{\partial \mathcal{L}}{\partial W_{hh}} = \sum_{t=1}^{T} \frac{\partial \mathcal{L}_t}{\partial W_{hh}} ∂Whh∂L=t=1∑T∂Whh∂Lt
其中 L t \mathcal{L}_t Lt 是 t 时刻的损失(如 CrossEntropy)。
Step 2:单个时间步的梯度分解
考虑 ∂ L t ∂ W h h \frac{\partial \mathcal{L}_t}{\partial W_{hh}} ∂Whh∂Lt。因为 W h h W_{hh} Whh 影响了从第 1 步到第 t 步的所有隐藏状态,所以:
∂ L t ∂ W h h = ∑ k = 1 t ∂ L t ∂ h t ⋅ ( ∏ i = k + 1 t ∂ h i ∂ h i − 1 ) ⋅ ∂ h k ∂ W h h \frac{\partial \mathcal{L}_t}{\partial W_{hh}} = \sum_{k=1}^{t} \frac{\partial \mathcal{L}_t}{\partial h_t} \cdot \left( \prod_{i=k+1}^{t} \frac{\partial h_i}{\partial h_{i-1}} \right) \cdot \frac{\partial h_k}{\partial W_{hh}} ∂Whh∂Lt=k=1∑t∂ht∂Lt⋅(i=k+1∏t∂hi−1∂hi)⋅∂Whh∂hk
这个公式有三层含义:
| 部分 | 含义 |
|---|---|
| ∂ L t ∂ h t \frac{\partial \mathcal{L}_t}{\partial h_t} ∂ht∂Lt | 当前损失对当前记忆的敏感度(误差信号) |
| ∏ i = k + 1 t ∂ h i ∂ h i − 1 \prod_{i=k+1}^{t} \frac{\partial h_i}{\partial h_{i-1}} ∏i=k+1t∂hi−1∂hi | 从 t 到 k 的“记忆传递链”(核心问题所在) |
| ∂ h k ∂ W h h \frac{\partial h_k}{\partial W_{hh}} ∂Whh∂hk | 当前记忆对参数的直接依赖 |
Step 3:聚焦“传递链”——消失的根源
我们重点看中间的连乘项:
∂
h
i
∂
h
i
−
1
=
∂
∂
h
i
−
1
tanh
(
W
h
x
x
i
+
W
h
h
h
i
−
1
+
b
h
)
=
diag
(
1
−
tanh
2
(
z
i
)
)
⏟
tanh
′
(
z
i
)
⋅
W
h
h
\frac{\partial h_i}{\partial h_{i-1}} = \frac{\partial}{\partial h_{i-1}} \tanh(W_{hx} x_i + W_{hh} h_{i-1} + b_h) = \underbrace{\text{diag}\left(1 - \tanh^2(z_i)\right)}_{\tanh'(z_i)} \cdot W_{hh}
∂hi−1∂hi=∂hi−1∂tanh(Whxxi+Whhhi−1+bh)=tanh′(zi)
diag(1−tanh2(zi))⋅Whh
- tanh ′ ( z ) = 1 − tanh 2 ( z ) ∈ ( 0 , 1 ] \tanh'(z) = 1 - \tanh^2(z) \in (0, 1] tanh′(z)=1−tanh2(z)∈(0,1],且在大多数情况下远小于 1(例如当 ∣ z ∣ > 1 |z| > 1 ∣z∣>1 时, tanh ′ ( z ) < 0.4 \tanh'(z) < 0.4 tanh′(z)<0.4)。
- W h h W_{hh} Whh 是一个矩阵,其谱范数(最大奇异值)若小于 1,也会进一步压缩梯度。
因此,每一步的 Jacobian 矩阵 ∂ h i ∂ h i − 1 \frac{\partial h_i}{\partial h_{i-1}} ∂hi−1∂hi 的范数通常小于 1。
假设平均每步的缩放因子为 λ < 1 \lambda < 1 λ<1,那么跨越 d = t − k d = t - k d=t−k 步的梯度大小约为:
∥ ∏ i = k + 1 t ∂ h i ∂ h i − 1 ∥ ≈ λ d \left\| \prod_{i=k+1}^{t} \frac{\partial h_i}{\partial h_{i-1}} \right\| \approx \lambda^d i=k+1∏t∂hi−1∂hi ≈λd
这是一个指数衰减!
3.2.3 数值模拟:看看梯度到底衰减多快?
我们做一个简化实验(标量 RNN):
- 隐藏状态维度 = 1(即 h t ∈ R h_t \in \mathbb{R} ht∈R)
- W h h = 0.8 W_{hh} = 0.8 Whh=0.8
- 所有输入 x t = 0 x_t = 0 xt=0,偏置 b h = 0 b_h = 0 bh=0
- 初始 h 0 = 0 h_0 = 0 h0=0
则:
h
t
=
tanh
(
0.8
⋅
h
t
−
1
)
h_t = \tanh(0.8 \cdot h_{t-1})
ht=tanh(0.8⋅ht−1)
由于 h 0 = 0 h_0 = 0 h0=0,所有 h t = 0 h_t = 0 ht=0,于是 tanh ′ ( z t ) = 1 \tanh'(z_t) = 1 tanh′(zt)=1。
此时:
∂
h
t
∂
h
t
−
1
=
0.8
×
1
=
0.8
\frac{\partial h_t}{\partial h_{t-1}} = 0.8 \times 1 = 0.8
∂ht−1∂ht=0.8×1=0.8
那么从
h
10
h_{10}
h10 到
h
1
h_1
h1 的梯度链为:
∂
h
10
∂
h
1
=
(
0.8
)
9
≈
0.134
\frac{\partial h_{10}}{\partial h_1} = (0.8)^9 \approx 0.134
∂h1∂h10=(0.8)9≈0.134
如果 W h h = 0.5 W_{hh} = 0.5 Whh=0.5,则 ( 0.5 ) 9 ≈ 0.002 (0.5)^9 \approx 0.002 (0.5)9≈0.002 —— 几乎为零!
而在真实训练中:
- tanh ′ ( z ) \tanh'(z) tanh′(z) 通常 ≤ 0.5
- W h h W_{hh} Whh 经初始化和训练后,其有效增益可能更低
因此,超过 10~20 步后,梯度基本消失。这意味着:
“我昨天在银行取了钱” 中,“我” 对 “钱” 的影响,在反向传播中几乎为零。
3.2.4 可视化类比:梯度是一条穿越峡谷的河流
想象梯度是一条从山顶(当前损失)流向源头(初始输入)的河流:
- 每经过一个时间步,就穿过一道狭窄的峡谷(激活函数 + 权重);
- 峡谷两侧不断吸水( tanh ′ < 1 \tanh' < 1 tanh′<1),河道越来越窄( W h h W_{hh} Whh 范数 < 1);
- 流经 10 道峡谷后,河水只剩涓涓细流;
- 流经 20 道后,河床干涸。
此时,源头的土地(早期参数)得不到灌溉(梯度),无法生长(更新)。
💡 这就是为什么 RNN 在短序列(如 10 个词以内)表现尚可,但在长文本中束手无策。
3.2.5 副作用:梯度爆炸(Gradient Explosion)
有趣的是,如果 W h h W_{hh} Whh 的谱范数 大于 1,连乘项会指数增长,导致梯度爆炸:
∥ ∂ h t ∂ h k ∥ ≈ ( 1.2 ) d → ∞ as d ↑ \left\| \frac{\partial h_t}{\partial h_k} \right\| \approx (1.2)^d \to \infty \quad \text{as } d \uparrow ∂hk∂ht ≈(1.2)d→∞as d↑
这会导致参数更新剧烈震荡,训练崩溃。
✅ 工程对策:实践中常使用 梯度裁剪(Gradient Clipping),将梯度范数限制在阈值内(如 5.0),防止爆炸。但它无法解决消失问题。
3.2.6 历史演进:如何绕过这座大山?
梯度消失问题直接推动了深度学习两大里程碑:
-
LSTM(1997)
引入细胞状态 c t c_t ct 和门控机制,让信息可以通过加法路径( c t = f t ⋅ c t − 1 + i t ⋅ c ~ t c_t = f_t \cdot c_{t-1} + i_t \cdot \tilde{c}_t ct=ft⋅ct−1+it⋅c~t)无损传递。
→ 梯度可直接流过 c t c_t ct,避免连乘衰减。 -
Transformer(2017)
彻底抛弃循环结构,用自注意力直接计算任意两个 token 的关系。
→ 依赖建模与距离无关,且完全并行,成为大模型标准架构。
📉 现实结论:尽管 LSTM 在中小规模任务中仍有价值,但在当前 LLM(大语言模型)时代,RNN/LSTM 已被 Transformer 全面取代。原因正是:注意力机制天然规避了梯度消失。
3.2.7 如何记忆“梯度消失”的核心逻辑?
记住这个链条:
循环连接 → 隐藏状态依赖链 → 反向传播需连乘 → 激活函数导数 < 1 → 指数衰减 → 早期时间步学不到东西
或用一句话:
“RNN 的记忆像雪球下山:滚得越远,化得越快;还没到山脚,早已无影无踪。”
3.2.8 终极总结
梯度消失不是实现 bug,而是 RNN 架构在连续非线性变换 + 参数共享下的必然数学结果。它揭示了一个深刻事实:
顺序处理 ≠ 长期记忆。
正因如此,深度学习才从“循环思维”走向“全局注意力”,开启了大模型时代。理解梯度消失,不仅是掌握 RNN 的关键,更是理解现代 AI 架构演进的起点。
4. 解决方案:LSTM(长短期记忆网络)
4.1 LSTM 核心思想:三个门 + 细胞状态
- 遗忘门(Forget Gate):决定丢弃哪些旧信息
- 输入门(Input Gate):决定更新哪些新信息
- 输出门(Output Gate):决定输出什么
4.2 数学公式(逐行详解)
设输入为 x t x_t xt,上一时刻隐藏状态为 h t − 1 h_{t-1} ht−1,细胞状态为 c t − 1 c_{t-1} ct−1。
# 将上一时刻隐藏状态 h_{t-1} 和当前输入 x_t 拼接,作为门控网络的共同输入
# 形状:[batch_size, hidden_dim + input_dim]
combined = torch.cat((h_prev, x_t), dim=1)
# 遗忘门 f_t:sigmoid 输出 0~1,0 表示“完全遗忘”,1 表示“完全保留”
f_t = torch.sigmoid(W_f @ combined + b_f) # [batch, hidden]
# 输入门 i_t:控制新候选值 c_tilde 的写入程度
i_t = torch.sigmoid(W_i @ combined + b_i) # [batch, hidden]
# 候选细胞状态 c_tilde:tanh 保证数值稳定在 [-1,1]
c_tilde = torch.tanh(W_c @ combined + b_c) # [batch, hidden]
# 更新细胞状态:旧状态 * 遗忘比例 + 新信息 * 写入比例
c_t = f_t * c_prev + i_t * c_tilde # [batch, hidden]
# 输出门 o_t:决定最终输出多少细胞状态的信息
o_t = torch.sigmoid(W_o @ combined + b_o) # [batch, hidden]
# 隐藏状态 h_t:输出门控制下的细胞状态(经过 tanh 压缩)
h_t = o_t * torch.tanh(c_t) # [batch, hidden]
✅ 为什么能缓解梯度消失?
细胞状态 c t c_t ct 的更新是加法操作( c t = f t ⋅ c t − 1 + ⋯ c_t = f_t \cdot c_{t-1} + \cdots ct=ft⋅ct−1+⋯),梯度可以直接流过而不经过激活函数连乘,因此长期依赖得以保留。
5. 完整代码示例:用 PyTorch 实现 RNN 与 LSTM 文本分类(逐行注释)
# 导入 PyTorch 核心模块:张量计算、神经网络、优化器
import torch
import torch.nn as nn
import torch.optim as optim
# 导入 TorchText 数据集和工具:用于自然语言处理
from torchtext.datasets import IMDB # IMDB 影评数据集(正面/负面)
from torchtext.data.utils import get_tokenizer # 获取英文分词器
from torchtext.vocab import build_vocab_from_iterator # 从文本构建词汇表
from torch.utils.data import DataLoader # 批量加载数据
from collections import Counter # 统计词频(备用)
# 自动选择 GPU(若可用)或 CPU 运行,加速训练
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# -----------------------------
# 1. 数据预处理
# -----------------------------
# 使用基础英文分词器(按空格和标点分割)
tokenizer = get_tokenizer('basic_english')
# 定义一个生成器函数:每次从数据集中 yield 一个样本的分词结果
def yield_tokens(data_iter):
for _, text in data_iter: # 忽略标签,只取文本
yield tokenizer(text) # 返回分词后的词列表
# 加载 IMDB 训练集和测试集(返回迭代器)
train_iter, test_iter = IMDB(split=('train', 'test'))
# 从训练集构建词汇表:只保留频率最高的 10000 个词
vocab = build_vocab_from_iterator(
yield_tokens(train_iter),
min_freq=1, # 至少出现 1 次
max_tokens=10000 # 最多保留 10000 个词
)
# 设置未知词 <unk> 的索引(未登录词统一映射到此)
vocab.set_default_index(vocab['<unk>'])
# 定义批处理函数:将一批 (label, text) 转换为张量
def collate_batch(batch):
label_list, text_list = [], []
for _label, _text in batch:
# 将标签 "pos"/"neg" 转为 1/0
label_list.append(1 if _label == 'pos' else 0)
# 将文本分词后查词汇表,转为整数索引张量
processed_text = torch.tensor(vocab(tokenizer(_text)), dtype=torch.long)
text_list.append(processed_text)
# 对变长序列进行填充(padding),使同一批次长度一致
padded_texts = nn.utils.rnn.pad_sequence(
text_list,
batch_first=True, # 形状变为 [batch_size, seq_len]
padding_value=0 # 用 0 填充(对应词汇表中的 <pad>,但此处 vocab 未显式定义,0 默认为填充)
)
# 截断过长序列(超过 64 个词的部分丢弃)
padded_texts = padded_texts[:, :64]
# 若序列不足 64,右侧补 0
if padded_texts.size(1) < 64:
padded_texts = nn.functional.pad(
padded_texts,
(0, 64 - padded_texts.size(1)) # (左, 右) 填充
)
# 返回标签张量和文本张量
return torch.tensor(label_list, dtype=torch.long), padded_texts
# 将 IMDB 数据集转换为可迭代的列表(便于 DataLoader 处理)
train_data = list(IMDB(split='train'))
test_data = list(IMDB(split='test'))
# 创建数据加载器:每次返回一个 batch(32 条样本)
train_dataloader = DataLoader(
train_data,
batch_size=32,
shuffle=True, # 训练时打乱顺序
collate_fn=collate_batch
)
test_dataloader = DataLoader(
test_data,
batch_size=32,
shuffle=False, # 测试时不打乱
collate_fn=collate_batch
)
# -----------------------------
# 2. 定义 RNN 分类模型
# -----------------------------
class RNNClassifier(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_dim, output_dim, model_type='RNN'):
super().__init__()
# 词嵌入层:将词索引映射为稠密向量(维度 embed_dim)
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
# 根据 model_type 选择 RNN 或 LSTM 层
if model_type == 'RNN':
self.rnn = nn.RNN(embed_dim, hidden_dim, batch_first=True)
elif model_type == 'LSTM':
self.rnn = nn.LSTM(embed_dim, hidden_dim, batch_first=True)
# 全连接层:将隐藏状态映射到类别 logits(2 类:正面/负面)
self.fc = nn.Linear(hidden_dim, output_dim)
self.model_type = model_type
def forward(self, x):
# x: [batch_size, seq_len],每个元素是词索引
embedded = self.embedding(x) # [batch, seq_len, embed_dim]
if self.model_type == 'RNN':
# output: [batch, seq, hidden], hidden: [num_layers=1, batch, hidden]
output, hidden = self.rnn(embedded)
elif self.model_type == 'LSTM':
# LSTM 返回 (output, (hidden, cell))
output, (hidden, cell) = self.rnn(embedded)
# 取最后一个时间步的隐藏状态作为整个句子的表示
last_hidden = hidden.squeeze(0) # [batch, hidden]
# 映射到分类 logits
logits = self.fc(last_hidden) # [batch, 2]
return logits
# -----------------------------
# 3. 训练与评估函数
# -----------------------------
def train_model(model_type='LSTM'):
# 初始化模型并移至设备(GPU/CPU)
model = RNNClassifier(
vocab_size=len(vocab),
embed_dim=128,
hidden_dim=256,
output_dim=2,
model_type=model_type
).to(device)
# 交叉熵损失函数(内部包含 softmax)
criterion = nn.CrossEntropyLoss()
# Adam 优化器,学习率 0.001
optimizer = optim.Adam(model.parameters(), lr=0.001)
# 训练 5 个 epoch
for epoch in range(5):
model.train() # 开启训练模式(启用 dropout 等)
total_loss = 0
for labels, texts in train_dataloader:
labels, texts = labels.to(device), texts.to(device)
optimizer.zero_grad() # 清空上一步梯度
outputs = model(texts) # 前向传播
loss = criterion(outputs, labels) # 计算损失
loss.backward() # 反向传播
optimizer.step() # 更新参数
total_loss += loss.item() # 累加损失
print(f'Epoch {epoch+1}, Avg Loss: {total_loss/len(train_dataloader):.4f}')
# 测试阶段
model.eval() # 关闭训练模式(禁用 dropout)
correct = 0
total = 0
with torch.no_grad(): # 禁用梯度计算,节省内存
for labels, texts in test_dataloader:
labels, texts = labels.to(device), texts.to(device)
outputs = model(texts)
pred = outputs.argmax(dim=1) # 取概率最大的类别
correct += (pred == labels).sum().item()
total += labels.size(0)
print(f'{model_type} Test Accuracy: {100 * correct / total:.2f}%')
# 运行对比实验
print("Training RNN...")
train_model('RNN')
print("\nTraining LSTM...")
train_model('LSTM')
💡 运行结果预期:
- RNN 准确率约 75%~80%
- LSTM 准确率可达 85%+(因更好捕捉长距离依赖)
6. CNN vs RNN:应用场景对比表(大白话版)
| 特性 | CNN(卷积神经网络) | RNN(循环神经网络) |
|---|---|---|
| 擅长处理的数据类型 | 网格结构数据(如图像、视频帧) | 序列数据(如文本、语音、时间序列) |
| 核心能力 | 提取局部空间特征(边缘、纹理、物体部件) | 捕捉时间/顺序依赖(前文影响后文) |
| 典型例子 | - 识别猫狗图片 - 医学影像分割 - 人脸识别 | - 机器翻译(英→中) - 语音转文字 - 股票价格预测 |
| 能不能记住过去? | ❌ 每个区域独立处理,无记忆 | ✅ 隐藏状态携带历史信息 |
| 并行性 | ✅ 卷积操作可高度并行 | ❌ 必须按时间步顺序计算 |
| 现代地位 | 仍是 CV 领域基石(尤其结合 ViT) | 在大模型中基本被取代 |
7. RNN 在大模型时代的地位
明确结论:RNN(包括 LSTM)在当前主流大模型中已基本过时,不再使用。
-
原因 1:并行性差
Transformer 的自注意力机制允许所有 token 同时计算,极大提升训练速度;而 RNN 必须串行处理,无法利用现代 GPU 的并行能力。 -
原因 2:长程依赖仍有瓶颈
尽管 LSTM 缓解了梯度消失,但在数千 token 的上下文中,信息传递依然衰减严重;而 Transformer 通过注意力直接建模任意两个 token 的关系。 -
现实情况:
当前所有主流大模型(如 GPT、LLaMA、BERT、Claude)均基于 Transformer 架构,没有使用 RNN/LSTM。
RNN 仅在资源受限设备(如嵌入式系统)、短序列任务或教学场景中仍有价值。
8. 结语与预告
本文系统介绍了 RNN 的动机、结构、数学原理、梯度问题及 LSTM 解决方案,并附带逐行注释的完整 PyTorch 代码实现。虽然 RNN 已非最前沿,但其“循环”思想深刻影响了序列建模范式。
🔗 延伸阅读:关于 CNN 的详细解析,请参见笔者此前文章:
【人工智能】人工智能发展历程全景解析:从图灵测试到大模型时代(含CNN、Q-Learning深度实践)
✅ 下一篇预告:
【人工智能】【深度学习】 ② GAN核心算法介绍:生成器与判别器的博弈艺术
→ 将深入剖析 GAN 的对抗训练机制、损失函数推导、WGAN 改进,以及人脸生成实战代码!
敬请期待!
更多推荐
所有评论(0)