全球人名分类案例 - 搭建 LSTM 和 GRU 模型
1 课程概览
本课基于 RNN 模型代码,通过复制粘贴修改的方式搭建 LSTM 和 GRU 模型。LSTM 相比 RNN 多了细胞状态 $C$,需要传入和返回 $C$。GRU 与 RNN 接口一致,只需替换 API 名称。三个模型结构高度相似,体现了框架的复用性。
2 核心概念与定义
- LSTM(Long Short-Term Memory):长短时记忆网络,RNN 的变体,多了细胞状态。
- GRU(Gated Recurrent Unit):门控循环单元,RNN 的简化变体。
- 细胞状态(Cell State):LSTM 特有的状态,用 $C$ 表示,形状与隐藏状态相同。
- 框架复用:基于一个模型代码,通过简单修改搭建其他模型。
3 模型与算法详解
三个模型对比
| 对比项 | RNN | LSTM | GRU |
|---|---|---|---|
| API | nn.RNN | nn.LSTM | nn.GRU |
| 输入 | $X_t$, $H_{t-1}$ | $X_t$, $H_{t-1}$, $C_{t-1}$ | $X_t$, $H_{t-1}$ |
| 输出 | $H_t$, $Y_t$ | $H_t$, $C_t$, $Y_t$ | $H_t$, $Y_t$ |
| 初始化 | h0 | h0, c0 | h0 |
| 修改量 | 基础 | 改 API + 加 $C$ | 改 API |
LSTM 模型修改要点
- 类名:
MyRNN→MyLSTM - API:
nn.RNN→nn.LSTM - forward:传入 $C$,返回 $C$
- initHidden:初始化 $H$ 和 $C$
GRU 模型修改要点
- 类名:
MyRNN→MyGRU - API:
nn.RNN→nn.GRU - 其他不变:接口与 RNN 一致
LSTM 返回值形状
| 返回值 | 形状 | 说明 |
|---|---|---|
| output | (batch, output_size) | 预测的类别概率分布 |
| hn | (1, batch, hidden_size) | 最后一个时间步的隐藏状态 |
| cn | (1, batch, hidden_size) | 最后一个时间步的细胞状态 |
4 数学原理与推导
LSTM 前向传播
$$f_t = \sigma(W_f [h_{t-1}, x_t] + b_f)$$ $$i_t = \sigma(W_i [h_{t-1}, x_t] + b_i)$$ $$\tilde{C}t = \tanh(W_C [h{t-1}, x_t] + b_C)$$ $$C_t = f_t \cdot C_{t-1} + i_t \cdot \tilde{C}t$$ $$o_t = \sigma(W_o [h{t-1}, x_t] + b_o)$$ $$h_t = o_t \cdot \tanh(C_t)$$
GRU 前向传播
$$r_t = \sigma(W_r [h_{t-1}, x_t] + b_r)$$ $$z_t = \sigma(W_z [h_{t-1}, x_t] + b_z)$$ $$\tilde{h}t = \tanh(W [r_t \cdot h{t-1}, x_t] + b)$$ $$h_t = (1-z_t) \cdot h_{t-1} + z_t \cdot \tilde{h}_t$$
5 代码示例
LSTM 模型
import torch
import torch.nn as nn
class MyLSTM(nn.Module):
def __init__(self, input_size, hidden_size, output_size, n_layers=1):
super(MyLSTM, self).__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.output_size = output_size
self.n_layers = n_layers
# LSTM 层(替换 nn.RNN → nn.LSTM)
self.rnn = nn.LSTM(
self.input_size,
self.hidden_size,
self.n_layers
)
# 全连接层(不变)
self.linear = nn.Linear(self.hidden_size, self.output_size)
# LogSoftmax(不变)
self.softmax = nn.LogSoftmax(dim=1)
def forward(self, input, hidden, c):
# 添加批次维度
input = input.unsqueeze(1)
# LSTM 计算(需要传入 hidden 和 c)
output, (hn, cn) = self.rnn(input, (hidden, c))
# 取最后一个时间步
output = output[-1]
# 全连接层
output = self.linear(output)
# LogSoftmax
output = self.softmax(output)
# 返回 output, hn, cn(多返回 cn)
return output, hn, cn
def initHidden(self):
# 同时初始化 hidden 和 c
hidden = c = torch.zeros(
self.n_layers, 1, self.hidden_size
)
return hidden, c
GRU 模型
class MyGRU(nn.Module):
def __init__(self, input_size, hidden_size, output_size, n_layers=1):
super(MyGRU, self).__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.output_size = output_size
self.n_layers = n_layers
# GRU 层(替换 nn.RNN → nn.GRU)
self.rnn = nn.GRU(
self.input_size,
self.hidden_size,
self.n_layers
)
# 全连接层(不变)
self.linear = nn.Linear(self.hidden_size, self.output_size)
# LogSoftmax(不变)
self.softmax = nn.LogSoftmax(dim=1)
def forward(self, input, hidden):
# 添加批次维度
input = input.unsqueeze(1)
# GRU 计算(与 RNN 一致)
output, hn = self.rnn(input, hidden)
# 取最后一个时间步
output = output[-1]
# 全连接层
output = self.linear(output)
# LogSoftmax
output = self.softmax(output)
return output, hn
def initHidden(self):
# 与 RNN 一致
return torch.zeros(self.n_layers, 1, self.hidden_size)
6 重难点与易错提醒
- ❗重点:LSTM 相比 RNN 多了细胞状态 $C$,需要传入和返回。
- ❗重点:GRU 与 RNN 接口一致,只需替换 API 名称。
- ⚠️易错:LSTM 的
forward需要传入hidden和c两个参数。 - ⚠️易错:LSTM 返回
(hn, cn)元组,需要拆包。 - 💡深入理解:三个模型结构高度相似,体现了框架的复用性。
- 💡深入理解:
self.rnn名称可以不变,因为 LSTM/GRU 都属于 RNN 系列。
7 课堂问答精选
Q1:LSTM 相比 RNN 有什么不同?
A:LSTM 相比 RNN 多了细胞状态 $C$,需要传入 $C_{t-1}$ 和返回 $C_t$。initHidden 需要同时初始化 hidden 和 c。
Q2:GRU 相比 RNN 有什么不同?
A:GRU 与 RNN 接口完全一致,只需将 nn.RNN 替换为 nn.GRU,其他代码不变。
Q3:三个模型的返回值有什么区别?
A:RNN 和 GRU 返回 (output, hn)。LSTM 返回 (output, hn, cn),多了一个细胞状态 cn。
Q4:为什么 self.rnn 名称可以不变?
A:因为 LSTM 和 GRU 都属于 RNN 系列的变体,使用 self.rnn 作为属性名不影响功能,只需将 API 替换为 nn.LSTM 或 nn.GRU 即可。
8 本课小结
- LSTM:替换 API + 加细胞状态 $C$。
- GRU:只需替换 API,其他不变。
- LSTM 返回
(output, hn, cn),RNN/GRU 返回(output, hn)。 initHidden:LSTM 需初始化hidden和c。- 三个模型结构高度相似,体现框架复用性。
9 延伸思考
- LSTM 和 GRU 在什么场景下选择使用?
- 三个模型的训练效果会有什么差异?