全球人名分类案例 - LSTM 模型训练
1 课程概览
本课讲解 LSTM 模型的训练。基于 RNN 训练代码复制粘贴修改,主要修改点:①initHidden 返回 hidden 和 c;②前向传播传入 hidden 和 c;③返回 output, hn, cn。LSTM 内部结构最复杂,训练时间理论上最长。
2 核心概念与定义
- LSTM 训练:基于 RNN 训练代码修改,增加细胞状态 $C$ 的处理。
- initHidden:LSTM 需同时初始化
hidden和c。 - 训练时间:LSTM 内部结构最复杂,训练时间理论上最长。
3 模型与算法详解
LSTM 训练代码修改要点
| 修改点 | RNN | LSTM |
|---|---|---|
| 模型实例化 | my_rnn = MyRNN(...) | my_lstm = MyLSTM(...) |
| initHidden | hidden = my_rnn.initHidden() | hidden, c = my_lstm.initHidden() |
| 前向传播 | output, hn = my_rnn(x, hidden) | output, hn, cn = my_lstm(x, hidden, c) |
| 返回值 | output, hn | output, hn, cn |
修改步骤
- 复制 RNN 训练代码
- 函数名:
dm_train_rnn→dm_train_lstm - 模型实例化:
my_rnn→my_lstm(变量名建议不改) - initHidden:返回
hidden, c两个值 - 前向传播:传入
hidden, c,返回output, hn, cn - 保存模型:
my_rnn→my_lstm
训练时间对比
| 模型 | 训练时间(1轮) | 说明 |
|---|---|---|
| RNN | ~51 秒 | 结构最简单 |
| LSTM | ~52 秒 | 结构最复杂 |
| GRU | ~55 秒 | 介于两者之间 |
注:训练时间受电脑运行环境影响,建议关闭其他软件。
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)$$
损失函数
$$\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_lstm():
# 1. 数据准备(与 RNN 一致)
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_lstm = MyLSTM(input_size, hidden_size, output_size, n_layers=1)
criterion = nn.NLLLoss()
optimizer = optim.Adam(my_lstm.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, c = my_lstm.initHidden()
# 修改点3:前向传播传入 hidden 和 c,返回 output, hn, cn
output, hn, cn = my_lstm(x, hidden, c)
# 计算损失
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_lstm.state_dict(), 'my_lstm.pth')
return loss_list, acc_list
loss_list, acc_list = dm_train_lstm()
6 重难点与易错提醒
- ❗重点:LSTM 的
initHidden返回hidden和c两个值。 - ❗重点:前向传播传入
hidden, c,返回output, hn, cn。 - ⚠️易错:不能扩写传参,
my_lstm(x, hidden, c)而非my_lstm(x, (hidden, c))。 - ⚠️易错:变量名建议不改,减少修改量。
- 💡深入理解:LSTM 内部结构最复杂,训练时间理论上最长。
- 💡深入理解:训练时间受电脑运行环境影响,建议关闭其他软件。
7 课堂问答精选
Q1:LSTM 训练代码相比 RNN 有哪些修改?
A:①initHidden 返回 hidden 和 c;②前向传播传入 hidden, c;③返回 output, hn, cn;④保存模型时改为 my_lstm。
Q2:为什么变量名建议不改?
A:如果改了变量名(如 my_rnn → my_lstm),后边所有引用该变量的地方都要改,增加修改量。保持变量名不变,只改模型类名即可。
Q3:LSTM 的训练时间为什么理论上最长?
A:因为 LSTM 内部结构最复杂,有遗忘门、输入门、输出门和细胞状态,计算量最大。
Q4:为什么不能扩写传参?
A:因为 MyLSTM 是自己写的类,forward 方法定义为 forward(self, input, hidden, c),需要分别传参。如果用 PyTorch 官方的 nn.LSTM,则可以传元组 (hidden, c)。
8 本课小结
- LSTM 修改点:
initHidden返回hidden, c;前向传播传入hidden, c。 - 返回值:
output, hn, cn(多一个cn)。 - 变量名建议不改,减少修改量。
- LSTM 训练时间理论上最长(结构最复杂)。
- 训练时间受电脑运行环境影响。
9 延伸思考
- LSTM 相比 RNN 在人名分类任务上效果如何?
- 如何优化训练过程以减少时间?