RNN 人名分类案例 - 构建数据加载器对象
1 课程概览
本课讲解 RNN 人名分类案例的构建数据加载器对象。基于 Dataset 创建 DataLoader,实现分批次获取数据。DataLoader 的参数包括 dataset、batch_size、shuffle。
2 核心概念与定义
- DataLoader:数据加载器,分批次获取数据。
- get_data_loader 函数:获取数据加载器对象。
- batch_size:批次大小。
- shuffle:是否打乱数据。
3 模型与算法详解
构建流程
Dataset → DataLoader。
1. 读取数据文件,获取样本列表和标签列表
2. 创建数据集对象(Dataset)
3. 创建数据加载器对象(DataLoader)
4. 测试数据加载器
get_data_loader 函数
def get_data_loader():
# 1. 读取数据文件
my_list_x, my_list_y = read_data(file_path)
# 2. 创建数据集对象
dataset = NameClassDataset(my_list_x, my_list_y, ...)
# 3. 创建数据加载器对象
dataloader = DataLoader(dataset, batch_size=1, shuffle=True)
# 4. 测试数据加载器
for x, y in dataloader:
print(x.shape, y)
break
return dataloader
DataLoader 参数
| 参数 | 说明 |
|---|---|
| dataset | 数据集对象 |
| batch_size | 批次大小(1) |
| shuffle | 是否打乱(True) |
优化
返回 DataLoader。
return my_data_loader
4 数学原理与推导
DataLoader
$$\text{DataLoader} = \text{DataLoader}(\text{Dataset}, \text{batch_size}, \text{shuffle})$$
分批次获取数据
$$\text{batch} = \text{DataLoader}.\text{next}()$$
5 代码示例
from torch.utils.data import DataLoader
def get_data_loader(file_path, all_letters, all_countries, batch_size=1, shuffle=True):
"""
获取数据加载器对象
Args:
file_path: 数据文件路径
all_letters: 全局字母表
all_countries: 国家列表
batch_size: 批次大小
shuffle: 是否打乱
Returns:
my_data_loader: 数据加载器对象
"""
# 1. 读取数据文件,获取样本列表和标签列表
my_list_x, my_list_y = read_data(file_path)
# 2. 创建数据集对象
name_class_dataset = NameClassDataset(
my_list_x, my_list_y, all_letters, all_countries
)
# 3. 创建数据加载器对象,用于批量加载和处理数据
my_data_loader = DataLoader(
dataset=name_class_dataset, # 数据集对象
batch_size=batch_size, # 批次大小
shuffle=shuffle # 是否打乱数据
)
# 4. 测试数据加载器,打印第一批数据的形状和内容
for x, y in my_data_loader:
print(f"人名张量形状: {x.shape}")
print(f"人名张量内容: {x}")
print(f"国家索引: {y}")
break # 仅打印第一批次数据,避免全部输出
# 5. 优化:返回数据加载器
return my_data_loader
# 测试
if __name__ == "__main__":
print("=== RNN 人名分类案例 - 构建数据加载器对象 ===")
import string
all_letters = string.ascii_letters + " .,;'"
all_countries = ["Chinese", "English", "German", "French"]
file_path = "./data/names_classification.txt"
my_data_loader = get_data_loader(file_path, all_letters, all_countries, batch_size=1, shuffle=False)
print(f"数据加载器: {my_data_loader}")
print(f"批次数量: {len(my_data_loader)}")
代码说明
| 代码 | 说明 |
|---|---|
read_data(file_path) | 读取数据文件 |
NameClassDataset(...) | 创建数据集对象 |
DataLoader(...) | 创建数据加载器 |
batch_size=1 | 批次大小为 1 |
shuffle=True | 打乱数据 |
for x, y in my_data_loader: | 遍历数据加载器 |
break | 仅打印第一批次 |
return my_data_loader | 返回数据加载器 |
6 重难点与易错提醒
- ❗重点:构建流程:Dataset → DataLoader。
- ❗重点:DataLoader 的参数包括 dataset、batch_size、shuffle。
- ❗重点:实际开发中应该返回 DataLoader,而不是只测试。
- 💡技巧:测试时使用
break避免全部输出。
7 课堂问答精选
Q1:如何构建数据加载器对象?
A:基于 Dataset 创建 DataLoader,参数包括 dataset、batch_size、shuffle。
Q2:DataLoader 的参数有哪些?
A:参数1:dataset(数据集对象);参数2:batch_size(批次大小);参数3:shuffle(是否打乱)。
Q3:为什么要返回 DataLoader?
A:返回 DataLoader 后,后续可以直接调用,不用重新封装。
8 本课小结
- 构建流程:Dataset → DataLoader。
- DataLoader 的参数:dataset、batch_size、shuffle。
- 实际开发中应该返回 DataLoader。
- 测试时使用
break避免全部输出。
9 延伸思考
- 如何搭建 RNN 模型?
- 如何测试 RNN 模型?