LSTM 算法计算过程详解

该文章已生成可运行项目,

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:

参数数值
输入维度 d2
隐藏状态维度 h3
序列长度 T1

输入数据:

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 的核心基础之一。

本文章已经生成可运行项目
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值