LSTM 算法计算过程详解
本文详细讲解 LSTM(Long Short-Term Memory,长短期记忆网络)的完整计算流程,包括:
- 正向传播(Forward Pass)
- Cell 状态更新
- 隐藏状态输出
- 损失函数计算
- 反向传播(Backward Pass)
- 梯度下降参数更新
1. 什么是 LSTM?
LSTM(Long Short-Term Memory)是一种特殊的循环神经网络(RNN)。
它主要解决普通 RNN 的:
- 长期依赖问题
- 梯度消失问题
LSTM 通过:
- 遗忘门(Forget Gate)
- 输入门(Input Gate)
- 候选记忆(Candidate)
- 输出门(Output Gate)
来控制信息的保存与遗忘。
2. 本例网络结构
本次 Excel 示例中采用的是一个极简化的 LSTM:
| 参数 | 数值 |
|---|---|
| 输入维度 d | 2 |
| 隐藏状态维度 h | 3 |
| 序列长度 T | 1 |
输入数据:
x₁ = [0.6, 0.9]
上一时刻隐藏状态:
h₀ = [0.2, 0.1, 0.4]
上一时刻 Cell 状态:
C₀ = [0, 0, 0]
3. LSTM 的四个门
LSTM 的核心是四个门:
| 门 | 符号 | 作用 |
|---|---|---|
| 遗忘门 | f | 决定遗忘多少旧记忆 |
| 输入门 | i | 决定写入多少新信息 |
| 候选记忆 | u | 生成新的候选内容 |
| 输出门 | o | 决定输出多少隐藏状态 |
4. Step1:拼接输入向量
LSTM 首先将:
- 上一时刻隐藏状态 h₀
- 当前输入 x₁
拼接成一个长向量:
[h₀, x₁]
代入数据:
[0.2, 0.1, 0.4, 0.6, 0.9]
这个向量会同时输入到四个门。
5. Step2:计算四个门的 net 值
四个门都要先计算线性加权和:
net_f = wf · [h₀,x₁]
net_i = wi · [h₀,x₁]
net_u = wu · [h₀,x₁]
net_o = wo · [h₀,x₁]
其中:
- wf:遗忘门权重
- wi:输入门权重
- wu:候选记忆权重
- wo:输出门权重
遗忘门计算
假设:
wf[1] = [0.2,0.1,0.3,0.4,0.5]
则:
0.2×0.2
+0.1×0.1
+0.3×0.4
+0.4×0.6
+0.5×0.9
= 0.86
同理计算其它两行:
net_f = [0.86, 0.73, 0.58]
输入门计算
net_i = [0.32, 0.51, 0.77]
候选记忆计算
net_u = [1.15, 1.36, 1.52]
输出门计算
net_o = [0.92, 1.08, 1.26]
6. Step3:激活函数
LSTM 中:
- f、i、o 使用 sigmoid
- u 使用 tanh
sigmoid 函数
公式:
σ(x)=1/(1+e^(-x))
特点:
- 输出范围:(0,1)
- 表示“保留比例”
tanh 函数
公式:
tanh(x)
特点:
- 输出范围:(-1,1)
- 可以表示正负记忆
遗忘门输出
f₁ = σ(net_f)
结果:
f₁ = [0.703, 0.675, 0.641]
输入门输出
i₁ = [0.579, 0.625, 0.683]
候选记忆输出
u₁ = [0.818, 0.876, 0.909]
输出门输出
o₁ = [0.715, 0.746, 0.779]
7. Step4:更新 Cell 状态
LSTM 最核心的公式:
C₁ = f₁ ⊙ C₀ + i₁ ⊙ u₁
其中:
- ⊙ 表示逐元素乘法
- 第一项:保留旧记忆
- 第二项:加入新记忆
由于:
C₀=[0,0,0]
所以:
C₁ = i₁ ⊙ u₁
计算:
C₁[1] = 0.579 × 0.818 = 0.474
C₁[2] = 0.625 × 0.876 = 0.548
C₁[3] = 0.683 × 0.909 = 0.621
因此:
C₁=[0.474,0.548,0.621]
8. Step5:计算隐藏状态
隐藏状态输出公式:
h₁ = o₁ ⊙ tanh(C₁)
先计算:
tanh(C₁)
结果:
[0.441,0.499,0.552]
然后逐元素相乘:
h₁ =
[0.715,0.746,0.779]
⊙
[0.441,0.499,0.552]
得到:
h₁=[0.315,0.372,0.430]
这就是最终输出。
9. 损失函数计算
假设真实值:
y=[1,1,1]
损失函数采用平方误差:
L = Σ(h₁-y)²
代入:
L =
(0.315-1)²
+(0.372-1)²
+(0.430-1)²
计算得到:
L ≈ 1.25
10. 反向传播(Backward Pass)
LSTM 的反向传播本质上是:
链式求导
误差从:
Loss
→ h₁
→ o₁
→ C₁
→ i₁/u₁/f₁
→ 各权重矩阵
逐层传播。
11. 输出门梯度
因为:
h₁ = o₁ ⊙ tanh(C₁)
所以:
∂L/∂o₁
=
∂L/∂h₁
⊙ tanh(C₁)
再进一步:
∂L/∂wo
=
∂L/∂o₁
· [h₀,x₁]^T
得到输出门权重梯度矩阵。
12. Cell 状态梯度
Cell 状态是 LSTM 的核心。
因为:
h₁ = o₁ ⊙ tanh(C₁)
所以:
∂L/∂C₁
=
∂L/∂tanh(C₁)
⊙ tanh'(C₁)
其中:
tanh'(x)=1-tanh²(x)
13. 输入门梯度
由于:
C₁ = f₁⊙C₀ + i₁⊙u₁
因此:
∂C₁/∂i₁ = u₁
得到:
∂L/∂i₁
=
∂L/∂C₁
⊙ u₁
再经过 sigmoid 导数:
σ'(x)=σ(x)(1-σ(x))
最终求得:
∂L/∂wi
14. 候选记忆梯度
同理:
∂C₁/∂u₁ = i₁
因此:
∂L/∂u₁
=
∂L/∂C₁
⊙ i₁
再经过:
tanh'(x)=1-tanh²(x)
得到:
∂L/∂wu
15. 遗忘门梯度
由于:
∂C₁/∂f₁ = C₀
而本例:
C₀=[0,0,0]
因此:
∂L/∂f₁ = 0
所以:
∂L/∂wf = 0
这说明:
- 当前样本中
- 遗忘门没有贡献梯度
- 因为没有历史记忆可遗忘
这是完全正常的。
16. 参数更新(Gradient Descent)
最后使用梯度下降更新参数:
W_new = W_old - η × gradient
其中:
η = 0.01
例如:
wo(new)
=
wo(old)
-
0.01×∂L/∂wo
其它参数:
- wf
- wi
- wu
也都按同样方式更新。
17. LSTM 的核心思想总结
LSTM 的本质:
让网络“记住重要信息”
并“忘掉无用信息”
核心公式只有两个:
Cell 状态更新
Cₜ = fₜ⊙Cₜ₋₁ + iₜ⊙uₜ
含义:
- f:保留多少旧记忆
- i:写入多少新记忆
- u:新记忆内容
隐藏状态输出
hₜ = oₜ⊙tanh(Cₜ)
含义:
- o:输出多少记忆
- tanh©:当前长期记忆
18. 本文总结
本文完整展示了:
- LSTM 四个门的计算
- Cell 状态更新
- 隐藏状态输出
- Loss 计算
- 反向传播链式求导
- 梯度矩阵计算
- 参数更新
通过 Excel 表格,可以非常直观地理解:
- LSTM 为什么能记忆长期信息
- 各个门的真实作用
- 梯度如何在网络中传播
这也是学习深度学习与 NLP 的核心基础之一。

572

被折叠的 条评论
为什么被折叠?



