📖目录

前言

💡 提示:本系列按《全解深度学习——九大核心算法》目录顺序展开。由于笔者此前已在 这篇博客 中系统讲解了 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+Whhht1+bh)=Wyhht+by

乍看之下,这只是带了“上一时刻隐藏状态”的普通神经网络。但它的设计背后有清晰的推导逻辑工程直觉。下面我们一步步拆解。


2.2.1 公式是怎么来的?——从全连接网络到循环结构

  1. 起点:普通神经网络
    对于静态输入 x x x,单层网络输出为:
    h = tanh ⁡ ( W x + b ) h = \tanh(W x + b) h=tanh(Wx+b)
    每个样本独立处理,无历史依赖。

  2. 需求:序列需要“记忆”
    在处理“我吃了一个___”时,模型必须知道前面出现了“吃”,才能预测“苹果”。
    → 因此,当前状态 h t h_t ht 应同时依赖 当前输入 x t x_t xt历史状态 h t − 1 h_{t-1} ht1

  3. 自然扩展:拼接输入
    最直接的做法是把 x t x_t xt h t − 1 h_{t-1} ht1 拼成一个大向量:
    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[xtht1]+bh)

  4. 拆解权重矩阵
    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} ht1

    于是:
    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[xtht1]=Whxxt+Whhht1

  5. 最终形式
    代入后即得 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+Whhht1+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 如何快速记忆这个公式?

记住这个口诀:

“新输入 + 旧记忆 → 融合 → 激活 → 新记忆”

对应步骤:

  1. 新输入 W h x x t W_{hx} x_t Whxxt —— 当前时刻看到的内容
  2. 旧记忆 W h h h t − 1 W_{hh} h_{t-1} Whhht1 —— 上一时刻留下的印象
  3. 融合 + 偏置:两者相加再加 b h b_h bh
  4. 激活 tanh ⁡ ( ⋅ ) \tanh(\cdot) tanh() 压缩到 [-1, 1],引入非线性
  5. 新记忆:结果就是 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} Whhht1保留并转换历史记忆(“循环”的来源)
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=1TCrossEntropy(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 L0

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}} WhhL 的结构。

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}} WhhL=t=1TWhhLt

其中 L t \mathcal{L}_t Lt 是 t 时刻的损失(如 CrossEntropy)。

Step 2:单个时间步的梯度分解

考虑 ∂ L t ∂ W h h \frac{\partial \mathcal{L}_t}{\partial W_{hh}} WhhLt。因为 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}} WhhLt=k=1thtLt(i=k+1thi1hi)Whhhk

这个公式有三层含义:

部分含义
∂ L t ∂ h t \frac{\partial \mathcal{L}_t}{\partial h_t} htLt当前损失对当前记忆的敏感度(误差信号)
∏ i = k + 1 t ∂ h i ∂ h i − 1 \prod_{i=k+1}^{t} \frac{\partial h_i}{\partial h_{i-1}} i=k+1thi1hi从 t 到 k 的“记忆传递链”(核心问题所在)
∂ h k ∂ W h h \frac{\partial h_k}{\partial W_{hh}} Whhhk当前记忆对参数的直接依赖
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} hi1hi=hi1tanh(Whxxi+Whhhi1+bh)=tanh(zi) diag(1tanh2(zi))Whh

  • tanh ⁡ ′ ( z ) = 1 − tanh ⁡ 2 ( z ) ∈ ( 0 , 1 ] \tanh'(z) = 1 - \tanh^2(z) \in (0, 1] tanh(z)=1tanh2(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}} hi1hi 的范数通常小于 1

假设平均每步的缩放因子为 λ < 1 \lambda < 1 λ<1,那么跨越 d = t − k d = t - k d=tk 步的梯度大小约为:

∥ ∏ 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+1thi1hi λd

这是一个指数衰减


3.2.3 数值模拟:看看梯度到底衰减多快?

我们做一个简化实验(标量 RNN):

  • 隐藏状态维度 = 1(即 h t ∈ R h_t \in \mathbb{R} htR
  • 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.8ht1)

由于 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 ht1ht=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 h1h10=(0.8)90.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)90.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 hkht (1.2)das d

这会导致参数更新剧烈震荡,训练崩溃。

工程对策:实践中常使用 梯度裁剪(Gradient Clipping),将梯度范数限制在阈值内(如 5.0),防止爆炸。但它无法解决消失问题


3.2.6 历史演进:如何绕过这座大山?

梯度消失问题直接推动了深度学习两大里程碑:

  1. 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=ftct1+itc~t)无损传递。
    → 梯度可直接流过 c t c_t ct,避免连乘衰减。

  2. 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} ht1,细胞状态为 c t − 1 c_{t-1} ct1

# 将上一时刻隐藏状态 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=ftct1+),梯度可以直接流过而不经过激活函数连乘,因此长期依赖得以保留。


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 改进,以及人脸生成实战代码!

敬请期待!

更多推荐