第 6 章 微调数据准备
学习目标
- 能把原始文本数据集变成「模型能吃的」token 数据集
- 掌握
map(..., batched=True, remove_columns=...)的标准写法 - 掌握
DataCollatorWithPadding的作用 - 掌握
labels列与id2label/label2id映射
6.1 从文本到输入:一条完整链路
模型需要的是 input_ids、attention_mask、labels 这样的张量,而不是字符串。微调前必须把原始数据集处理成:
原始列: sentence, label
↓ tokenize
处理列: input_ids, attention_mask, labels标准做法:用 map 批量分词,然后删除原始列,只留模型需要的列。
from datasets import load_dataset
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained("google-bert/bert-base-uncased")
ds = load_dataset("nyu-mll/glue", "sst2", split="train[:4]")
def tokenize(examples):
return tokenizer(examples["sentence"], truncation=True, max_length=128)
tokenized = ds.map(tokenize, batched=True, remove_columns=ds.column_names)
print("处理后的列:", list(tokenized.features.keys()))
print("input_ids 长度:", len(tokenized[0]["input_ids"]))输出:
处理后的列: ['input_ids', 'token_type_ids', 'attention_mask']
input_ids 长度: 128关键点:
batched=True:一批一批分词,速度快几个数量级。truncation=True, max_length=128:控制序列上限,防止过长样本爆显存。remove_columns=ds.column_names:删掉原始列——这不是为了省空间,而是为了不让字符串列漏进训练。
6.2 不删原始列会怎样
如果只分词不删列,DataCollator 会把整行(含 sentence 字符串)当输入:
ValueError: Unable to create tensor, you should probably activate truncation and/or padding with 'padding=True' 'truncation=True' to have batched tensors with the same length. Perhaps your features (`text` in this case) have excessive nesting (inputs type `list` where type `int` is expected).报错的关键句:Perhaps your features (...) have excessive nesting——字符串列混进了批处理。解决办法就是 remove_columns。
6.3 labels 列
分类任务的标签列要求叫 labels(复数)。SST-2 的标签列叫 label,需要改名:
def tokenize(examples):
return tokenizer(examples["sentence"], truncation=True, max_length=128)
tokenized = ds.map(tokenize, batched=True, remove_columns="sentence")
tokenized = tokenized.rename_column("label", "labels")
print(list(tokenized.features.keys()))输出:
['labels', 'input_ids', 'token_type_ids', 'attention_mask']Trainer 会在前向时把 labels 自动传给损失函数,列名必须是 labels。
6.4 DataCollator:把样本拼成 batch
DataCollator 负责把一批样本「整理」成模型输入:对齐长度、补 [PAD]、转成张量。最常用的是 DataCollatorWithPadding,它按 batch 内最长序列动态 padding(而不是全局 padding,省显存):
from transformers import DataCollatorWithPadding
collator = DataCollatorWithPadding(tokenizer=tokenizer, return_tensors="pt")
batch = collator([tokenized[i] for i in range(4)])
print("input_ids shape:", tuple(batch["input_ids"].shape))
print("attention_mask shape:", tuple(batch["attention_mask"].shape))输出:
input_ids shape: (4, 128)
attention_mask shape: (4, 128)形状 (4, 128) = 4 条样本 × 128 个 token,全部对齐。
6.5 id2label / label2id:让类别可读
模型输出的是编号,报告时需要名称。在构造模型 config 时声明映射:
from transformers import BertConfig, BertForSequenceClassification
config = BertConfig.from_pretrained(
"prajjwal1/bert-tiny",
num_labels=2,
id2label={0: "negative", 1: "positive"},
label2id={"negative": 0, "positive": 1},
)
model = BertForSequenceClassification.from_pretrained(
"prajjwal1/bert-tiny", config=config)
print(model.config.id2label)输出:
{0: 'negative', 1: 'positive'}没有这套映射,推理时 pipeline 只能输出 LABEL_0、LABEL_1(第 8 章会看到真实例子)。
6.6 规模选择
微调起步阶段用小数据跑通流程,确认无 bug 后再放大:
small_train = load_dataset("nyu-mll/glue", "sst2", split="train[:200]")
small_eval = load_dataset("nyu-mll/glue", "sst2", split="validation[:100]")先 200 条训练、100 条验证跑通,再换 train[:2000] 看趋势,最后才全量。这条原则能帮你省下大量调试时间。
动手实践
- 对 SST-2 的
train[:1000]做完整预处理:分词、删列、改名labels、DataCollatorWithPadding。 - 打印一个 collate 后的 batch,人工核对:每行的
attention_mask末尾 0 的个数 = 该行[PAD]个数。 - 构造 config 时故意不写
id2label,训练后预测,观察输出变成什么(记下现象,第 8 章解释)。
常见错误
错误 1:忘了 remove_columns,报 "too many dimensions 'str'"。
解决:在 map 里删掉原始文本列,或只保留需要的列。
错误 2:标签列名不是 labels。
Trainer 找不到 labels 列时,训练会报错或直接忽略标签。检查 column_names。
错误 3:只设置 truncation=True 不设 max_length。
会截断到模型默认最大长度;若你的任务文本很长,记得显式设 max_length,并统计截断比例。
章末练习
基础
- 写出微调数据准备的完整代码:SST-2
train[:200]→ 分词 → 删列 → 改名 → collator。 - 打印预处理后数据集的列名,确认没有字符串列残留。
提高
- 统计 SST-2
train[:2000]中,截断到 64 token 的样本占比(提示:比较分词前字符长度与分词后长度)。 - 解释为什么
DataCollatorWithPadding比全局padding="max_length"省显存。
挑战
- 写一个
prepare(split, tokenizer, max_length)函数,返回预处理后的 DatasetDict(训练/验证),并在函数内断言:不含字符串列、列名含labels、样本数正确。
章末自测
- 预处理时
remove_columns的作用是什么? - 分类任务中标签列的标准名称是什么?
DataCollatorWithPadding的作用是什么?id2label的作用是什么?- 判断:
batched=True只影响速度,不影响结果。 - 预处理后的
attention_mask中,[PAD]位置对应的值是 0 还是 1?
