RNN 人名分类案例 - 构建数据集对象
1 课程概览
本课讲解 RNN 人名分类案例的构建数据集对象。自定义 NameClassDataset 类,继承 Dataset,实现 __init__、__len__、__getitem__ 三个方法。在 __getitem__ 中进行 one-hot 编码和张量转换,并进行索引边界校验。
2 核心概念与定义
- NameClassDataset:自定义数据集类,继承 Dataset。
__init__:初始化函数,接收样本和标签数据。__len__:获取样本总数。__getitem__:根据索引获取样本,进行 one-hot 编码。- 索引边界校验:使用 min-max 组合确保索引合法。
3 模型与算法详解
自定义 Dataset 流程
继承 Dataset,实现三个方法。
1. __init__: 初始化,接收样本和标签
2. __len__: 返回样本总数
3. __getitem__: 根据索引获取样本,进行 one-hot 编码
NameClassDataset 类
class NameClassDataset(Dataset):
def __init__(self, my_list_x, my_list_y):
...
def __len__(self):
...
def __getitem__(self, index):
...
索引边界校验
使用 min-max 组合确保索引合法。
index = min(max(index, 0), self.sample_len - 1)
max(index, 0):确保索引不小于 0min(..., self.sample_len - 1):确保索引不大于最大值
one-hot 编码
在
__getitem__中进行。
tensor_x = one_hot(name) # 人名的 one-hot 编码
tensor_y = tensor(country) # 国家的张量表示
4 数学原理与推导
索引边界校验
$$\text{index} = \min(\max(\text{index}, 0), N - 1)$$
其中:
- $N$ 是样本总数
- $\max(\text{index}, 0)$ 确保索引不小于 0
- $\min(..., N - 1)$ 确保索引不大于最大值
one-hot 编码
$$\text{tensor_x} = \text{one_hot}(\text{name})$$
$$\text{tensor_y} = \text{country_to_index}(\text{country})$$
5 代码示例
import torch
from torch.utils.data import Dataset
class NameClassDataset(Dataset):
"""人名分类数据集"""
def __init__(self, my_list_x, my_list_y, all_letters, all_countries):
"""
初始化函数
Args:
my_list_x: 样本数据列表(人名)
my_list_y: 标签数据列表(国家)
all_letters: 全局字母表
all_countries: 国家列表
"""
# 1. 调用父类初始化
super(NameClassDataset, self).__init__()
# 2. 存储样本数据列表
self.my_list_x = my_list_x
# 3. 存储标签数据列表
self.my_list_y = my_list_y
# 4. 计算样本总数并存储
self.sample_len = len(my_list_x)
# 5. 存储字母表和国家列表
self.all_letters = all_letters
self.all_countries = all_countries
self.n_letters = len(all_letters)
self.n_countries = len(all_countries)
def __len__(self):
"""获取样本总数"""
return self.sample_len
def __getitem__(self, index):
"""
根据索引获取样本
并进行 one-hot 编码和张量转换
Args:
index: 样本索引
Returns:
tensor_x: 人名的 one-hot 编码张量
tensor_y: 国家的张量表示
"""
# 1. 索引边界校验
index = min(max(index, 0), self.sample_len - 1)
# 2. 获取人名和国家
name = self.my_list_x[index]
country = self.my_list_y[index]
# 3. one-hot 编码
tensor_x = torch.zeros(len(name), self.n_letters)
for i, letter in enumerate(name):
tensor_x[i][self.all_letters.find(letter)] = 1
# 4. 国家转索引
tensor_y = torch.tensor(self.all_countries.index(country))
return tensor_x, tensor_y
# 测试
if __name__ == "__main__":
print("=== RNN 人名分类案例 - 构建数据集对象 ===")
import string
all_letters = string.ascii_letters + " .,;'"
all_countries = ["Chinese", "English", "German", "French"]
my_list_x = ["Zhang", "Smith", "Muller", "Martin"]
my_list_y = ["Chinese", "English", "German", "French"]
dataset = NameClassDataset(my_list_x, my_list_y, all_letters, all_countries)
print(f"样本总数: {len(dataset)}")
tensor_x, tensor_y = dataset[0]
print(f"人名张量形状: {tensor_x.shape}")
print(f"国家索引: {tensor_y}")
代码说明
| 代码 | 说明 |
|---|---|
super(NameClassDataset, self).__init__() | 调用父类初始化 |
self.sample_len = len(my_list_x) | 计算样本总数 |
min(max(index, 0), self.sample_len - 1) | 索引边界校验 |
torch.zeros(len(name), self.n_letters) | one-hot 张量 |
self.all_letters.find(letter) | 字符转索引 |
self.all_countries.index(country) | 国家转索引 |
6 重难点与易错提醒
- ❗重点:自定义 Dataset 需要继承 Dataset 并实现三个方法。
- ❗重点:在
__getitem__中进行 one-hot 编码。 - ❗重点:使用 min-max 组合进行索引边界校验。
- ⚠️易错:
max(index, 0)在内,min(..., N-1)在外。
7 课堂问答精选
Q1:如何构建数据集对象?
A:自定义 NameClassDataset 类,继承 Dataset,实现 __init__、__len__、__getitem__ 三个方法。在 __getitem__ 中进行 one-hot 编码和张量转换。
Q2:如何进行索引边界校验?
A:使用 min-max 组合:index = min(max(index, 0), self.sample_len - 1),确保索引在 0 到 sample_len-1 之间。
Q3:在哪个方法中进行 one-hot 编码?
A:在 __getitem__ 方法中进行 one-hot 编码和张量转换。
8 本课小结
- 自定义 NameClassDataset 类,继承 Dataset。
- 实现三个方法:
__init__、__len__、__getitem__。 - 在
__getitem__中进行 one-hot 编码。 - 使用 min-max 组合进行索引边界校验。
9 延伸思考
- 如何构建数据加载器?
- 如何分批次获取数据?