SFT
有监督微调 (Supervised Fine-Tuning, SFT) 用「输入 -> 目标输出」示范,把只擅长续写的 预训练 模型变成能够遵循指令、保持对话角色、输出固定格式和调用工具的策略。InstructGPT 展示了「SFT -> 奖励模型 -> 强化学习」的经典对齐流水线,其中 SFT 负责建立可用的初始助手策略。它仍使用 next-token 交叉熵,却把监督位置限制在希望模型学会生成的内容上。
SFT 的意义在于建立稳定的行为接口,而非向模型灌入所有知识。高质量 SFT 让后续 强化学习 有一个可用的初始策略,也让 蒸馏 的学生具备基本生成能力;如果模板、掩码或示范本身错误,后续优化通常只会放大这些问题。
数据流与符号约定¶
SFT 的数据流从聊天模板开始,再到数据和 loss mask:
首次实现不要直接从全量训练开始。随机抽取几十条样本,打印模板渲染后的文本、token、角色边界与 loss mask;运行若干步确认 loss 下降;最后用与上线相同的模板做生成。三项都正确后,再扩大数据与并行规模。
全文使用以下符号:
| 符号 | 含义 | 形状 |
|---|---|---|
| \(B\) | batch 中的对话数 | 标量 |
| \(S\) | padding 后的序列长度 | 标量 |
| \(V\) | 词表大小 | 标量 |
| \(X\) | 完整对话 token id | \([B,S]\) |
| \(Y\) | next-token 标签 | \([B,S]\) |
| \(M\) | SFT loss mask | \([B,S]\) |
| \(Z\) | 模型 logits | \([B,S,V]\) |
| \(L_b\) | 第 \(b\) 条样本有效监督 token 数 | 标量 |
阶段意义¶
预训练语料大多是连续文本,模型学到的是「给定前缀,什么文本更可能出现」。助手却必须理解 system、user、assistant、tool 等角色,并在约束下完成任务。SFT 用示范明确这些行为:
- 指令遵循:区分问题、上下文和目标;
- 交互协议:学习消息角色、结束标记和多轮结构;
- 输出约束:生成 JSON、代码补丁、引用或工具参数;
- 能力激活:把预训练中已有但不稳定的能力组织成可复用行为;
- 安全起点:学习基本拒答边界和替代建议。
SFT 只能模仿示范。它不知道两个合理回答哪个更受偏好,也不会因为完整 Agent 任务成功或失败而重新评价早期动作。前者需要偏好数据,后者需要奖励与信用分配。
数据序列化¶
一条结构化消息可能是:
[
{"role": "system", "content": "只返回 JSON。"},
{"role": "user", "content": "北京的城市代码是什么?"},
{"role": "assistant", "content": "{\"city\": \"BJS\"}"}
]
Chat Template 将其渲染为模型约定的特殊 token 序列。训练、验证和部署必须使用同一模板、Tokenizer 与特殊 token 配置;只要角色前缀或结束标记不同,模型面对的条件分布就已经改变。
数据处理后的张量可表示为:
常见实现直接令 \(Y_{b,t}=-100\) 表示 \(M_{b,t}=0\),并使用 ignore_index=-100。注意框架可能在模型内部完成 label shift,因此要先检查模型 API,避免数据侧和模型侧各 shift 一次。
SFT loss¶
Token 级计算¶
模型输出 \(Z=f_\theta(X)\in\mathbb{R}^{B\times S\times V}\)。和预训练相同,位置 \(t\) 的 logits 预测位置 \(t+1\) 的 token:
令
则每个未屏蔽目标 token 的负对数似然为:
按有效 token 归一化的 SFT loss 是:
结果是一个标量。反向传播从该标量经过 \(Z:[B,S,V]\) 回到所有可训练参数,但只有 \(M=1\) 的 token 直接提供监督信号。
哪些 token 参与 loss¶
通用助手数据通常屏蔽 system、user 和外部工具 observation,只训练 assistant 实际产生的内容。若 assistant 发出结构化工具调用,则 tool name 与 arguments 也属于模型动作,通常应参与 loss;工具返回值是环境观察,不应让模型学习伪造。
| 内容 | 常见 mask | 原因 |
|---|---|---|
| system 指令 | 0 | 是条件,不是模型回复 |
| user 消息 | 0 | 是输入,不应训练模型复述 |
| assistant 文本 | 1 | 是目标生成行为 |
| assistant 的 tool call | 1 | 是模型需要生成的动作 |
| tool observation | 0 | 是环境返回,不由模型生成 |
| padding | 0 | 不包含训练语义 |
这不是所有任务的绝对规则。若目标是训练纯续写模型,可对整段文本计算 loss;若要隐藏内部推理,可单独屏蔽某些 assistant 区间。关键是把策略写成显式、可测试的 span 规则,而不是依赖字符串查找猜测角色边界。
样本平均与 token 平均¶
上式让每个有效 token 权重相同,长回答对梯度贡献更大。另一种做法是先计算每条样本的平均 loss,再对样本平均:
其中 \(B'\) 是至少有一个监督 token 的样本数。它让长短样本权重相同。两种方式对应不同的数据加权策略,不能当作纯粹的数值实现细节;训练日志必须注明 reduction 方式,分布式训练时还要在全局有效 token 数上正确归一化。
多轮对话与 packing¶
多轮对话可以只监督最后一轮 assistant,也可以监督所有 assistant 轮次。监督所有轮次能利用更多 token,但早期回复可能不是在完整后文条件下采集的理想行为;只监督最后一轮更贴近当前回答,却浪费前序 assistant 的示范。应根据数据来源和任务定义选择,并通过 mask 单元测试固定行为。
packing 把多条短对话装入 \([B,S]\)。必须同时处理三类边界:
- label 边界:上一条样本结尾不能预测下一条样本开头;
- attention 边界:若要求样本隔离,构造 block-diagonal causal mask;
- position 边界:位置编号是否重置要与模型和内核支持一致。
只添加 EOS 并不总能阻止跨样本注意力。若训练中能看到下一条样本的 prompt,模型可能靠泄漏完成任务,训练 loss 会下降但验证失败。
数据质量与配比¶
SFT 数据通常比预训练数据少得多,因此单条错误样本的影响更大。质量检查至少覆盖:
- 正确性:答案事实、代码和工具参数可验证;
- 一致性:同类任务的术语、风格与拒答标准不冲突;
- 覆盖面:问答、写作、代码、多轮、结构化输出和工具场景比例合理;
- 去重:避免少数模板被重复放大;
- 完整性:没有截断回答、空 assistant 或角色错位;
- 污染:训练样本不与最终评测集重合;
- 安全与隐私:不包含不应记忆或复现的数据。
「少而精」通常比「多而杂」更可靠,但仍需要足够覆盖。为缓解灾难性遗忘,可以混合一部分通用指令或重放数据;例如从 5%—20% 做小规模配比扫描,而不是把某个比例当成通用结论。每个配方都要同时看目标任务和通用回归指标。
Self-Instruct 探索了由模型生成并筛选指令数据以扩大任务覆盖,LIMA 则强调少量精选示范对行为对齐的价值。两条路线并不矛盾:合成数据负责扩展覆盖,人工精选数据负责建立质量锚点;无论来源如何,都要经过正确性、重复度和模板一致性检查。
全参数微调与参数高效微调¶
SFT 描述训练数据和目标,不限定哪些参数被更新。全参数微调更新模型全部参数,表达能力强但显存和 checkpoint 成本高;LoRA 冻结大部分基座参数,只训练低秩 adapter,QLoRA 再量化冻结权重,完整机制见 参数高效微调。
不论参数更新范围如何,前向仍产生 \(Z:[B,S,V]\),SFT loss 也相同。LoRA 改变的是梯度流向和可训练参数集合,不是一个独立的数据阶段。小数据先用 LoRA 验证闭环,大数据或需要显著行为改变时再比较更高秩与全参数方案。
训练与监控¶
首轮训练建议执行以下检查:
- 抽样解码模板,逐 token 展示角色与 mask;
- 让单个 batch 过拟合,确认 loss 可明显下降;
- 检查非零梯度只出现在预期的可训练参数;
- 比较训练前后固定样例,确认模型不是只复述 prompt;
- 使用独立验证集记录目标能力与通用能力;
- 保存并重新加载 checkpoint,使用生产模板生成。
学习率通常低于预训练,但没有脱离模型规模、参数更新方式与有效 batch 的固定区间。轮数也不是越多越好;小型高重复数据很容易在 1—3 轮后过拟合。应优先依据验证指标、生成样例和遗忘程度早停。
案例:结构化工具调用¶
假设模型需要根据天气问题输出工具调用:
训练时把 system、user 与工具 observation 的 label 设为 -100,只监督 assistant 生成的工具名、参数和最终回答。对一个 batch:
| 张量 | 示例形状 | 内容 |
|---|---|---|
input_ids |
\([8,2048]\) | padding 后的完整多轮对话 |
logits |
\([8,2048,V]\) | 每个位置的词表 logits |
labels |
\([8,2048]\) | 非目标位置为 -100 |
loss_mask |
\([8,2048]\) | assistant 与 tool call 位置为 1 |
token_loss |
\([8,2047]\) | shift 后逐 token 交叉熵 |
loss |
\([]\) | 有效 token 平均后的标量 |
训练后至少评测 JSON 可解析率、schema 合法率、工具选择准确率、参数准确率与工具执行后的最终答案正确率。只有 JSON 可解析,不能证明工具选择或任务完成正确。
故障定位¶
| 现象 | 常见原因 | 检查方法 |
|---|---|---|
| 模型复述用户问题 | user token 未 mask、样本角色反转 | 可视化 token 与 mask |
| 输出到一半停止 | 结束标记或最大生成长度错误 | 对比训推模板和停止条件 |
| JSON 格式不稳定 | 数据格式冲突、模板不一致 | 分任务统计解析率 |
| loss 极低但验证差 | packing 泄漏、重复数据、评测污染 | 隔离 attention 与查重 |
| 多轮后角色混乱 | 特殊 token、轮次边界或 mask 错误 | 解码完整训练序列 |
| 目标任务升而通用能力跌 | 数据过窄、训练过久 | 加重放数据并早停 |
| 多卡 loss 与单卡不一致 | 分母按 rank 局部 token 数计算 | 全局汇总 loss numerator 与 mask sum |
与其他阶段的边界¶
| 方法 | 数据来源 | 直接监督信号 | 是否依赖当前策略 rollout |
|---|---|---|---|
| 预训练 | 大规模连续语料 | 真实下一个 token | 否 |
| SFT | 专家或筛选后的目标回答 | 目标回答 token | 否 |
| DPO | chosen / rejected 回答对 | 相对偏好 | 否 |
| RL | 当前策略生成的回答或轨迹 | 奖励与优势 | 是 |
| OPD | 当前学生 rollout | 教师 token 分布 | 是 |
SFT、DPO 与蒸馏都可能写成交叉熵或 log-prob 运算,但数据由谁产生、在哪些状态计算目标、监督是硬标签还是分布,决定了它们的真实差别。
个人经验¶
- 少而精的数据比多而烂的数据重要得多。
- 避免灾难性遗忘的方法是混合一定比例(例如 5% ~ 10%)的通用数据
- loss 计算时要 mask 掉:
system_prompt、user_prompt和tool_observation,只有assistant reasoning content、assistant content和tool_call才进入 loss 计算。 - chat template 要保持训推一致。