Datasets 数据集

datasets 库与 Transformers 配套:从 Hub 拉数据、用 map 做分词、再交给 Trainer 或 DataLoader。本章用公开的 rotten_tomatoes(烂番茄影评,二分类),与 微调训练 共用。


load_dataset

from datasets import load_dataset

ds = load_dataset("rotten_tomatoes")
print(ds)
print(ds["train"][0])

典型结构:

DatasetDict({
  train: Dataset({ features: ['text', 'label'], num_rows: 8530 })
  validation: Dataset({ ... num_rows: 1066 })
  test: Dataset({ ... num_rows: 1066 })
})

label0 / 1(负面 / 正面)。指定子集:load_dataset("rotten_tomatoes", split="train[:500]") 便于 CPU 试跑。

其他常见来源:Hub 上的 org/name、本地 data/*.json、CSV。门控数据集需 hf auth login 并在网页同意条款。


查看特征与标签

print(ds["train"].features)
print(ds["train"].features["label"].names)  # 若为 ClassLabel
print(ds["train"].unique("label"))

写训练脚本前先确认:文本列名(这里是 text)和 标签列名label)。换数据集时这两处最容易写错。


map:批量分词

from transformers import AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained("distilbert/distilbert-base-uncased")

def tokenize(batch):
    return tokenizer(batch["text"], truncation=True)

tokenized = ds.map(tokenize, batched=True, remove_columns=["text"])
tokenized = tokenized.rename_column("label", "labels")
print(tokenized["train"][0].keys())

建议:

  • batched=True:一次处理多条,远快于逐条 Python 循环。
  • truncation=True:截到模型 model_max_length,避免超长影评炸显存。
  • 动态 padding 交给 DataCollator,不要在整个数据集上 padding="max_length"(除非你有意对齐到固定长)。
  • Trainerlabels 列;有的版本也能认 label,显式改名更稳。

DataCollator:组 batch

from transformers import DataCollatorWithPadding
from torch.utils.data import DataLoader

collator = DataCollatorWithPadding(tokenizer=tokenizer)
loader = DataLoader(
    tokenized["train"].with_format("torch"),
    batch_size=8,
    collate_fn=collator,
)
batch = next(iter(loader))
print({k: v.shape for k, v in batch.items()})

DataCollatorWithPadding当前 batch 内最长句补 pad,比全局固定长度省算力。生成任务常用 DataCollatorForLanguageModelingDataCollatorForSeq2Seq

张量形状([batch, seq])的含义见本站 PyTorch 教程


设格式、缓存与小样本

small = tokenized["train"].shuffle(seed=42).select(range(256))
small.set_format(type="torch", columns=["input_ids", "attention_mask", "labels"])

map 的结果会缓存到磁盘,改了函数要 load_from_cache_file=False 或换 cache_file_name。评测指标不要写在本章的 map 里,用 evaluate 库在 Trainercompute_metrics 中计算,见下一章。


和 Hub 数据集页

每个数据集都有介绍页(例如 rotten_tomatoes):许可、字段、引用。上线前读许可;仅学术可用的语料不要直接塞进商业产品。

也可以 ds.push_to_hub("your-name/my-reviews") 发布自己的预处理版本(需登录)。


下一步

评论