第 5 章 datasets 数据集
学习目标
- 掌握
load_dataset的加载语法与 Dataset/DatasetDict 结构 - 掌握
map/filter/select等变换操作 - 掌握
train_test_split与显式Features类型 - 了解流式加载与本地持久化
5.1 load_dataset:一行代码加载数据集
datasets 库提供 load_dataset(repo_id, config, split) 接口,从 Hub 下载并解析数据。
from datasets import load_dataset
dataset = load_dataset("nyu-mll/glue", "sst2", split="train[:10]")
print(type(dataset).__name__)
print(dataset)
print("features:", dataset.features)
print("row0:", dataset[0])输出:
Dataset
Dataset({
features: ['sentence', 'label', 'idx'],
num_rows: 10
})
features: {'sentence': Value('string'), 'label': ClassLabel(names=['negative', 'positive']), 'idx': Value('int32')}
row0: {'sentence': 'hide new secretions from the parental units ', 'label': 0, 'idx': 0}解读:
"nyu-mll/glue"是数据集仓库 id,"sst2"是该仓库的一个 config(子任务,斯坦福情感二分类)。split="train[:10]"表示只取训练集前 10 条——先小样本试代码,再放大。features声明了每列类型:sentence是字符串,label是带名字的类别(negative/positive)。
不写 split 时返回 DatasetDict(train / validation / test 三个子集):
full = load_dataset("nyu-mll/glue", "sst2")
print(full.keys())
print(full["train"].num_rows)5.2 手工构造 Dataset
from datasets import Dataset
manual = Dataset.from_dict({
"text": ["good", "bad", "okay"],
"label": [1, 0, 1],
})
print(manual.to_dict())输出:
{'text': ['good', 'bad', 'okay'], 'label': [1, 0, 1]}Dataset.from_dict 从 Python 字典创建数据集;也可以用 from_list。手工构造对单元测试和 TRL 小实验(第 11 章)非常有用。
5.3 map / filter / select
map 对每一行(或一批行)应用函数,返回新列:
upper = manual.map(lambda x: {"upper": x["text"].upper()})
print(upper.to_dict())输出:
{'text': ['good', 'bad', 'okay'], 'label': [1, 0, 1], 'upper': ['GOOD', 'BAD', 'OKAY']}filter 按条件保留行:
positive = manual.filter(lambda x: x["label"] == 1)
print(positive["text"])输出:
Column(['good', 'okay'])select 按下标取行:
print(manual.select([0, 2])["text"])输出:
Column(['good', 'okay'])map 加 batched=True 可一次处理一批(第 6 章批量分词的关键)。
5.4 划分训练集与验证集
split = manual.train_test_split(test_size=0.3, seed=42)
print(split.keys())
print("train:", len(split["train"]), "| test:", len(split["test"]))输出:
dict_keys(['train', 'test'])
train: 2 | test: 1train_test_split 返回 DatasetDict,seed 保证切分可复现。
5.5 显式 Features:类型即文档
from datasets import Features, ClassLabel, Value
features = Features({
"text": Value("string"),
"label": ClassLabel(names=["neg", "pos"]),
})
typed = Dataset.from_dict({"text": ["bad", "good"], "label": [0, 1]},
features=features)
print(typed.features)
print(typed.features["label"].int2str(1))输出:
{'text': Value('string'), 'label': ClassLabel(names=['neg', 'pos'])}
posClassLabel 让编号与名称互相转换:int2str(1) → "pos";训练时模型输出编号,报告时转成名称。
5.6 流式加载:大数据集不求全
stream = load_dataset("nyu-mll/glue", "sst2", split="train", streaming=True)
it = iter(stream)
row = next(it)
print(list(row.keys()))输出:
['sentence', 'label', 'idx']streaming=True 时数据集不会全部下载,而是按需拉取——适合超大语料(如网络爬虫数据)。
5.7 本地持久化
import tempfile, os
tmp = tempfile.mkdtemp()
typed.save_to_disk(os.path.join(tmp, "ds"))
back = Dataset.load_from_disk(os.path.join(tmp, "ds"))
print("行内容一致:", back["text"] == typed["text"])输出:
行内容一致: Truesave_to_disk 把数据集(含 features)存到磁盘,load_from_disk 原样读回,适合缓存预处理结果。
动手实践
- 加载
nyu-mll/glue的sst2,统计训练集与验证集的类别分布(collections.Counter)。 - 对
manual数据集写一个map:新增一列text_len = len(text)。 - 用
select取出训练集前 3 条,再和filter结果对比。
常见错误
错误 1:datasets 5 要求仓库 id 带命名空间。
旧文档里的 load_dataset("imdb") 在 v5 会报:
huggingface_hub.errors.HfUriError: Repository id must be 'namespace/name', got 'imdb'.必须用完整仓库 id(如 stanfordnlp/imdb、nyu-mll/glue)。有些老仓库数据本身有质量问题,加载后先检查 features 与标签分布,再决定是否使用。
错误 2:split 语法写错。
split="train[:10]" 是切片,split="train[10:]" 是跳过前 10 条,split="train" 是全部。写错下标范围会得到空数据集——先打印 num_rows 确认。
错误 3:map 的函数作用域。
map 的函数必须返回 dict;若函数内部用了外部大对象,注意它是被序列化执行的,复杂闭包可能报错。保持函数「纯」:输入一行,输出新列。
章末练习
基础
- 加载 SST-2 训练集,打印 features 与前 3 条 sentence。
- 对
Dataset.from_dict({"a": [1, 2, 3]})用map新增一列b = a * 10。
提高
- 用
train_test_split(test_size=0.2, seed=7)划分 SST-2 前 1000 条,验证 train/test 行数之和等于 1000。 - 用
filter找出 SST-2 验证集中长度(字符数)小于 20 的句子,打印数量。
挑战
- 写一个函数
describe(ds):打印行数、列名、每列类型、label 分布;用它对 SST-2 训练集调用,输出一份结构化摘要。
章末自测
load_dataset("nyu-mll/glue", "sst2")返回的对象类型是?Dataset.map的返回值是什么?filter(lambda x: ...)的作用是什么?train_test_split返回什么类型?streaming=True与默认加载的区别是什么?ClassLabel.names字段存放什么?
