LSTM & GRU —— 门控循环神经网络
一、引言:RNN的“记忆困境”与门控革命的到来
在上一篇文章中,我们推导了BPTT的数学公式,并看到了核心问题所在:
∂hk∂ht=i=k+1∏t(Whh⊤⋅diag(ϕ′(hi−1)))这个矩阵连乘要么让梯度指数级衰减到零(梯度消失),要么指数级膨胀到无穷(梯度爆炸)。RNN的“记忆”在数学上被这个连乘公式锁死了——距离越远的信息,梯度贡献越小,最终完全消失。
门控循环神经网络(Gated Recurrent Neural Network)的提出,正是为了更好地捕捉时间序列中时间步距离较大的依赖关系。它通过可以学习的门来控制信息的流动,让网络自行决定哪些信息该记住、哪些该遗忘、哪些该输出。
本文将深入剖析两种最经典的门控架构——LSTM(Long Short-Term Memory) 和 GRU(Gated Recurrent Unit) ,从数学公式到设计哲学,从梯度传播到工程实践。
二、LSTM:从“单轨”到“双轨”的架构革命
2.1 核心思想:给记忆修一条“高速公路”
常规RNN的问题是它内部状态的更新方式是“粗暴”的——每一步的新信息都会与旧信息无差别地混合。LSTM的设计哲学是赋予网络自行决定信息取舍的能力。
与RNN只有一个隐藏状态 ht 在时间步之间传递不同,LSTM引入了两个独立的状态向量在时间轴上并行传递:
-
细胞状态(Cell State, ct) :这是LSTM的核心,原始论文中称之为 “恒定误差旋转木马”(Constant Error Carousel, CEC) 。可以把它想象成一条“信息高速公路”或“传送带”,负责在整个序列中传递长期记忆。
-
隐藏状态(Hidden State, ht) :与RNN中的隐藏状态类似,代表了当前时间步的短期记忆和最终输出。
💡 关键洞察:在普通RNN中,信息在时间步之间传递必须经过矩阵乘法(Whhht−1)。而在LSTM中,细胞状态的传递路径是按元素的加法和乘法(ct=ft⊙ct−1+it⊙c~t),没有额外的矩阵连乘,信息可以直接在这条传送带上流动。
2.2 三个门:信息流动的“智能开关”
LSTM中引入了3个门,即输入门(input gate)、遗忘门(forget gate)和输出门(output gate) 。
LSTM中的“门”是一种让信息选择性通过的结构,设计灵感来源于数字电路中的逻辑门。它的实现非常简单:一个以Sigmoid为激活函数的全连接层,输入通常是当前时间步的输入 xt 和上一个时间步的隐藏状态 ht−1 的拼接向量。
Sigmoid函数将元素值映射到 (0, 1) 区间内:
- 输出接近 1 → “允许”对应维度的信息完全通过
- 输出接近 0 → “阻止”对应维度的信息通过,即“遗忘”或“忽略”它
三个门的分工如下:
| 门 | 符号 | 作用 |
|---|---|---|
| 遗忘门 | ft | 决定是否让上一时刻学到的信息通过或部分通过 |
| 输入门 | it | 计算出候选值,决定哪些新信息写入细胞状态 |
| 输出门 | ot | 决定哪些信息输出到隐藏状态 |
2.3 LSTM的完整数学公式
假设隐藏单元个数为 h,给定时间步 t 的小批量输入 Xt∈Rn×d 和上一时间步隐藏状态 Ht−1∈Rn×h。
Step 1:三个门的计算
It=σ(XtWxi+Ht−1Whi+bi)Ft=σ(XtWxf+Ht−1Whf+bf)Ot=σ(XtWxo+Ht−1Who+bo)其中 Wxi,Wxf,Wxo∈Rd×h 和 Whi,Whf,Who∈Rh×h 是权重参数,bi,bf,bo∈R1×h 是偏差参数。
Step 2:候选记忆细胞
C~t=tanh(XtWxc+Ht−1Whc+bc)这里使用值域在 [−1,1] 的 tanh 函数作为激活函数。
Step 3:细胞状态更新(核心!)
Ct=Ft⊙Ct−1+It⊙C~t其中 ⊙ 表示按元素乘法。
这个公式是LSTM的灵魂:
- Ft⊙Ct−1:遗忘旧信息(遗忘门控制保留多少)
- It⊙C~t:写入新信息(输入门控制写入多少)
Step 4:隐藏状态计算
Ht=Ot⊙tanh(Ct)输出门控制从细胞状态中读取多少信息到隐藏状态。
2.4 为什么LSTM能缓解梯度消失?——数学证明
这是理解LSTM最核心的部分。让我们从梯度传播的角度来看。
在普通RNN中,隐藏状态的更新是:
ht=ϕ(Wxhxt+Whhht−1+bh)梯度传播的关键项是:
∂ht−1∂ht=Whh⊤⋅diag(ϕ′(ht−1))这是矩阵乘法——谱范数决定了梯度是指数衰减还是爆炸。
而在LSTM中,细胞状态的更新是:
Ct=Ft⊙Ct−1+It⊙C~t关键差异:∂Ct/∂Ct−1 是什么?
∂Ct−1∂Ct=diag(Ft)这是一个对角矩阵,而不是满矩阵!
这意味着:
- 没有矩阵乘法:梯度在细胞状态路径上的传播是逐元素的,不存在谱范数导致的全局指数衰减
- 遗忘门可以学习:如果网络需要长期记忆,遗忘门 Ft 可以学习为接近1的值,让梯度几乎无损地通过
- 加法路径:信息可以直接在细胞状态的“传送带”上流动,仅经过按元素的加权与相加
🔑 核心结论:LSTM通过将信息传递路径从矩阵乘法改为按元素加法,从根本上改变了梯度传播的数学性质。梯度不再需要经过一连串的矩阵相乘,而是可以通过细胞状态的“高速公路”直接流回早期时间步。
从另一个角度看,门控机制也是为了解决权重冲突问题——输入门保护细胞状态不受无关输入的干扰,输出门则保护其他单元不受当前细胞状态中无关记忆的干扰。
三、GRU:LSTM的“精简版”
3.1 为什么需要GRU?
LSTM成功解决了长时依赖问题,但代价是三个门 + 一个细胞状态,结构复杂、参数众多。GRU(门控循环单元)由Cho等人于2014年提出,是LSTM的一个更简单的变体。
GRU的设计目标很明确:在保持LSTM性能的同时,减少参数数量和计算复杂度。
3.2 GRU的两大简化
简化一:三门→两门
GRU将LSTM中的三个门(遗忘门、输入门、输出门)合并为两个门——重置门(Reset Gate)和更新门(Update Gate) 。
具体来说,GRU把LSTM的输入门和遗忘门组合在一起,少了一个门。更新门 z 的角色相当于LSTM里的遗忘门,而 1−z 相当于LSTM中的输入门。
简化二:双状态→单状态
LSTM有两个状态向量在时间轴上传递——细胞状态 ct(长期记忆)和隐藏状态 ht(短期记忆)。
GRU将细胞状态和隐藏状态合并,只传递一个隐藏状态 ht。在GRU里,ht 的角色比较像LSTM中的 ct,可以保留得比较久。
💡 设计哲学:GRU中遗忘门和输入门是联动的——如果有新的信息进来,才会忘掉之前的信息;如果没有新信息进来,就不会忘记信息。这个逻辑比LSTM的独立三门更简洁。
3.3 GRU的完整数学公式
Step 1:重置门和更新门
Rt=σ(XtWxr+Ht−1Whr+br)Zt=σ(XtWxz+Ht−1Whz+bz)其中 Wxr,Wxz∈Rd×h 和 Whr,Whz∈Rh×h 是权重参数。
Step 2:候选隐藏状态
H~t=tanh(XtWxh+(Rt⊙Ht−1)Whh+bh)重置门 Rt 控制着过去信息的丢弃程度:
- 当重置门的值接近 0 时,意味着对应的隐藏状态元素将被重置为0,从而丢弃上一时间步的历史信息
- 当接近 1 时,表示保留上一时间步的隐藏状态
Step 3:最终隐藏状态
Ht=Zt⊙Ht−1+(1−Zt)⊙H~t最终的隐藏状态是候选隐藏状态和前一隐藏状态的加权组合,权重由更新门控制:
- 当更新门接近 1 时,新状态几乎完全继承过去状态
- 当接近 0 时,新状态主要由候选状态决定
3.4 GRU vs LSTM:参数量的定量对比
LSTM的参数由3个门 + 1个候选细胞状态组成,每个都需要独立的权重矩阵:
LSTM参数量=4×(d×h+h×h+h)GRU只有2个门 + 1个候选隐藏状态:
GRU参数量=3×(d×h+h×h+h)GRU的参数量约为LSTM的 43 。
📌 实际表现:在很多时候,人们更愿意使用GRU来替换LSTM,因为GRU比LSTM少一个门,参数更少,相对容易训练且可以防止过拟合(尤其是在训练样本少的时候)。而且,GRU的性能和LSTM几乎一样。
不过需要注意的是,虽然GRU参数更少,但由于重置门的计算中并行性较低,某些情况下LSTM的执行时间反而更短。
四、双向LSTM与深层堆叠
4.1 双向LSTM:同时看到“过去”和“未来”
标准的LSTM是单向的——信息只能从过去流向未来。但在很多任务中(如机器翻译、文本分类),未来的上下文同样重要。
双向LSTM(Bidirectional LSTM) 使用两个独立的LSTM层:
- 前向LSTM:按时间正序处理序列(从 t=1 到 t=T)
- 后向LSTM:按时间逆序处理序列(从 t=T 到 t=1)
然后将两个方向的隐藏状态拼接起来作为最终的表示。
💡 直观理解:就像我们在做阅读理解时,不仅看前面的词,也会看后面的词来确定当前词的含义。双向LSTM让模型同时拥有了“回顾过去”和“展望未来”的能力。
典型应用场景:
- 自然语言处理:情感分析、命名实体识别、机器翻译
- 语音识别:利用前后音素信息提高识别准确率
- 蛋白质结构预测:利用序列上下文信息
4.2 深层堆叠LSTM:增加网络的“深度”
堆叠LSTM(Stacked LSTM / Deep LSTM) 将多个LSTM层垂直堆叠在一起:
- 第1层LSTM接收原始输入序列
- 第2层LSTM接收第1层的输出作为输入
- 依此类推…
每一层LSTM都在不同的时间抽象层次上学习特征:
- 底层:捕捉局部的、短期的模式
- 高层:捕捉全局的、长期的依赖关系
📌 实践建议:堆叠的层数足够大时,多层RNN的效果可能会比单层好。但堆叠层数增加会带来更高的计算负荷,且需要更多数据来避免过拟合。
4.3 PyTorch实现
1import torch2import torch.nn as nn3
4# ============ LSTM ============5# 单层LSTM6lstm = nn.LSTM(input_size=10, hidden_size=20, num_layers=1, batch_first=True)7
8# 双层堆叠LSTM9stacked_lstm = nn.LSTM(input_size=10, hidden_size=20, num_layers=2, batch_first=True)10
11# 双向LSTM12bidirectional_lstm = nn.LSTM(13 input_size=10,14 hidden_size=20,15 num_layers=2,16 bidirectional=True, # 开启双向17 batch_first=True18)19
20# ============ GRU ============21# 单层GRU22gru = nn.GRU(input_size=10, hidden_size=20, num_layers=1, batch_first=True)23
24# 双层堆叠GRU25stacked_gru = nn.GRU(input_size=10, hidden_size=20, num_layers=2, batch_first=True)26
27# 双向GRU28bidirectional_gru = nn.GRU(29 input_size=10,30 hidden_size=20,31 num_layers=2,32 bidirectional=True,33 batch_first=True34)35
36# ============ 前向传播示例 ============37batch_size, seq_len, input_size = 32, 50, 1038x = torch.randn(batch_size, seq_len, input_size)39
40# 双向双层LSTM41output, (h_n, c_n) = bidirectional_lstm(x)42# output: (batch_size, seq_len, hidden_size * 2) # 双向 → 2倍43# h_n: (num_layers * 2, batch_size, hidden_size)44# c_n: (num_layers * 2, batch_size, hidden_size)45
46print(f"输出形状: {output.shape}") # (32, 50, 40)47print(f"最终隐藏状态形状: {h_n.shape}") # (4, 32, 20)五、总结
| 特性 | 标准RNN | LSTM | GRU |
|---|---|---|---|
| 状态数量 | 1个(ht) | 2个(ct,ht) | 1个(ht) |
| 门控数量 | 0 | 3(输入/遗忘/输出) | 2(重置/更新) |
| 参数量 | 基准 | ~4倍于RNN | ~3倍于RNN(LSTM的3/4) |
| 梯度消失 | 严重 | ✅ 极大缓解 | ✅ 极大缓解 |
| 长时依赖 | 差 | ✅ 优秀 | ✅ 优秀 |
| 计算效率 | 高 | 低 | 中等 |
| 适用场景 | 短序列 | 长序列、复杂任务 | 长序列、资源受限 |
选择建议:
- 数据量大、任务复杂、计算资源充足 → 选择 LSTM
- 数据量适中、需要快速迭代、资源有限 → 选择 GRU
- 需要利用未来上下文信息 → 使用 双向LSTM/GRU
- 需要建模多层次的时间抽象 → 使用 堆叠LSTM/GRU
LSTM和GRU的门控机制,是深度学习历史上最重要的架构创新之一。它们不仅让循环神经网络真正具备了处理长序列的能力,其设计哲学——用可学习的“门”来控制信息流动——也深刻地影响了后来的Transformer、扩散模型等现代架构。
延伸阅读:
- Long Short-Term Memory(Hochreiter & Schmidhuber, 1997)
- Learning Phrase Representations using RNN Encoder-Decoder for Statistical Machine Translation(GRU原始论文,Cho et al., 2014)
- Dive into Deep Learning - LSTM
- Dive into Deep Learning - GRU
Some information may be outdated