
LSTM & GRU —— 门控循环神经网络
系统讲解LSTM与GRU的门控机制设计哲学与数学原理。从LSTM的“双轨”架构(细胞状态+隐藏状态)与三门结构(遗忘门、输入门、输出门)出发,推导细胞状态更新公式如何通过按元素加法路径缓解梯度消失,并深入剖析GRU的两大简化及其参数量优势。附带讨论双向LSTM与深层堆叠LSTM的应用场景。
阅读文章ZHY's Blog
A UNIVERSE OF IDEAS · BY ZHANG HAOYI
让好奇心 点亮知识宇宙
在代码、模型与思想之间自由漫游。这里持续记录人工智能、机器学习、软件工程与成长实践,让每次阅读都成为一次新的发现。
ARTICLE NOTE
在前两篇文章中,我们完整走过了空间维度上深度网络的优化之路:CNN通过局部连接与权重共享建模空间结构,ResNet借助恒等映射的直通路径让梯度无损穿越数百层。然而,当我们从图像转向文本、语音等序列数据时,一个新的维度悄然浮现——时间。今天的单词依赖于昨天的上下文,当前的股价受制于前序走势——这种“记忆”需求,要求网络不仅要在层间传播信息,更要在时间步之间传递状态。
本篇博客正是应对这一“时序建模”挑战的核心章节。我们从RNN的循环结构与共享参数出发,揭示隐藏状态 ht=ϕ(Wxhxt+Whhht−1+bh) 如何通过 Whh 在时间维度上传递“记忆”。随后,我们将深入BPTT(随时间反向传播) 的数学推导——将RNN按时间步展开后,误差不仅要在层间传播,更要在 T 个时间步上反向流动,最终导出包含双重求和的梯度表达式。
这一推导将直面RNN最致命的困境:梯度消失与梯度爆炸。核心原因在于 ∂ht/∂hk 是一个从 k+1 到 t 的矩阵连乘 —— 当谱范数小于1时,梯度指数级衰减,网络“记不住”10步前的信息;当大于1时,梯度指数级爆炸,训练瞬间崩溃。我们将介绍梯度裁剪作为应对梯度爆炸的工程解法,并前瞻LSTM如何通过门控机制(加法路径替代乘法路径)从根本上缓解梯度消失。
值得注意的是,本篇对“时序梯度传播”的深刻理解,将直接服务于后续LSTM与GRU —— 当您理解了BPTT为何失效,才能真正读懂门控机制的设计初衷。现在,请带着“误差如何跨时间流动”的疑问进入正文——理解了BPTT,您就掌握了序列模型训练的核心密码。
一个标准的循环神经网络在时间步 t 的隐藏状态 ht 和输出 ot 可以表示为:
ht=ϕ(Wxhxt+Whhht−1+bh)ot=ψ(Whoht+bo)其中:
💡 关键洞察:Whh 是RNN的“记忆矩阵”——它决定了当前隐藏状态如何继承上一时刻的信息。正是这个矩阵的循环使用,让RNN拥有了处理变长序列的能力。
BPTT的核心思想是:将循环网络按时间步“展开”成一个前馈网络。
对于一个长度为 T 的序列,展开后的计算图为:
1x₁ → h₁ → o₁ → L₁2↑ ↑3Wxh Whh4 ↓5x₂ → h₂ → o₂ → L₂6↑ ↑7Wxh Whh8 ↓9...10x_T → h_T → o_T → L_T注意:所有时间步共享同一组参数(Wxh,Whh,Who)。这是RNN与前馈网络的关键区别——也是BPTT梯度推导的难点所在。
对于序列数据,总损失是所有时间步损失之和:
L=t=1∑TLt(ot,yt)其中 yt 是时间步 t 的真实标签。
在BPTT中,损失 L 对参数 Whh 的梯度需要沿着两条路径传播:
对于参数 Whh,总梯度为所有时间步贡献之和:
∂Whh∂L=t=1∑T∂Whh∂Lt根据链式法则,∂Lt/∂Whh 可以展开为:
∂Whh∂Lt=k=1∑t∂ht∂Lt⋅∂hk∂ht⋅∂Whh∂hk这个公式的含义是:时间步 t 的损失,受到从时间步 k 到 t 的所有隐藏状态的影响。
其中最关键的是 ∂ht/∂hk,它表示了隐藏状态在时间上的传播。
根据隐藏状态的更新公式 hi=ϕ(Wxhxi+Whhhi−1+bh),我们有:
∂hi−1∂hi=Whh⊤⋅diag(ϕ′(hi−1))其中 diag(ϕ′(hi−1)) 是对角矩阵,对角线上是激活函数导数在各隐藏神经元上的值。
因此,从时间步 k 到 t 的梯度传播为:
∂hk∂ht=i=k+1∏t∂hi−1∂hi=i=k+1∏t(Whh⊤⋅diag(ϕ′(hi−1)))这就是BPTT的核心公式——一个从 k+1 到 t 的矩阵连乘。
将上述结果代入,得到 ∂Lt/∂Whh 的完整表达式:
∂Whh∂Lt=k=1∑t∂ht∂Lt⋅(i=k+1∏tWhh⊤diag(ϕ′(hi−1)))⋅∂Whh∂hk最终,总梯度为:
∂Whh∂L=t=1∑Tk=1∑t∂ht∂Lt⋅(i=k+1∏tWhh⊤diag(ϕ′(hi−1)))⋅∂Whh∂hk🔑 核心公式的直观理解:这个双重求和告诉我们——每一个时间步的损失,都受到所有历史时间步的影响,而影响的强度由一连串矩阵乘积决定。距离越远(t−k 越大),连乘的项越多,梯度就越容易出问题。
观察 ∂ht/∂hk 的表达式:
∂hk∂ht=i=k+1∏t(Whh⊤⋅diag(ϕ′(hi−1)))这个连乘的行为取决于矩阵 Whh⊤⋅diag(ϕ′(hi−1)) 的谱范数(即最大奇异值):
对于 tanh 激活函数,ϕ′(h)=1−tanh2(h)∈(0,1],其最大值恰好为1。这意味着:
∥Whh⊤⋅diag(ϕ′(h))∥≤∥Whh∥⋅max(ϕ′)=∥Whh∥当 ∥Whh∥<1 时:
∂hk∂ht≈(γ⋅∥Whh∥)t−k其中 γ=max(ϕ′)≤1。当 t−k 较大时(即依赖关系较长),梯度指数级趋近于0。
直观表现:
📌 实例:在一个100个单词的句子中,第1个单词对第100个单词的影响,在BPTT中需要经过99次矩阵连乘。如果每次乘法的谱范数为0.9,那么 0.999≈3×10−5——梯度几乎完全消失了。
当 ∥Whh∥>1 时,梯度会指数级增长。
梯度爆炸虽然不如梯度消失常见(因为权重初始化通常会控制 ∥Whh∥ 在1附近),但一旦发生,后果极其严重:
NaN(无穷大)在实际训练中,梯度爆炸更容易被检测到(梯度值突然变得极大),而梯度消失更隐蔽(训练缓慢但不崩溃)。
梯度爆炸的解决方案相对直接——在反向传播过程中限制梯度的大小。
梯度裁剪(Gradient Clipping)的核心思想是:如果梯度的范数超过某个阈值,就将其缩放到该阈值以内,同时保持梯度方向不变。
最常用的方法是基于全局范数的梯度裁剪。
设所有参数的梯度为 g=[g1,g2,…,gN],其 L2 范数为:
∥g∥2=i=1∑Ngi2设定阈值 θ(通常称为 max_norm),裁剪后的梯度为:
关键特性:
💡 为什么是全局裁剪而非逐层裁剪? 如果逐层独立裁剪,会破坏不同层之间梯度的相对比例,影响优化 dynamics。全局裁剪保持了梯度的方向信息,是更合理的选择。
阈值的选择是梯度裁剪中最关键的超参数:
经验法则:
常见实践:
PyTorch中的实现:
1import torch2import torch.nn as nn3
4# 在反向传播之后、优化器更新之前进行梯度裁剪5loss.backward()6torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)7optimizer.step()clip_grad_norm_ 函数会计算所有参数的梯度总范数,如果超过 max_norm,则统一缩放。
1import numpy as np2
3def clip_gradients(parameters, gradients, max_norm):4 """5 基于全局L2范数的梯度裁剪6
7 Args:8 parameters: 模型参数列表9 gradients: 与parameters对应的梯度列表10 max_norm: 梯度范数的阈值11
12 Returns:13 clipped_gradients: 裁剪后的梯度列表14 """15 # 计算所有梯度的L2范数平方16 total_norm_sq = 0.017 for grad in gradients:18 total_norm_sq += np.sum(grad ** 2)19 total_norm = np.sqrt(total_norm_sq)20
21 # 如果范数超过阈值,进行缩放22 if total_norm > max_norm:23 clip_coef = max_norm / (total_norm + 1e-6)24 clipped_gradients = [grad * clip_coef for grad in gradients]25 else:26 clipped_gradients = gradients27
28 return clipped_gradients, total_norm29
30# 使用示例31# 假设有一个3层RNN的参数32params = [np.random.randn(10, 10) for _ in range(3)] # 模拟参数33grads = [np.random.randn(10, 10) * 5 for _ in range(3)] # 模拟梯度34
35clipped_grads, norm = clip_gradients(params, grads, max_norm=1.0)36print(f"原始梯度范数: {np.sqrt(sum(np.sum(g**2) for g in grads)):.4f}")37print(f"裁剪后梯度范数: {np.sqrt(sum(np.sum(g**2) for g in clipped_grads)):.4f}")| 方式 | 操作 | 适用场景 |
|---|---|---|
按值裁剪 (clip_by_value) | 将每个梯度值截断到 [min, max] 区间 | 需要严格控制每个梯度值范围 |
按范数裁剪 (clip_by_norm / clip_by_global_norm) | 缩放梯度向量使其范数不超过阈值 | RNN/ LSTM 的推荐做法,保持梯度方向 |
对于RNN训练,强烈推荐使用基于全局范数的裁剪,因为它保持了梯度的方向信息。
虽然梯度裁剪能有效应对梯度爆炸,但对于梯度消失,裁剪无能为力——它只能限制过大的梯度,无法“放大”消失的梯度。
真正从根本上缓解梯度消失的是LSTM和GRU等门控RNN架构。
LSTM引入了一个细胞状态(Cell State) ct,其更新方式为:
ct=ft⊙ct−1+it⊙c~t其中 ft 是遗忘门,it 是输入门,⊙ 是逐元素乘法。
关键差异:在普通RNN中,∂ht/∂ht−1 是矩阵乘法(Whh⊤⋅diag(ϕ′));而在LSTM中,细胞状态的传递路径上是逐元素乘法(ft⊙ct−1)。
逐元素乘法的好处是:
GRU(门控循环单元)是LSTM的简化版本,合并了细胞状态和隐藏状态:
ht=(1−zt)⊙ht−1+zt⊙h~t其中 zt 是更新门。与LSTM类似,GRU通过 (1−zt)⊙ht−1 这一加法路径,为梯度提供了直接传播的通道。
1import numpy as np2
3class RNN:4 def __init__(self, input_dim, hidden_dim, output_dim):5 # 参数初始化(使用较小的值,减少梯度爆炸风险)6 self.W_xh = np.random.randn(hidden_dim, input_dim) * 0.017 self.W_hh = np.random.randn(hidden_dim, hidden_dim) * 0.018 self.W_ho = np.random.randn(output_dim, hidden_dim) * 0.019 self.b_h = np.zeros((hidden_dim, 1))10 self.b_o = np.zeros((output_dim, 1))11
12 self.hidden_dim = hidden_dim13
14 def forward(self, x_sequence):15 """16 前向传播:按时间步展开17 x_sequence: list of (input_dim, 1) 向量18 """19 T = len(x_sequence)20 h = np.zeros((T + 1, self.hidden_dim, 1))21 o = np.zeros((T, self.hidden_dim, 1)) # 简化:输出=隐藏状态22
23 for t in range(T):24 h[t] = np.tanh(self.W_xh @ x_sequence[t] +25 self.W_hh @ h[t-1] + self.b_h)26 o[t] = self.W_ho @ h[t] + self.b_o27
28 return o, h29
30 def bptt(self, x_sequence, y_sequence, max_norm=1.0):31 """32 随时间反向传播 + 梯度裁剪33 """34 T = len(x_sequence)35
36 # 前向传播37 o, h = self.forward(x_sequence)38
39 # 初始化梯度40 dW_xh = np.zeros_like(self.W_xh)41 dW_hh = np.zeros_like(self.W_hh)42 dW_ho = np.zeros_like(self.W_ho)43 db_h = np.zeros_like(self.b_h)44 db_o = np.zeros_like(self.b_o)45
46 # 输出层梯度47 delta_o = o - y_sequence # 假设MSE损失48
49 # 从最后一个时间步反向传播50 delta_h = np.zeros((self.hidden_dim, 1))51 for t in range(T-1, -1, -1):52 # 输出层参数梯度53 dW_ho += delta_o[t] @ h[t].T54 db_o += delta_o[t]55
56 # 从输出层到隐藏层的误差57 delta_h = self.W_ho.T @ delta_o[t] + delta_h58 # 通过tanh激活函数59 delta_h = delta_h * (1 - h[t]**2)60
61 # 隐藏层参数梯度(当前时间步)62 dW_xh += delta_h @ x_sequence[t].T63 db_h += delta_h64
65 # 对Whh的梯度:需要累加所有历史时间步66 # 这里简化处理,实际BPTT需要循环k67 dW_hh += delta_h @ h[t-1].T68
69 # 准备传递给下一层(更早的时间步)70 delta_h = self.W_hh.T @ delta_h71
72 # 梯度裁剪(基于全局范数)73 grads = [dW_xh, dW_hh, dW_ho, db_h, db_o]74 total_norm = np.sqrt(sum(np.sum(g**2) for g in grads))75 if total_norm > max_norm:76 scale = max_norm / total_norm77 dW_xh *= scale78 dW_hh *= scale79 dW_ho *= scale80 db_h *= scale81 db_o *= scale82
83 return dW_xh, dW_hh, dW_ho, db_h, db_o| 概念 | 核心要点 | 关键公式 |
|---|---|---|
| RNN前向传播 | 隐藏状态循环传递,所有时间步共享参数 | ht=ϕ(Wxhxt+Whhht−1+bh) |
| BPTT | 将RNN按时间步展开,沿时间反向传播误差 | ∂L/∂Whh=∑t=1T∑k=1t∂Lt/∂ht⋅(∏i=k+1tWhh⊤diag(ϕ′(hi−1)))⋅∂hk/∂Whh |
| 梯度消失 | 矩阵连乘谱范数<1,梯度指数衰减 | ∥∂ht/∂hk∥≈(γ⋅∥Whh∥)t−k |
| 梯度爆炸 | 矩阵连乘谱范数>1,梯度指数增长 | 同上,条件相反 |
| 梯度裁剪 | 限制梯度范数,保持方向不变 | gclipped=θ⋅g/∥g∥2 当 ∥g∥2>θ |
BPTT的三大启示:
理解BPTT的数学原理,不仅有助于理解为什么RNN难以训练长序列,也为理解更先进的序列模型(如Transformer中的自注意力机制)奠定了基础——Transformer彻底抛弃了循环结构,用并行化的注意力机制替代了时序上的串行传递,从根本上规避了BPTT的梯度困境。
延伸阅读:
按顺序完成这组文章,循序渐进地掌握主题
发现错误、内容过时或有改进想法?欢迎告诉我
根据本文分类与标签,为你推荐可能感兴趣的内容

系统讲解LSTM与GRU的门控机制设计哲学与数学原理。从LSTM的“双轨”架构(细胞状态+隐藏状态)与三门结构(遗忘门、输入门、输出门)出发,推导细胞状态更新公式如何通过按元素加法路径缓解梯度消失,并深入剖析GRU的两大简化及其参数量优势。附带讨论双向LSTM与深层堆叠LSTM的应用场景。
阅读文章
系统讲解ResNet残差网络的设计哲学与数学原理:从退化问题的本质出发,推导残差块F(x)+x如何通过恒等映射直通路径缓解梯度消失,详解Bottleneck块如何将参数量减少约94%,并对比ResNet(加法融合)与DenseNet(拼接融合)的梯度流动差异。
阅读文章
系统讲解激活函数与权重初始化的协同演进关系。从Sigmoid/Tanh的梯度饱和问题出发,推导ReLU及其变体(LeakyReLU、PReLU、ELU、GELU、Swish)如何解决梯度消失,并深入推导Xavier初始化(适用于Sigmoid/Tanh)和Kaiming初始化(适用于ReLU)的方差守恒数学原理。
阅读文章请使用微信扫描二维码分享
当前文章会保持在原页面