全球人名分类案例 - GRU 模型训练
1 课程概览
本课讲解 GRU 模型的训练。基于 LSTM 训练代码修改,只需删除细胞状态 $C$ 相关代码,并将模型类名改为 MyGRU。GRU 接口与 RNN 一致,训练时间介于 RNN 和 LSTM 之间。编程只有第一次是难的,后续模型训练代码复用性高。
2 核心概念与定义
- GRU 训练:基于 LSTM 训练代码修改,删除细胞状态 $C$。
- 代码复用:三个模型训练代码高度相似,只需少量修改。
3 模型与算法详解
GRU 训练代码修改要点
| 修改点 | LSTM | GRU |
|---|---|---|
| 模型实例化 | my_lstm = MyLSTM(...) | my_gru = MyGRU(...) |
| initHidden | hidden, c = my_lstm.initHidden() | hidden = my_gru.initHidden() |
| 前向传播 | output, hn, cn = my_lstm(x, hidden, c) | output, hn = my_gru(x, hidden) |
| 返回值 | output, hn, cn | output, hn |
修改步骤
- 复制 LSTM 训练代码
- 函数名:
dm_train_lstm→dm_train_gru - 模型实例化:
my_lstm→my_gru - initHidden:删除
c,只返回hidden - 前向传播:删除
c参数和cn返回值 - 保存模型:
my_lstm→my_gru
训练时间对比
| 模型 | 训练时间(1轮) | 结构复杂度 |
|---|---|---|
| RNN | ~51 秒 | 最简单 |
| GRU | ~55 秒 | 介于两者 |
| LSTM | ~52 秒 | 最复杂 |
注:训练时间受电脑运行环境影响,建议多跑几轮取平均值。
编程感悟
编程只有第一次做是难的,你做完以后后边再来写,所以类都很简单。
- 第一个模型训练代码:~40 分钟
- 后两个模型训练代码:~4 分钟(复制粘贴修改)
4 数学原理与推导
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$$
损失函数
$$\text{NLLLoss} = -\sum_{i=1}^{N} y_i \log(\hat{y}_i)$$
5 代码示例
import torch
import torch.nn as nn
import torch.optim as optim
from tqdm import tqdm
my_lr = 1e-3
epochs = 1
def dm_train_gru():
# 1. 数据准备(与 RNN/LSTM 一致)
from name_classification import read_data, NameClassDataset
from torch.utils.data import DataLoader
my_list_x, my_list_y = read_data()
name_class_dataset = NameClassDataset(my_list_x, my_list_y)
my_data_loader = DataLoader(name_class_dataset, batch_size=1, shuffle=True)
# 2. 模型与优化器初始化
input_size = 57
hidden_size = 128
output_size = 18
# 修改点1:模型实例化
my_gru = MyGRU(input_size, hidden_size, output_size, n_layers=1)
criterion = nn.NLLLoss()
optimizer = optim.Adam(my_gru.parameters(), lr=my_lr)
# 3. 参数初始化
import time
start_time = time.time()
loss_list = []
acc_list = []
# 4. 训练过程
for epoch in range(epochs):
current_loss = 0
step = 0
correct = 0
for x, y in tqdm(my_data_loader):
# 修改点2:initHidden 只返回 hidden(删除 c)
hidden = my_gru.initHidden()
# 修改点3:前向传播(删除 c 参数和 cn 返回值)
output, hn = my_gru(x, hidden)
# 计算损失
loss = criterion(output, y)
# 三剑客
optimizer.zero_grad()
loss.backward()
optimizer.step()
current_loss += loss.item()
step += 1
pred = output.argmax(dim=1)
correct += (pred == y).sum().item()
if step % 100 == 0:
avg_loss = current_loss / 100
loss_list.append(avg_loss)
current_loss = 0
if step % 2000 == 0:
avg_acc = correct / 2000
acc_list.append(avg_acc)
print(f"轮数: {epoch+1}, 步数: {step}, "
f"损失: {avg_loss:.4f}, 正确率: {avg_acc:.4f}")
correct = 0
end_time = time.time()
print(f"训练总时间: {end_time - start_time:.2f} 秒")
# 修改点4:保存模型
torch.save(my_gru.state_dict(), 'my_gru.pth')
return loss_list, acc_list
loss_list, acc_list = dm_train_gru()
6 重难点与易错提醒
- ❗重点:GRU 训练代码只需删除 LSTM 中的细胞状态 $C$ 相关代码。
- ❗重点:GRU 接口与 RNN 一致,
initHidden只返回hidden。 - ⚠️易错:删除
c时要同步删除前向传播的参数和返回值。 - 💡深入理解:三个模型训练代码高度相似,体现代码复用性。
- 💡深入理解:训练时间受电脑运行环境影响,建议多跑几轮取平均值。
7 课堂问答精选
Q1:GRU 训练代码相比 LSTM 有哪些修改?
A:①initHidden 只返回 hidden(删除 c);②前向传播删除 c 参数和 cn 返回值;③模型实例化改为 MyGRU。
Q2:GRU 的训练时间为什么介于 RNN 和 LSTM 之间?
A:GRU 的结构复杂度介于 RNN 和 LSTM 之间。GRU 有重置门和更新门(2 个门),比 RNN 复杂,比 LSTM(4 个门)简单。
Q3:为什么第一个模型训练代码写了 40 分钟,后两个只用了 4 分钟?
A:因为三个模型训练代码高度相似,只需复制粘贴并做少量修改。编程只有第一次是难的,后续代码复用性高。
Q4:如何获得更准确的训练时间?
A:将训练轮数从 1 轮改为 10 轮、20 轮或 100 轮,计算总时长,取平均值。一次不能论英雄,需要多跑几次。
8 本课小结
- GRU 修改点:删除 LSTM 中的细胞状态 $C$ 相关代码。
initHidden只返回hidden。- 前向传播:
output, hn = my_gru(x, hidden)。 - GRU 训练时间介于 RNN 和 LSTM 之间。
- 代码复用性高,后续模型修改量少。
9 延伸思考
- GRU 相比 LSTM 在性能上有什么差异?
- 如何选择合适的模型(RNN/LSTM/GRU)?