数据截至 (上游 commit 67dfbe211a07)
02 · SFT:数据管线、掩码与 packing
这一章讲什么:
SFTTrainer是 TRL 里最「简单」的 Trainer——loss 就是下一个 token 预测。但它的数据管线是全家最讲究的:怎么把三种格式的数据变成input_ids + labels、怎么只给回答部分算 loss、怎么用装箱(packing)消灭 padding。读完你会理解为什么「SFT 谁都会写」和「写好一个 SFT 管线」之间差着两千行代码。
1. 它要解决的小问题
SFT 的算法一句话说完:对示范数据做语言建模 loss。但工程上有四个真问题:
- 数据格式五花八门——有的是整篇文本,有的是「prompt + completion」两半,有的是对话消息列表。
- 只想学「回答」——prompt 部分不该贡献 loss,否则模型浪费容量去模仿提问。
- 序列长短不齐——直接 padding 会让 GPU 大量时间花在补零上。
- 词表巨大——logits 张量(batch × 序列 × 词表)是 SFT 的显存大头。
TRL 把这四个问题分别解掉:统一管线(§2)、掩码(§3)、packing(§4-5)、分块交叉熵(§6)。
2. 三种数据格式,一条管线
SFT 接受三种输入格式(判定的核心是 is_conversational,trl/data_utils.py:160——看值是不是消息列表):
| 格式 | 列 | loss 默认作用于 |
|---|---|---|
| 语言建模 | text(整篇文本) | 全部 token |
| prompt-completion | prompt + completion | completion(见 §3) |
| 对话 | messages([{"role": ..., "content": ...}]) | 全部 token,或仅 assistant(见 §3) |
整条管线在 _prepare_dataset(trl/trainer/sft_trainer.py:1424)里,是一串 dataset.map:
原始数据集
│ ① formatting_func(可选)→ 统一出 "text" 列
│ ② maybe_convert_to_chatml:旧格式({"conversations": [...]})转成标准消息
│ ③ add_eos:非对话数据在末尾补 EOS(trl/trainer/sft_trainer.py:1488)
│ ④ tokenize_fn:tokenize + 产出掩码列(trl/trainer/sft_trainer.py:1506)
│ ⑤ build_labels:掩码折进 labels,不学的位置写 -100(sft_trainer.py:1605)
│ ⑥ truncate:截到 max_length,丢掉被截成全 -100 的样本
│ ⑦ pack_dataset:装箱(可选,§4)
▼
训练就绪: input_ids / labels / (seq_lengths)
两个容易忽略的细节:
- 第 ④ 步对 prompt-completion 数据做了一次「前缀校验」:分别 tokenize
prompt和prompt+completion,然后检查前者确实是后者的前缀(trl/trainer/sft_trainer.py:1543-1549)。不是就打 warning——tokenizer 在拼接处的行为不一致是静默挂掉的高发区。 - 第 ⑥ 步会丢掉「prompt 就占满 max_length」的样本(
trl/trainer/sft_trainer.py:1643-1646):这种样本的 labels 全是 -100,留下来只会贡献零 loss 和噪声指标。
3. 只学回答:两种掩码
「不给 prompt 算 loss」由两个开关控制,可以叠加:
| 开关 | 作用对象 | 掩码来源 |
|---|---|---|
completion_only_loss | prompt-completion 数据 | tokenize 时算出的 completion_mask |
assistant_only_loss | 对话数据 | chat template 的 {% generation %} 标记产出 assistant_masks |
completion_mask 的构造直白(trl/trainer/sft_trainer.py:1551):prompt 段标 0、其余标 1。
assistant_masks 则依赖 chat template 支持「代标记」(generation keyword)——如果模板不支持,assistant_only_loss=True 会直接报错并提示检查模板(trl/trainer/sft_trainer.py:1573-1578)。
最后 build_labels 把掩码折进 labels——所有掩码位取与,任一掩码为 0 就写 -100:
# 摘自 trl/trainer/sft_trainer.py:1605-1612(build_labels)
labels = [
token_id if all(bits) else -100
for token_id, *bits in zip(example["input_ids"], *masks, strict=False)
]
-100 是 PyTorch 交叉熵的 ignore_index。所以到了训练侧,loss 就是标准的因果 LM loss,没有任何定制——掩码的全部魔法都在数据侧完成。这是个值得抄的分层:让「学什么」留在数据里,让训练循环保持 dumb。
4. Packing:把 padding 消灭掉
4.1 小问题
batch 里序列长短不一,补零到最长会浪费算力。SFT 数据(尤其多轮对话)长度方差极大,浪费可达数倍。