全球人名分类案例 - RNN 模型训练
1 课程概览
本课讲解 RNN 模型的训练过程。包括数据准备、模型与优化器初始化、训练循环(前向传播、计算损失、梯度清零、反向传播、参数更新)。使用 NLLLoss 损失函数和 Adam 优化器,每 100 次求平均损失用于绘图,每 2000 次打印日志。数据总量 20074 条人名,训练 1 轮约 50 秒。
2 核心概念与定义
- 学习率(Learning Rate, lr):控制参数更新步长,本课为
1e-3(0.001)。 - 训练轮数(Epochs):完整遍历数据集的次数,本课为 1 轮。
- NLLLoss(Negative Log Likelihood Loss):负对数似然损失,与 LogSoftmax 组合使用。
- Adam 优化器:自适应矩估计优化器。
- 三剑客:梯度清零、反向传播、参数更新。
3 模型与算法详解
训练流程
- 数据准备:读取数据 → 构建数据集 → 构建数据加载器
- 模型与优化器初始化:定义参数 → 创建模型 → 定义损失函数和优化器
- 参数初始化:记录开始时间、损失列表、准确率列表
- 训练循环:
- 外循环:每轮遍历
- 内循环:每批次遍历
- 获取结果 → 计算损失 → 梯度清零 → 反向传播 → 参数更新 → 累加损失
- 预测对错:每 100 次求平均损失,每 2000 次打印日志
训练参数
| 参数 | 值 | 说明 |
|---|---|---|
| lr | 1e-3 | 学习率 0.001 |
| epochs | 1 | 训练轮数 |
| 数据总量 | 20074 | 人名数量 |
| 训练时间 | ~50 秒 | 1 轮训练时间 |
损失函数选择
| 损失函数 | 对应激活函数 | 说明 |
|---|---|---|
| NLLLoss | LogSoftmax | 需手动加 LogSoftmax |
| CrossEntropyLoss | 无 | 内部已含 LogSoftmax |
日志打印频率
| 频率 | 操作 |
|---|---|
| 每 100 次 | 求平均损失,存入列表用于绘图 |
| 每 2000 次 | 打印轮数、损失、正确率 |
4 数学原理与推导
损失函数
$$\text{NLLLoss} = -\sum_{i=1}^{N} y_i \log(\hat{y}_i)$$
其中 $\hat{y}_i$ 为 LogSoftmax 的输出。
Adam 优化器
$$m_t = \beta_1 m_{t-1} + (1-\beta_1) g_t$$ $$v_t = \beta_2 v_{t-1} + (1-\beta_2) g_t^2$$ $$\hat{m}_t = \frac{m_t}{1-\beta_1^t}, \quad \hat{v}t = \frac{v_t}{1-\beta_2^t}$$ $$\theta_t = \theta{t-1} - \frac{\eta}{\sqrt{\hat{v}_t} + \epsilon} \hat{m}_t$$
平均损失
$$\text{avg_loss} = \frac{1}{100} \sum_{i=1}^{100} \text{loss}_i$$
5 代码示例
import torch
import torch.nn as nn
import torch.optim as optim
from tqdm import tqdm
# 1. 定义变量,记录学习率和训练轮数
my_lr = 1e-3 # 学习率 0.001
epochs = 1 # 训练轮数
def dm_train_rnn():
# ========== 1. 数据准备 ==========
# 1.1 读取数据
from name_classification import read_data
my_list_x, my_list_y = read_data()
# 1.2 构建数据集对象
from name_classification import NameClassDataset
name_class_dataset = NameClassDataset(my_list_x, my_list_y)
# 1.3 构建数据加载器
from torch.utils.data import DataLoader
my_data_loader = DataLoader(
name_class_dataset,
batch_size=1,
shuffle=True
)
# ========== 2. 模型与优化器初始化 ==========
# 2.1 定义模型参数
input_size = 57 # 字符表大小
hidden_size = 128 # 隐藏层维度
output_size = 18 # 国家数量
# 2.2 创建模型对象
my_rnn = MyRNN(input_size, hidden_size, output_size, n_layers=1)
# 2.3 定义损失函数和优化器
# 如果用了 LogSoftmax,则用 NLLLoss
criterion = nn.NLLLoss()
# 如果用 CrossEntropyLoss,则模型不需要 LogSoftmax
# criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(my_rnn.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):
# 4.1 获取结果
# x 形状:(seq_len, 1, 57)
# y 形状:(1,)
# 4.2 初始化隐藏状态
hidden = my_rnn.initHidden()
# 4.3 前向传播
output, hn = my_rnn(x, hidden)
# 4.4 计算损失
loss = criterion(output, y)
# 4.5 三剑客:梯度清零、反向传播、参数更新
optimizer.zero_grad() # 梯度清零
loss.backward() # 反向传播
optimizer.step() # 参数更新
# 4.6 累加损失
current_loss += loss.item()
step += 1
# 4.7 预测对错
pred = output.argmax(dim=1)
correct += (pred == y).sum().item()
# 4.8 每 100 次求平均损失,存入列表
if step % 100 == 0:
avg_loss = current_loss / 100
loss_list.append(avg_loss)
current_loss = 0
# 4.9 每 2000 次打印日志
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
# ========== 5. 训练结束 ==========
end_time = time.time()
print(f"训练总时间: {end_time - start_time:.2f} 秒")
return loss_list, acc_list
# 运行训练
loss_list, acc_list = dm_train_rnn()
输出示例
轮数: 1, 步数: 2000, 损失: 2.1234, 正确率: 0.4521
轮数: 1, 步数: 4000, 损失: 1.5678, 正确率: 0.5832
...
训练总时间: 50.23 秒
6 重难点与易错提醒
- ❗重点:三剑客——梯度清零、反向传播、参数更新。
- ❗重点:NLLLoss 与 LogSoftmax 组合使用,等价于 CrossEntropyLoss。
- ⚠️易错:使用 CrossEntropyLoss 时,模型不需要加 LogSoftmax。
- ⚠️易错:每 100 次求平均损失后,需重置
current_loss=0。 - 💡深入理解:每 2000 次打印日志,方便观察训练进度。
- 💡深入理解:
tqdm用于显示进度条,方便监控训练过程。
7 课堂问答精选
Q1:训练流程的核心步骤是什么?
A:①数据准备(读取数据、构建数据集、构建数据加载器);②模型与优化器初始化;③训练循环(前向传播、计算损失、梯度清零、反向传播、参数更新);④日志打印(每 100 次平均损失,每 2000 次打印日志)。
Q2:三剑客是什么?
A:梯度清零(optimizer.zero_grad())、反向传播(loss.backward())、参数更新(optimizer.step())。
Q3:NLLLoss 和 CrossEntropyLoss 有什么区别?
A:NLLLoss 需要与 LogSoftmax 组合使用,模型需要加 LogSoftmax 层。CrossEntropyLoss 内部已含 LogSoftmax,模型不需要加 LogSoftmax 层。
Q4:为什么每 100 次求平均损失?
A:数据总量 20074 条,每 100 次求平均损失并存入列表,方便后续绘制损失曲线图。20074 / 100 ≈ 200 个损失值。
8 本课小结
- 训练流程:数据准备 → 模型初始化 → 训练循环 → 日志打印。
- 三剑客:梯度清零、反向传播、参数更新。
- NLLLoss + LogSoftmax = CrossEntropyLoss。
- 每 100 次平均损失,每 2000 次打印日志。
- 数据总量 20074,训练 1 轮约 50 秒。
9 延伸思考
- 如何选择合适的学习率和训练轮数?
- 如何通过损失曲线判断模型是否收敛?