文章目录
LSTM (长短期记忆网络) 原理解析笔记
1. 什么是LSTM?
- LSTM (Long Short-Term Memory Network) 是循环神经网络(RNN)的一种变体
- 主要解决长期依赖问题(即模型能够记住长期信息)
- 与GRU相比,LSTM参数更多,通常在处理更复杂的长期依赖任务时表现更好
2. LSTM 总体结构
核心概念:
- 长期记忆:存储在 细胞状态(cell state),记为 C
- 短期记忆:存储在 隐藏状态(hidden state),记为 h
LSTM 通过三个“门”结构来控制信息的流动和记忆的更新:
- 遗忘门(Forget Gate)
- 输入门(Input Gate)
- 输出门(Output Gate)
如下图所示:

3. LSTM 执行步骤详解
(1) 遗忘门(Forget Gate)

作用:决定从长期记忆(cell state)中丢弃哪些信息
公式:
ft = σ ( Wf ⋅ [ ht − 1 , xt ] + bf )
- ft 是一个介于0~1之间的值(通过Sigmoid函数得到)
- 0 表示“完全遗忘”,1 表示“完全保留”
(2) 输入门(Input Gate)

作用:决定哪些新信息将被存储到细胞状态中
分为两部分:
-
it:输入重要性系数(缩放因子)
it = σ ( Wi ⋅ [ ht − 1 , xt ] + bi )
-
C~t:候选更新向量(表示新的信息)
C~t = tanh ( Wc ⋅ [ ht − 1 , xt ] + bc )
最终输入向量为:
输入向量=it⋅C~t
(3) 更新长期记忆(Cell State Update)

作用:结合遗忘门和输入门,更新细胞状态
Ct= ft ⋅ Ct −1+ it ⋅ C~t
(4) 输出门(Output Gate)

作用:基于当前细胞状态,决定输出什么信息(即更新隐藏状态)
-
输出比例系数:
ot = σ ( Wo ⋅ [ht−1,xt] + bo)
-
更新隐藏状态:
ht = ot ⋅ tanh(Ct)
4. LSTM 变体
(1) 窥视孔连接(Peephole Connections)

- 在计算门控信号(ft,it,ot)时,除了使用 ht−1 和 xt,还引入细胞状态 Ct−1 或 Ct
(2) 耦合遗忘门与输入门

- 使用 1−ft 代替输入门 it,使遗忘与输入形成互补关系
(3) GRU(门控循环单元)
- LSTM的简化版本,将遗忘门和输入门合并为“更新门”
- 只有两个门:更新门和重置门
5. 总结对比
| 组件 | 作用 | 输出范围 | 计算公式 |
|---|---|---|---|
| 遗忘门 ft | 控制历史信息的保留比例 | (0, 1) | σ(Wf[ht−1,xt]+bf) |
| 输入门 it | 控制新信息的重要性 | (0, 1) | σ(Wi[ht−1,xt]+bi) |
| 输出门 ot | 控制输出信息的比例 | (0, 1) | σ(Wo[ht−1,xt]+bo) |
| 候选记忆C~t | 生成候选更新信息 | (-1, 1) | tanh(Wc[ht−1,xt]+bc) |
6.代码实现:
import torch
from torch import nn
class LSTMCell(nn.Module):
def __init__(self,input_size,hidden_size):
super().__init__()
self.hidden_size = hidden_size
self.linear_ft = nn.Linear(input_size + hidden_size, hidden_size)
self.linear_it = nn.Linear(input_size + hidden_size, hidden_size)
self.linear_ct = nn.Linear(input_size + hidden_size, hidden_size)
self.linear_ot = nn.Linear(input_size + hidden_size, hidden_size)
self.sigmoid = nn.Sigmoid()
self.tanh = nn.Tanh()
def forward(self, x,c=None,h=None):
_x = torch.concat([h,x],dim=1)
#遗忘门
ft = self.sigmoid(self.linear_ft(_x))
#输入门
it = self.sigmoid(self.linear_it(_x))
ct = self.tanh(self.linear_ct(_x))
#更新长期记忆
c = ft*c + it*ct
#输出门
ot = self.sigmoid(self.linear_ot(_x))
h = ot*torch.tanh(c)
return c,h
class LSTM(nn.Module):
def __init__(self,input_size,hidden_size):
super().__init__()
self.cell = LSTMCell(input_size,hidden_size)
self.hidden_size = hidden_size
self.fc_out = nn.Linear(hidden_size,hidden_size)
def forward(self, x,c=None,h=None):
N,L,input_size = x.shape
if h is None:
h = torch.zeros(N,self.hidden_size)
if c is None:
c = torch.zeros(N,self.hidden_size)
outputs = []
for i in range(L):
c, h = self.cell(x[:,i],c,h)
out = self.fc_out(h)
outputs.append(out)
outputs = torch.stack(outputs,dim=1)
return outputs,c,h
if __name__ == '__main__':
x = torch.rand(5,6,10)
model = LSTM(10,20)
y,c,h = model(x)
print(y.shape)
print(c.shape)
print(h.shape)
:LSTM长短期记忆网络&spm=1001.2101.3001.5002&articleId=151407309&d=1&t=3&u=1025f3f4ce614b1389e1802e26e086f1)
5249

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



