LSTM 模型 - 原理图解(下)
1 课程概览
本课讲解 LSTM 各个门的具体作用和公式。重点讲解遗忘门、输入门、细胞状态、输出门的作用和公式。门值接近 1 表示保留,接近 0 表示忘掉。
2 核心概念与定义
- 遗忘门(Forget Gate):决定记忆细胞中哪些信息可以丢掉。
- 输入门(Input Gate):决定哪些新信息被加入到记忆细胞。
- 细胞状态(Cell State):长期记忆。
- 输出门(Output Gate):决定输出哪些信息。
- 门值:接近 1 表示保留,接近 0 表示忘掉。
3 模型与算法详解
遗忘门
决定记忆细胞中哪些信息可以丢掉。
- 作用:通过当前输入和上一时刻的隐藏状态,决定记忆细胞中哪些信息可以丢掉
- 举例:记日记时发现昨天记的天气预报对今天没用,从长期记忆本中划掉
遗忘门公式
$$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$$
其中:
- $f_t$ 是遗忘门的门值
- $\sigma$ 是 sigmoid 激活函数
- $W_f$ 是权重矩阵
- $h_{t-1}$ 是上一时刻的隐藏状态
- $x_t$ 是当前输入
- $b_f$ 是偏置
门值的含义
| 门值 | 含义 |
|---|---|
| 接近 1 | 保留这条记忆(有用) |
| 接近 0 | 忘掉这条记忆(没用) |
门的比喻
门可以理解为水龙头。
- 门开得大(接近 1):信息流过的多
- 门关得小(接近 0):信息流过的少
4 数学原理与推导
遗忘门
$$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$$
细胞状态更新
$$C_t = f_t \cdot C_{t-1} + i_t \cdot \tilde{C}_t$$
其中:
- $C_t$ 是当前细胞状态
- $f_t$ 是遗忘门的门值
- $C_{t-1}$ 是上一时刻的细胞状态
- $i_t$ 是输入门的门值
- $\tilde{C}_t$ 是候选记忆
5 代码示例
import torch
import torch.nn as nn
# LSTM 模型
lstm = nn.LSTM(input_size=5, hidden_size=6, num_layers=1, batch_first=True)
# 输入数据
input_data = torch.randn(1, 3, 5) # [batch_size, seq_len, input_size]
# 初始隐藏状态和细胞状态
h0 = torch.zeros(1, 1, 6) # [num_layers, batch_size, hidden_size]
c0 = torch.zeros(1, 1, 6) # [num_layers, batch_size, hidden_size]
# 前向传播
output, (hn, cn) = lstm(input_data, (h0, c0))
print(f"输入: {input_data.shape}")
print(f"输出: {output.shape}")
print(f"隐藏状态: {hn.shape}")
print(f"细胞状态: {cn.shape}")
6 重难点与易错提醒
- ❗重点:遗忘门决定记忆细胞中哪些信息可以丢掉。
- ❗重点:门值接近 1 表示保留,接近 0 表示忘掉。
- ❗重点:遗忘门公式:$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$。
- 💡技巧:门可以理解为水龙头,开得大信息流过的多。
7 课堂问答精选
Q1:遗忘门的作用是什么?
A:遗忘门通过当前输入和上一时刻的隐藏状态,决定记忆细胞中哪些信息可以丢掉。
Q2:门值的含义是什么?
A:门值接近 1 表示保留这条记忆(有用),接近 0 表示忘掉这条记忆(没用)。
Q3:遗忘门的公式是什么?
A:$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$,其中 $\sigma$ 是 sigmoid 激活函数。
8 本课小结
- 遗忘门决定记忆细胞中哪些信息可以丢掉。
- 门值接近 1 表示保留,接近 0 表示忘掉。
- 遗忘门公式:$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$。
9 延伸思考
- 输入门和输出门的作用是什么?
- Bi-LSTM 是什么?