Skip to content

第 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 下载并解析数据。

python
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 三个子集):

python
full = load_dataset("nyu-mll/glue", "sst2")
print(full.keys())
print(full["train"].num_rows)

5.2 手工构造 Dataset ​

python
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 对每一行(或一批行)应用函数,返回新列:

python
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 按条件保留行:

python
positive = manual.filter(lambda x: x["label"] == 1)
print(positive["text"])

输出:

Column(['good', 'okay'])

select 按下标取行:

python
print(manual.select([0, 2])["text"])

输出:

Column(['good', 'okay'])

map 加 batched=True 可一次处理一批(第 6 章批量分词的关键)。

5.4 划分训练集与验证集 ​

python
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: 1

train_test_split 返回 DatasetDict,seed 保证切分可复现。

5.5 显式 Features:类型即文档 ​

python
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'])}
pos

ClassLabel 让编号与名称互相转换:int2str(1) → "pos";训练时模型输出编号,报告时转成名称。

5.6 流式加载:大数据集不求全 ​

python
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 本地持久化 ​

python
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"])

输出:

行内容一致: True

save_to_disk 把数据集(含 features)存到磁盘,load_from_disk 原样读回,适合缓存预处理结果。

动手实践 ​

  1. 加载 nyu-mll/glue 的 sst2,统计训练集与验证集的类别分布(collections.Counter)。
  2. 对 manual 数据集写一个 map:新增一列 text_len = len(text)。
  3. 用 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;若函数内部用了外部大对象,注意它是被序列化执行的,复杂闭包可能报错。保持函数「纯」:输入一行,输出新列。

章末练习 ​

基础

  1. 加载 SST-2 训练集,打印 features 与前 3 条 sentence。
  2. 对 Dataset.from_dict({"a": [1, 2, 3]}) 用 map 新增一列 b = a * 10。

提高

  1. 用 train_test_split(test_size=0.2, seed=7) 划分 SST-2 前 1000 条,验证 train/test 行数之和等于 1000。
  2. 用 filter 找出 SST-2 验证集中长度(字符数)小于 20 的句子,打印数量。

挑战

  1. 写一个函数 describe(ds):打印行数、列名、每列类型、label 分布;用它对 SST-2 训练集调用,输出一份结构化摘要。

章末自测 ​

  1. load_dataset("nyu-mll/glue", "sst2") 返回的对象类型是?
  2. Dataset.map 的返回值是什么?
  3. filter(lambda x: ...) 的作用是什么?
  4. train_test_split 返回什么类型?
  5. streaming=True 与默认加载的区别是什么?
  6. ClassLabel.names 字段存放什么?