RNN 人名分类案例 - 读取数据到内存
1 课程概览
本课讲解 RNN 人名分类案例的读取数据到内存。读取原数据文件,将人名和国家分别存储到两个列表中。过滤长度小于等于 5 的无效数据。
2 核心概念与定义
- read_data 函数:读取原数据到内存。
- 过滤无效数据:整行长度小于等于 5 的数据过滤掉。
- my_list_x:存储人名(特征)。
- my_list_y:存储国家名(标签)。
3 模型与算法详解
读取数据流程
1. 创建两个列表存储人名和国家名
2. 关联文件并逐行读取
3. 过滤无效数据(长度 <= 5)
4. 按制表符分割
5. 添加到列表
6. 返回列表
数据格式
每行数据格式。
人名\t国家
- 中间是制表符
\t,不是空格 - 每行末尾有换行符
\n
过滤无效数据
整行长度小于等于 5 的过滤掉。
if len(line) <= 5:
continue
continue vs break
| 语句 | 作用 |
|---|---|
| continue | 跳过本次循环,进入下次循环 |
| break | 终止循环 |
4 数学原理与推导
数据分割
$$\text{name}, \text{country} = \text{line}.\text{split}(\text{'\t'})$$
其中:
- $\text{line}$ 是每行数据
- $\text{name}$ 是人名
- $\text{country}$ 是国家名
5 代码示例
def read_data(file_path):
"""
读取原数据到内存中
并把特征(人名)和标签(国家)分别存储到两个列表中
Args:
file_path: 原数据文件路径
Returns:
my_list_x: 存储人名的列表
my_list_y: 存储国家名的列表
"""
# 1. 创建两个列表,分别存储人名和国家名
my_list_x, my_list_y = [], []
# 2. 关联文件并逐行读取
with open(file_path, 'r', encoding='utf-8') as f:
# 3. 遍历获取每一行数据
for line in f.readlines():
# 4. 过滤无效数据(整行长度 <= 5)
if len(line) <= 5:
continue
# 5. 按制表符分割
# 注意:中间是 \t,不是空格
# 每行末尾有 \n,需要去掉
name, country = line.strip().split('\t')
# 6. 添加到列表
my_list_x.append(name)
my_list_y.append(country)
return my_list_x, my_list_y
# 测试
if __name__ == "__main__":
print("=== RNN 人名分类案例 - 读取数据到内存 ===")
file_path = "./data/names_classification.txt"
my_list_x, my_list_y = read_data(file_path)
print(f"人名数量: {len(my_list_x)}")
print(f"国家数量: {len(my_list_y)}")
print(f"前 5 个人名: {my_list_x[:5]}")
print(f"前 5 个国家: {my_list_y[:5]}")
代码说明
| 代码 | 说明 |
|---|---|
with open(...) | 关联文件,自动释放 |
encoding='utf-8' | UTF-8 编码 |
f.readlines() | 逐行读取 |
line.strip() | 去掉换行符 |
split('\t') | 按制表符分割 |
if len(line) <= 5: continue | 过滤无效数据 |
6 重难点与易错提醒
- ❗重点:中间是制表符
\t,不是空格。 - ❗重点:每行末尾有换行符
\n,需要用strip()去掉。 - ❗重点:过滤整行长度小于等于 5 的无效数据。
- ⚠️易错:
continue是跳过本次循环,break是终止循环。
7 课堂问答精选
Q1:如何读取数据到内存?
A:使用 with open(...) 关联文件,逐行读取,按制表符分割,将人名和国家分别存储到两个列表中。
Q2:为什么要过滤无效数据?
A:如果人名长度太短(整行长度 <= 5),可能无法区分是哪个国家,所以过滤掉。
Q3:continue 和 break 的区别?
A:continue 是跳过本次循环,进入下次循环;break 是终止循环。
8 本课小结
- 读取原数据到内存,将人名和国家分别存储到两个列表。
- 过滤整行长度小于等于 5 的无效数据。
- 中间是制表符
\t,不是空格。 - 每行末尾有换行符
\n,需要用strip()去掉。
9 延伸思考
- 如何构建数据集对象?
- 如何实现 one-hot 编码?