跳到主要内容

数据截至 (上游 commit 67dfbe211a07)

02 · SFT:数据管线、掩码与 packing

这一章讲什么: SFTTrainer 是 TRL 里最「简单」的 Trainer——loss 就是下一个 token 预测。但它的数据管线是全家最讲究的:怎么把三种格式的数据变成 input_ids + labels、怎么只给回答部分算 loss、怎么用装箱(packing)消灭 padding。读完你会理解为什么「SFT 谁都会写」和「写好一个 SFT 管线」之间差着两千行代码。


1. 它要解决的小问题

SFT 的算法一句话说完:对示范数据做语言建模 loss。但工程上有四个真问题:

  1. 数据格式五花八门——有的是整篇文本,有的是「prompt + completion」两半,有的是对话消息列表。
  2. 只想学「回答」——prompt 部分不该贡献 loss,否则模型浪费容量去模仿提问。
  3. 序列长短不齐——直接 padding 会让 GPU 大量时间花在补零上。
  4. 词表巨大——logits 张量(batch × 序列 × 词表)是 SFT 的显存大头。

TRL 把这四个问题分别解掉:统一管线(§2)、掩码(§3)、packing(§4-5)、分块交叉熵(§6)。


2. 三种数据格式,一条管线

SFT 接受三种输入格式(判定的核心是 is_conversationaltrl/data_utils.py:160——看值是不是消息列表):

格式loss 默认作用于
语言建模text(整篇文本)全部 token
prompt-completionprompt + completioncompletion(见 §3)
对话messages[{"role": ..., "content": ...}]全部 token,或仅 assistant(见 §3)

整条管线在 _prepare_datasettrl/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 promptprompt+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_lossprompt-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 数据(尤其多轮对话)长度方差极大,浪费可达数倍。

4.2 思路:装箱问题

把定长训练序列当「桶」、样本当「物品」,就是经典的装箱问题(bin packing)。TRL 用 Best Fit Decreasing(BFD,最佳适配递减)启发式:先按长度降序排,每个样本放进「剩余空间刚好够、且最紧」的桶。

packing_strategy 有三种(trl/trainer/sft_config.py:228-235):

策略超长样本是否跨样本截断适用
bfd(默认)截断不跨样本SFT/对话,保样本边界
bfd_split切成多段继续装切但不丢 token预训练式长文档
wrapped直接拼接后等长切随意切断最快,不管边界

4.3 原理演示

# 示意,非源码:BFD 在做什么
samples = sorted(lengths, reverse=True) # 先降序
bins = [] # 每个 bin 容量 = seq_length
for s in samples:
b = find_bin_with_tightest_fit(s) # 找剩余空间 ≥ s 且最小的桶
if b is None:
b = new_bin()
b.add(s) # 桶里记下样本 id 与已用长度
# 输出: 每条训练序列 = 若干样本拼接 + seq_lengths 记录边界

4.4 真实实现:线段树加速的 _pack_bfd

「找最紧的桶」如果每次线性扫是 O(n²)。TRL 用一棵线段树_SegmentTreetrl/data_utils.py:694)把「找 ≥ length 的最小剩余空间」降到 O(log n),注释里直接给了出处论文(Fewer Truncations Improve Language Modeling)。核心循环在 _pack_bfdtrl/data_utils.py:739):

# 摘自 trl/data_utils.py:790-797(BFD 主循环)
for length, idx in zip(lengths.field(0).to_numpy(), lengths.field(1).to_numpy(), strict=True):
space = segment_tree.search(length) # 最紧的可用剩余空间
if space < seq_length:
bin = space_to_bin[space].popleft() # 复用已有桶
else:
bin = {"ids": [], "length": 0} # 开新桶
bins.append(bin)
bin["ids"].append(idx)
bin["length"] += length
...

输出除了拼接后的列,还多一列 seq_lengths(每条训练序列里各样本的长度),它是后面「不让注意力跨样本泄漏」的关键。

整个功能对外就是 pack_datasettrl/data_utils.py:845),独立于 SFTTrainer 也可以直接 from trl import pack_dataset 用。


5. Padding-free:装了箱,干脆把 batch 维也拍平

5.1 小问题

装箱后每条训练序列已接近等长,但 batch 内仍有少量 padding;更重要的是,拼接后的样本之间不该互相做注意力——「样本 A 的结尾」不应该能注意到「样本 B 的开头」。

5.2 思路

FlashAttention 支持 varlen 模式:把整个 batch 拍平成一条超长序列,靠 position_ids(或 cu_seqlens)告诉 kernel 每段的边界,段与段之间注意力隔离。

5.3 真实实现

DataCollatorForLanguageModelingtrl/trainer/sft_trainer.py:400)在 padding_free=True 时把整个 batch 拼成一条,并用 seq_lengths 生成「每段从 0 重新开始」的 position_ids(get_position_ids_from_packed_seq_lengthstrl/trainer/sft_trainer.py:525):

# 摘自 trl/trainer/sft_trainer.py:538-546:用 cumsum 技巧批量生成分段 position_ids
position_ids = torch.ones(sum(example_lengths), dtype=batch_seq_lengths.dtype)
position_ids[0] = 0
position_ids[batch_seq_lengths[:-1].cumsum(0)] = -(batch_seq_lengths[:-1] - 1)
position_ids = position_ids.cumsum(0)

这个写法很妙:先全填 1,在每段起点放一个「修正值」,一次 cumsum 就得到 [0,1,2, 0,1, 0,1,2,3, ...]。collator 还有一行容易看漏的细节——padding-free 模式下把 position_ids == 0(每段第一个 token)的 label 也置为 -100(trl/trainer/sft_trainer.py:513),因为段首 token 的「上一个 token」属于另一个样本,不该被预测。

开启方式:SFT 里 packing=True 且策略为 bfd 时会自动启用 padding-free(trl/trainer/sft_config.py:82-88 的说明)。


6. 分块交叉熵:不物化全量 logits

6.1 小问题

最后一层隐藏态(batch×seq×d)乘 lm_head(d×vocab)得到的 logits 是 SFT 显存峰值——7B 模型、8k 序列、150k 词表,一个 batch 的 logits 就要几十 GB。

6.2 思路

沿序列维分块:每次只把一小段 hidden states 投到词表、算完交叉熵就扔掉 logits。loss 不需要完整的 logits 张量,只需要标量。

6.3 真实实现

本 commit 里这是默认路径SFTConfig.loss_type 默认 "chunked_nll"trl/trainer/sft_config.py:331-332:不用 Liger kernel 时默认 chunked)。_patch_chunked_ce_lm_headtrl/trainer/sft_trainer.py:233)替换模型的 forward,使其在拿到 hidden states 后走 _chunked_cross_entropy_losstrl/trainer/sft_trainer.py:119),每 256 个 token 一块(_CHUNKED_LM_HEAD_CHUNK_SIZEtrl/trainer/sft_trainer.py:89),逐块算 log_softmax + nll_chunktrl/trainer/sft_trainer.py:101),顺便把 token 准确率和熵也算出来(反正 logits 在手,不浪费)。

同族的省显存工具还有 selective_log_softmaxtrl/trainer/utils.py:480):log_softmax 后只 gather 需要的 token,DPO/GRPO/KTO 全在用它。

另外还有一个值得知道的 loss 变体:dft_losstrl/trainer/sft_trainer.py:788)实现 DFT 论文(2508.05629),把每个 token 的 loss 乘上它自己的概率(per_token_loss = -logprobs.exp().detach() * logprobs),一行之差把 SFT 改写成「奖励 = 1 的强化学习」形式。


7. 关键细节与坑

  • assistant_only_loss 依赖 chat template。 模板没有 {% generation %} 标记就直接 RuntimeError(trl/trainer/sft_trainer.py:1573-1578);换模型族时先确认模板支持。
  • vision 数据集限制多。 packing、padding-free、assistant-only loss、truncation_mode="keep_end" 对视觉数据全部禁用(trl/trainer/sft_trainer.py:1044-1067)——图像 token 住在 prompt 段里,这些机制会把它切坏。
  • with_transform 的数据集会被拒。 惰性 transform 会被 map() 固化进缓存(trl/trainer/sft_trainer.py:1434-1441);要么先 map 物化,要么 dataset_kwargs={"skip_prepare_dataset": True} 自己准备。
  • packing 必须先打乱。 _prepare_dataset 在装箱前 shuffle(trl/trainer/sft_trainer.py:1658-1660),否则同源样本总被装进同一条序列,引入相关性。
  • prompt 前缀不匹配的 warning 别忽略trl/trainer/sft_trainer.py:1543-1549)——它意味着 completion_mask 可能错位,模型在偷学 prompt 或漏学回答。

8. 代码地图(本章)

主题文件路径符号名
数据管线总入口trl/trainer/sft_trainer.pySFTTrainer._prepare_dataset
tokenize + 掩码trl/trainer/sft_trainer.pytokenize_fn(局部函数)、build_labels(局部函数)
对话格式判定trl/data_utils.pyis_conversationalmaybe_convert_to_chatml
装箱trl/data_utils.pypack_dataset_pack_bfd_pack_wrapped_SegmentTree
LM collator / padding-freetrl/trainer/sft_trainer.pyDataCollatorForLanguageModelingget_position_ids_from_packed_seq_lengths
分块交叉熵trl/trainer/sft_trainer.py_patch_chunked_ce_lm_head_chunked_cross_entropy_loss_chunk
DFT 损失trl/trainer/sft_trainer.pydft_loss
省显存 log-softmaxtrl/trainer/utils.pyselective_log_softmax
SFT 训练侧 losstrl/trainer/sft_trainer.pySFTTrainer.compute_loss