在 本地 LLM 部署 解决「模型能不能跑起来」之后,下一步往往是「用领域数据把模型拉向自己的任务」。系列索引:微调技术路线导读 | 框架对比:03-00 横向对比SFT(Supervised Fine-Tuning,监督微调)是最常见、也最容易上手的后训练手段;在 Python 生态里,TRL(Transformer Reinforcement Learning,Transformer 强化学习工具库)提供的 SFTTrainer 把 tokenize、loss 掩码、训练循环、PEFT 挂载等细节封装成一条可复用的流水线。
段末注释:SFT 用标注的输入—输出对继续优化因果语言模型的下一 token 预测;TRL 是 Hugging Face 维护的训练工具库,除 SFT 外还提供 DPO、GRPO 等对齐方法。
本文聚焦 SFTTrainer 本身:它在训练栈里扮演什么角色、接受什么数据、内部怎么算 loss、关键参数怎么配、如何与 LoRA 结合。实操案例可参考 AMD ROCm 上 LoRA 微调 Gemma。

一、SFTTrainer 在微调栈中的位置
一次典型的 LLM 监督微调可以拆成四层:
| 层级 | 职责 | 本篇对应组件 |
|---|---|---|
| 数据 | 标注样本 → 训练格式 | datasets + map 预处理 |
| 模型 | 可求导的 Causal LM | AutoModelForCausalLM 或模型 ID 字符串 |
| 训练器 | tokenize、collate、loss、循环 | SFTTrainer |
| 参数效率 | 只更新少量 adapter | peft.LoraConfig → peft_config |
SFTTrainer 继承自 Hugging Face Trainer,因此 checkpoint、日志、分布式、Hub 推送 等能力与标准 Transformers 训练一致;它在 Trainer 之上额外处理:
- 数据集格式识别(纯文本 / 对话 / prompt-completion)
- Chat template 自动套用(对话格式样本)
- Completion / Assistant 段 loss 掩码(不对 prompt 或 system/user 浪费梯度)
- 与 PEFT 的原生集成(构造时传入
peft_config即可挂 LoRA)
最小可运行示例:
1 | from datasets import load_dataset |
只传 model 与 train_dataset 时,TRL 会用默认 SFTConfig 完成其余配置(含 gradient_checkpointing=True、bf16 等 sensible defaults)。
二、依赖关系
| 包 | 作用 |
|---|---|
trl |
提供 SFTTrainer、SFTConfig |
transformers |
基座模型、Trainer 基类、tokenizer |
datasets |
加载与 map 预处理 |
peft(可选) |
LoRA / QLoRA 等 adapter |
accelerate |
设备抽象,多由 trl 间接依赖 |
安装:
1 | pip install "trl>=0.15" transformers datasets peft accelerate |
三、支持的数据格式
SFTTrainer 兼容 标准(standard) 与 对话(conversational) 两类数据集,并可细分为 language modeling 与 prompt-completion 四种形态:
3.1 四种样本形态
| 类型 | 字段结构 | 典型用途 |
|---|---|---|
| 标准 LM | {"text": "The sky is blue."} |
续写、领域语料继续预训练 |
| 对话 LM | {"messages": [{"role":"user",...}, {"role":"assistant",...}]} |
多轮指令微调 |
| 标准 prompt-completion | {"prompt": "The sky is", "completion": " blue."} |
单轮指令—回答 |
| 对话 prompt-completion | {"prompt": [...], "completion": [...]} |
Chat 模型 SFT(最常用) |
对话 prompt-completion 示例:
1 | { |
3.2 预处理:从业务表到 SFT 格式
原始分类数据 (text, label) 不能直接用,需映射为 prompt-completion:
1 | def to_prompt_completion(example): |
原则:训练用的 chat template、system prompt、用户侧 prompt 模板,必须与推理/评估阶段完全一致;否则模型学的是 A 分布,打分时测的是 B 分布。
3.3 扩展格式
- Tool Calling:样本需含
messages(含tool_calls)与tools列(JSON schema) - Vision-Language Model(VLM):需
image或images列;截断可能删掉 image token,VLM 场景常设max_length=None
四、内部机制:Tokenize 与 Loss
4.1 预处理与分词
训练时,每条样本经 tokenizer 转为 input_ids。若提供分离的 prompt 与 completion,先拼接再 tokenize(或由 collator 按模板拼接)。
对话格式会自动调用模型 tokenizer 上的 chat template(Jinja 模板),把 role/content 转成模型期望的特殊 token 序列。
4.2 损失函数
SFT 使用 token 级交叉熵(负对数似然),对序列中每个目标 token (y_t) 最大化条件概率:
$$
\mathcal{L}{\text{SFT}}(\theta) = - \sum{t=1}^{T} \log p_\theta(y_t \mid y_{<t})
$$
段末注释:上式即「给定前文,预测下一个 token」的标准 Causal LM 目标;SFT 只是在基座预训练权重上,用标注样本继续优化同一目标。
实现细节:
- Label shifting:输入序列右移一位作为 label,预测「下一个 token」
- Padding 掩码:padding 位置 label 设为
-100,不参与 loss - 默认 loss 类型:
loss_type="chunked_nll",按 chunk 计算交叉熵以降低vocab × seq_len峰值显存;可显式设为"nll"回退标准路径
4.3 只对 completion / assistant 算 loss
这是指令微调的核心配置,避免模型在 system/user 前缀上浪费梯度:
| 数据形态 | 配置项 | 行为 |
|---|---|---|
| prompt-completion | completion_only_loss=True(默认 True) |
仅 completion token 参与 loss |
| 对话 messages | assistant_only_loss=True |
仅 assistant 回复参与 loss |
1 | from trl import SFTConfig |
assistant_only_loss=True 要求 chat template 含 {% generation %}` / `{% endgeneration %} 标记;Qwen3 等常见模型 TRL 会自动 patch 模板。
若要对整段序列算 loss(如纯文本续写),设 completion_only_loss=False。
五、SFTTrainer 构造参数详解
1 | from trl import SFTConfig, SFTTrainer |
5.1 model
三种传入方式:
| 形式 | 说明 |
|---|---|
str |
Hub 模型 ID 或本地路径;按 args.model_init_kwargs 调用 from_pretrained |
PreTrainedModel |
已加载的 Causal LM;仅支持因果语言模型 |
PeftModel |
继续训练已有 adapter 时传入;不再传 peft_config |
字符串方式等价于:
1 | SFTTrainer(model="Qwen/Qwen2.5-0.5B-Instruct", ...) |
训练态手动加载时需注意:
1 | base_model.config.use_cache = False # 训练必须关闭 KV cache |
5.2 processing_class
Tokenizer 或 Processor。未传则从 model 名自动加载。必须设置 pad_token(缺省时用 eos_token)。
指令模型应确认 chat template 存在;缺失时可从官方仓库补 chat_template.jinja,或通过 SFTConfig(chat_template_path="...") 克隆模板。
5.3 peft_config
传入 LoraConfig 等 PEFT 配置后,Trainer 在初始化时把 adapter 注入模型,只优化 adapter 参数。LoRA 场景学习率通常比全量微调高一个量级(约 1e-4 vs 2e-5)。
训练前务必检查可训练参数量:
1 | trainable = sum(p.numel() for p in trainer.model.parameters() if p.requires_grad) |
若为 0,说明 target_modules 未命中模型层名。
5.4 主要方法
| 方法 | 作用 |
|---|---|
train(resume_from_checkpoint=None) |
主训练入口;可传 checkpoint 路径或 True 续训 |
evaluate() |
在 eval_dataset 上评估 |
save_model() |
保存权重(LoRA 时存 adapter) |
push_to_hub() |
推送到 Hugging Face Hub |
六、SFTConfig 关键参数
SFTConfig 继承 TrainingArguments,默认值与通用 TrainingArguments 有几处不同:
| 参数 | SFTConfig 默认 | 含义 |
|---|---|---|
logging_steps |
10 | 日志间隔 |
gradient_checkpointing |
True | 用重算换显存 |
bf16 |
True(未设 fp16 时) | BF16 训练 |
learning_rate |
2e-5 | 全量微调常用;LoRA 常改 1e-4 |
6.1 SFT 专属参数
| 参数 | 说明 |
|---|---|
max_length |
单条样本最大 token 数;VLM 慎用,可能截断 image token |
packing |
多条短样本拼进同一序列,提高 GPU 利用率 |
completion_only_loss |
prompt-completion 时仅对 completion 算 loss |
assistant_only_loss |
messages 格式时仅对 assistant 算 loss |
model_init_kwargs |
传给 from_pretrained 的参数,如 {"dtype": torch.bfloat16} |
chat_template_path |
从其他模型克隆 chat template(Base → Instruct 场景) |
eos_token |
对齐 chat template 的结束符 |
loss_type |
"chunked_nll"(默认)/ "nll" / "dft" |
dataset_text_field |
标准 LM 格式的文本列名,默认 "text" |
6.2 常用训练超参(LoRA 单卡参考)
1 | training_args = SFTConfig( |
显存不足时的调参顺序:per_device_train_batch_size → max_length → 训练样本量 → 收窄 LoRA target_modules。
七、完整示例:LoRA + Chat SFT
与 AMD 情绪分类微调 同构的最小完整链路:
1 | import torch |
产物通常为 adapter_model.safetensors + adapter_config.json,体积远小于全量权重;推理时用 PeftModel.from_pretrained(base, OUTPUT_DIR) 挂载。
八、训练日志指标
SFTTrainer 在训练/评估中会记录:
| 指标 | 含义 |
|---|---|
loss |
当前 logging 区间内非掩码 token 的平均交叉熵 |
mean_token_accuracy |
非掩码 token 上 top-1 预测命中率 |
entropy |
预测分布的平均熵 |
grad_norm |
梯度 L2 范数(clip 前) |
learning_rate |
当前学习率 |
num_tokens |
累计处理 token 数 |
对「只生成一个短标签」的分类式 SFT,mean_token_accuracy 与「学没学会」高度相关;但仍建议用任务级指标(accuracy、F1、人工抽检)做最终验收,不能只看 train loss。
九、SFTTrainer vs 其他 TRL Trainer
| Trainer | 数据要求 | 优化目标 | 典型场景 |
|---|---|---|---|
| SFTTrainer | 标注 input-output | 监督 NLL | 指令跟随、领域适配 |
| DPOTrainer | chosen / rejected 对 | 偏好对比损失 | RLHF 轻量替代 |
| GRPOTrainer | prompt + 可验证 reward | 组相对策略优化 | 推理、代码、数学 |
| RewardTrainer | 文本 + 标量 reward | 训练 reward model | RLHF 流水线中间段 |
SFT 通常是后训练链路的第一步:先把模型拉到能按格式输出,再视需要做 DPO / GRPO 等对齐。
十、常见踩坑
| 现象 | 原因 | 对策 |
|---|---|---|
| Trainable params = 0 | LoRA target_modules 未匹配 |
打印 named_modules() 核对层名 |
| 微调后推理格式乱 | 训练/推理 chat template 不一致 | 共用同一 tokenizer 与 SYSTEM_PROMPT |
| loss 下降但任务指标不涨 | 只对前缀过拟合或评估脚本不同 | 检查 completion_only_loss;统一 generate 流程 |
| OOM | 激活 + 优化器 > 推理显存 | 降 batch、max_length;开 gradient_checkpointing |
| padding 方向错误 | decoder-only 需左 padding | tokenizer.padding_side = "left"(生成场景) |
| VLM 训练报 image token 错 | max_length 截断删 image token |
max_length=None 或验证数据集长度 |
| ROCm / 特殊硬件优化器失败 | adamw_bnb_8bit 等不兼容 |
改用 optim="adamw_torch" |
| 续训 adapter 失败 | 同时传了 PeftModel 与 peft_config |
只传已加载的 PeftModel,去掉 peft_config |
十一、进阶能力(按需启用)
| 能力 | 配置 | 说明 |
|---|---|---|
| Packing | SFTConfig(packing=True) |
短样本拼序列,提升吞吐 |
| DFT Loss | loss_type="dft" |
Dynamic Fine-Tuning,改善泛化(见 TRL 论文索引) |
| Liger Kernel | use_liger_kernel=True |
Triton 内核加速,降显存 |
| Unsloth | 按 TRL 文档集成 | 2× 训练加速、约 70% 省显存 |
| Hub 推送 | push_to_hub() |
训练完一键上传 adapter 或全量 |
十二、小结
SFTTrainer 把监督微调里最容易写错的环节——格式转换、loss 掩码、与 PEFT 的衔接——收敛到一个与 Trainer 兼容的入口。实践时记住四条:
- 数据:优先用与推理一致的 chat prompt-completion;system prompt 训练/评估共用。
- Loss:指令任务开
completion_only_loss或assistant_only_loss,别让模型背 user 前缀。 - 配置:LoRA 用较高 lr +
gradient_checkpointing;OOM 先动 batch 与max_length。 - 验收:看任务指标而不只看 loss;保存的是 adapter 时,部署需
PeftModel或合并权重。
下一步可结合具体任务阅读 AMD ROCm LoRA 实操;其他工具见 03-00 框架对比;前置:01-01 Chat Template、02-01 LoRA;验收见 04 评估指标系列。