全球人名分类案例 - 测试三种模型
1 课程概览
本课测试 RNN、LSTM、GRU 三种模型。流程包括:定义变量、加载数据、创建数据集和数据加载器、模型初始化。使用真实数据(从文件加载)而非随机生成的数据。数据加载器的批次大小只能为 1,因为人名长度不一致。
2 核心概念与定义
- 数据加载器(DataLoader):按批次获取数据的工具。
- 批次大小(Batch Size):每次获取的样本数,人名分类中只能为 1。
- NLLLoss(Negative Log Likelihood Loss):负对数似然损失,与 LogSoftmax 组合使用。
3 模型与算法详解
测试流程
- 定义变量:记录输入维度、隐藏状态、输出维度
- 加载数据:从文件读取人名和国家标签
- 创建数据集对象:
NameClassDataset - 创建数据加载器:
DataLoader - 模型初始化:实例化 RNN、LSTM、GRU 模型
变量定义
| 变量 | 值 | 含义 |
|---|---|---|
| n_letters | 57 | 字符表大小 |
| category_num | 18 | 国家数量 |
| input_size | 57 | 输入维度 |
| hidden_size | 128 | 隐藏层维度 |
| output_size | 18 | 输出维度 |
数据加载器参数
| 参数 | 值 | 说明 |
|---|---|---|
| dataset | NameClassDataset | 数据集对象 |
| batch_size | 1 | 批次大小(必须为 1) |
| shuffle | True | 是否打乱 |
批次大小为什么为 1
人名长度不一致,无法批量处理。
- 人名 "丁" → 4 个字符
- 人名 "欧阳" → 6 个字符
- 人名 "张" → 15 个字符
不同长度的人名无法组成批次,因此 batch_size=1。
4 数学原理与推导
本课为测试流程,无数学推导。
5 代码示例
import torch
from torch.utils.data import DataLoader
# 1. 定义变量,记录输入维度、隐藏状态、输出维度
n_letters = 57 # 字符表大小
category_num = 18 # 国家数量
input_size = n_letters # 57
hidden_size = 128 # 隐藏层维度
output_size = category_num # 18
# 2. 加载数据
from name_classification import read_data
my_list_x, my_list_y = read_data()
# 3. 创建数据集对象
from name_classification import NameClassDataset
name_class_dataset = NameClassDataset(my_list_x, my_list_y)
# 4. 创建数据加载器
my_data_loader = DataLoader(
name_class_dataset,
batch_size=1, # 批次大小必须为 1(人名长度不一致)
shuffle=True # 打乱数据
)
# 5. 模型初始化
# RNN
my_rnn = MyRNN(input_size, hidden_size, output_size, n_layers=1)
# LSTM
my_lstm = MyLSTM(input_size, hidden_size, output_size, n_layers=1)
# GRU
my_gru = MyGRU(input_size, hidden_size, output_size, n_layers=1)
# 6. 测试数据加载器
for x, y in my_data_loader:
print("x shape:", x.shape)
print("y:", y)
break # 只取第一个批次测试
数据加载流程
文件 → read_data() → my_list_x, my_list_y
→ NameClassDataset → name_class_dataset
→ DataLoader → my_data_loader
→ for x, y in my_data_loader
6 重难点与易错提醒
- ❗重点:批次大小
batch_size=1,因为人名长度不一致。 - ❗重点:使用真实数据(从文件加载),而非随机生成的数据。
- ⚠️易错:数据加载器需要传入数据集对象,而非直接传数据。
- 💡深入理解:
n_letters和category_num可替代硬编码的数字,提高代码可读性。 - 💡深入理解:
shuffle=True可以打乱数据顺序,提高训练效果。
7 课堂问答精选
Q1:为什么批次大小只能为 1?
A:因为人名长度不一致(如"丁"4 个字符,"欧阳"6 个字符),无法组成批次。不同长度的序列无法批量处理,因此 batch_size=1。
Q2:测试三种模型的流程是什么?
A:①定义变量(输入维度、隐藏状态、输出维度);②加载数据;③创建数据集对象;④创建数据加载器;⑤模型初始化(RNN、LSTM、GRU)。
Q3:如何使用变量替代硬编码的数字?
A:定义 n_letters=57 和 category_num=18,然后用 input_size=n_letters 和 output_size=category_num 替代直接写数字,提高代码可读性。
Q4:数据加载器的作用是什么?
A:数据加载器按批次获取数据,支持 shuffle 打乱数据顺序,方便模型训练。
8 本课小结
- 测试流程:定义变量 → 加载数据 → 创建数据集 → 创建数据加载器 → 模型初始化。
- 批次大小:
batch_size=1(人名长度不一致)。 - 使用真实数据:从文件加载,而非随机生成。
- 变量替代数字:
n_letters、category_num。 - 三个模型:RNN、LSTM、GRU 同时初始化。
9 延伸思考
- 如何处理不同长度的序列以支持批次训练?
- 数据加载器的
shuffle参数对训练有什么影响?