全球人名分类案例 - RNN 模型预测
1 课程概览
本课讲解全球人名分类案例的 RNN 模型预测。训练完成后,加载模型参数,输入人名,预测可能来自哪个国家,返回前 3 名结果。预测时需要将人名转换为 one-hot 张量。
2 核心概念与定义
- 模型预测:加载训练好的模型参数,输入人名,预测国家。
- name_to_tensor 函数:将人名转换为 one-hot 张量。
- top_k:返回前 k 名结果。
3 模型与算法详解
预测流程
1. 将人名转换为 one-hot 张量
2. 加载模型参数
3. 预测
4. 返回前 3 名结果
训练过程总结
以 RNN 模型为例。
1. 读数据到内存
2. 构建 Dataset
3. 构建 DataLoader
4. 实例化模型、损失函数、优化器
5. 外循环控制轮数,内循环控制迭代次数
6. 算预测值 → 算损失 → 梯度清零 → 反向传播 → 梯度更新
7. 每 100 次算一次平均损失
8. 每 2400 次打日志
9. 保存模型
模型选择
综合分析选择 GRU。
| 指标 | RNN | LSTM | GRU |
|---|---|---|---|
| 收敛速度 | 最快 | 较慢 | 中等 |
| 训练时长 | 最短 | 较长 | 中等 |
| 准确率 | 一般 | 一般 | 最好 |
4 数学原理与推导
预测
$$\text{output} = \text{model}(\text{name_to_tensor}(\text{name}))$$
Top-K
$$\text{top_k} = \text{topk}(\text{output}, k=3)$$
5 代码示例
import torch
def name_to_tensor(name, all_letters, n_letters):
"""
将人名转换为 one-hot 张量
Args:
name: 人名
all_letters: 全局字母表
n_letters: 字母表大小
Returns:
tensor: one-hot 张量
"""
tensor = torch.zeros(len(name), 1, n_letters)
for i, letter in enumerate(name):
tensor[i][0][all_letters.find(letter)] = 1
return tensor
def my_predict(name, model, all_letters, n_letters, all_countries, top_k=3):
"""
预测人名来自哪个国家
Args:
name: 人名
model: 训练好的模型
all_letters: 全局字母表
n_letters: 字母表大小
all_countries: 国家列表
top_k: 返回前 k 名
Returns:
前 k 名结果
"""
# 1. 将人名转换为 one-hot 张量
input_tensor = name_to_tensor(name, all_letters, n_letters)
# 2. 加载模型参数
model.load_state_dict(torch.load('model_params.pth'))
model.eval()
# 3. 预测
with torch.no_grad():
output = model(input_tensor)
# 4. 返回前 k 名
topv, topi = output.topk(top_k, dim=1, largest=True)
for i in range(top_k):
value = topv[0][i]
index = topi[0][i]
country = all_countries[index]
print(f"第 {i+1} 名: {country} (概率: {value:.4f})")
# 测试
if __name__ == "__main__":
print("=== 全球人名分类案例 - RNN 模型预测 ===")
import string
all_letters = string.ascii_letters + " .,;'"
n_letters = len(all_letters)
all_countries = ["Chinese", "English", "German", "French"]
# 假设 model 已定义
# model = RNNModel(n_letters, 64, n_countries)
# 预测
# my_predict("Zhang", model, all_letters, n_letters, all_countries, top_k=3)
# my_predict("Smith", model, all_letters, n_letters, all_countries, top_k=3)
# my_predict("Muller", model, all_letters, n_letters, all_countries, top_k=3)
代码说明
| 代码 | 说明 |
|---|---|
name_to_tensor(name, ...) | 将人名转换为 one-hot 张量 |
model.load_state_dict(...) | 加载模型参数 |
model.eval() | 设置为评估模式 |
torch.no_grad() | 不计算梯度 |
output.topk(top_k, ...) | 返回前 k 名 |
6 重难点与易错提醒
- ❗重点:预测时需要将人名转换为 one-hot 张量。
- ❗重点:使用
model.load_state_dict()加载模型参数。 - ❗重点:使用
model.eval()设置为评估模式。 - ❗重点:使用
torch.no_grad()不计算梯度。 - 💡技巧:LSTM 和 GRU 的预测代码与 RNN 类似。
7 课堂问答精选
Q1:如何进行模型预测?
A:将人名转换为 one-hot 张量,加载模型参数,预测,返回前 3 名结果。
Q2:为什么需要 name_to_tensor 函数?
A:模型在训练时使用 one-hot 张量,预测时也需要将人名转换为 one-hot 张量。
Q3:如何选择模型?
A:综合分析,GRU 的准确率最好,所以优先选择 GRU。
8 本课小结
- 预测流程:将人名转换为 one-hot 张量 → 加载模型参数 → 预测 → 返回前 3 名。
- LSTM 和 GRU 的预测代码与 RNN 类似。
- 综合分析,GRU 的准确率最好。
9 延伸思考
- 如何优化案例?
- 如何将损失、时间、准确率导出到文件?