SFTTrainer 详解:TRL 监督微调训练器

本地 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 监督微调流程概览


一、SFTTrainer 在微调栈中的位置

一次典型的 LLM 监督微调可以拆成四层:

层级 职责 本篇对应组件
数据 标注样本 → 训练格式 datasets + map 预处理
模型 可求导的 Causal LM AutoModelForCausalLM 或模型 ID 字符串
训练器 tokenize、collate、loss、循环 SFTTrainer
参数效率 只更新少量 adapter peft.LoraConfigpeft_config

SFTTrainer 继承自 Hugging Face Trainer,因此 checkpoint、日志、分布式、Hub 推送 等能力与标准 Transformers 训练一致;它在 Trainer 之上额外处理:

  1. 数据集格式识别(纯文本 / 对话 / prompt-completion)
  2. Chat template 自动套用(对话格式样本)
  3. Completion / Assistant 段 loss 掩码(不对 prompt 或 system/user 浪费梯度)
  4. 与 PEFT 的原生集成(构造时传入 peft_config 即可挂 LoRA)

最小可运行示例:

1
2
3
4
5
6
7
8
from datasets import load_dataset
from trl import SFTTrainer

trainer = SFTTrainer(
model="Qwen/Qwen3-0.6B",
train_dataset=load_dataset("trl-lib/Capybara", split="train"),
)
trainer.train()

只传 modeltrain_dataset 时,TRL 会用默认 SFTConfig 完成其余配置(含 gradient_checkpointing=Truebf16 等 sensible defaults)。


二、依赖关系

作用
trl 提供 SFTTrainerSFTConfig
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
2
3
4
5
6
7
8
9
{
"prompt": [
{"role": "system", "content": "你是情绪分类助手,只输出标签。"},
{"role": "user", "content": "Classify: I feel anxious about tomorrow."},
],
"completion": [
{"role": "assistant", "content": "fear"},
],
}

3.2 预处理:从业务表到 SFT 格式

原始分类数据 (text, label) 不能直接用,需映射为 prompt-completion:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
def to_prompt_completion(example):
label = label_names[example["label"]]
return {
"prompt": [
{"role": "system", "content": SYSTEM_PROMPT},
{"role": "user", "content": f"Classify:\n\n{example['text']}"},
],
"completion": [{"role": "assistant", "content": label}],
}

sft_dataset = raw_dataset.map(
to_prompt_completion,
remove_columns=raw_dataset["train"].column_names,
)

原则:训练用的 chat template、system prompt、用户侧 prompt 模板,必须与推理/评估阶段完全一致;否则模型学的是 A 分布,打分时测的是 B 分布。

3.3 扩展格式

  • Tool Calling:样本需含 messages(含 tool_calls)与 tools 列(JSON schema)
  • Vision-Language Model(VLM):需 imageimages 列;截断可能删掉 image token,VLM 场景常设 max_length=None

四、内部机制:Tokenize 与 Loss

4.1 预处理与分词

训练时,每条样本经 tokenizer 转为 input_ids。若提供分离的 promptcompletion先拼接再 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
2
3
4
5
6
from trl import SFTConfig

training_args = SFTConfig(
completion_only_loss=True, # prompt-completion 数据集
# assistant_only_loss=True, # 纯 messages 数据集时使用
)

assistant_only_loss=True 要求 chat template 含 {% generation %}` / `{% endgeneration %} 标记;Qwen3 等常见模型 TRL 会自动 patch 模板。

若要对整段序列算 loss(如纯文本续写),设 completion_only_loss=False


五、SFTTrainer 构造参数详解

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
from trl import SFTConfig, SFTTrainer
from peft import LoraConfig

trainer = SFTTrainer(
model=base_model, # 或模型 ID 字符串
args=training_args, # SFTConfig 实例
train_dataset=train_ds,
eval_dataset=eval_ds, # 可选
processing_class=tokenizer, # 推荐显式传入
peft_config=lora_config, # 可选,LoRA
formatting_func=None, # 可选,tokenize 前格式化
data_collator=None, # 可选,默认 LM collator
compute_metrics=None,
callbacks=None,
)

5.1 model

三种传入方式:

形式 说明
str Hub 模型 ID 或本地路径;按 args.model_init_kwargs 调用 from_pretrained
PreTrainedModel 已加载的 Causal LM;仅支持因果语言模型
PeftModel 继续训练已有 adapter 时传入;不再peft_config

字符串方式等价于:

1
2
SFTTrainer(model="Qwen/Qwen2.5-0.5B-Instruct", ...)
# 内部 ≈ AutoModelForCausalLM.from_pretrained(..., **model_init_kwargs)

训练态手动加载时需注意:

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
2
3
trainable = sum(p.numel() for p in trainer.model.parameters() if p.requires_grad)
total = sum(p.numel() for p in trainer.model.parameters())
print(f"Trainable: {trainable:,} / {total:,} ({100*trainable/total:.2f}%)")

若为 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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
training_args = SFTConfig(
output_dir="./output",

per_device_train_batch_size=4,
gradient_accumulation_steps=4, # 等效 batch = 16

learning_rate=1e-4,
weight_decay=0.01,
lr_scheduler_type="linear",
warmup_steps=50,
num_train_epochs=1,

logging_steps=5,
eval_strategy="steps",
eval_steps=25,
save_strategy="steps",
save_steps=25,

gradient_checkpointing=True,
bf16=True,

max_length=256,
completion_only_loss=True,
optim="adamw_torch",

seed=42,
)

显存不足时的调参顺序:per_device_train_batch_sizemax_length → 训练样本量 → 收窄 LoRA target_modules


七、完整示例:LoRA + Chat SFT

AMD 情绪分类微调 同构的最小完整链路:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
import torch
from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig
from trl import SFTConfig, SFTTrainer

MODEL_DIR = "./models/my-instruct-model"
OUTPUT_DIR = "./lora-output"

tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR, use_fast=True)
base_model = AutoModelForCausalLM.from_pretrained(
MODEL_DIR, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True,
)
base_model.to("cuda")
base_model.config.use_cache = False

lora_config = LoraConfig(
r=16, lora_alpha=32, lora_dropout=0.05,
task_type="CAUSAL_LM", target_modules="all-linear",
)

training_args = SFTConfig(
output_dir=OUTPUT_DIR,
per_device_train_batch_size=4,
gradient_accumulation_steps=4,
learning_rate=1e-4,
num_train_epochs=1,
max_length=256,
completion_only_loss=True,
gradient_checkpointing=True,
bf16=True,
logging_steps=5,
eval_strategy="steps",
eval_steps=25,
optim="adamw_torch",
)

trainer = SFTTrainer(
model=base_model,
train_dataset=sft_dataset["train"],
eval_dataset=sft_dataset["validation"],
peft_config=lora_config,
args=training_args,
processing_class=tokenizer,
)

train_result = trainer.train()
trainer.save_model(OUTPUT_DIR)

产物通常为 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 失败 同时传了 PeftModelpeft_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 兼容的入口。实践时记住四条:

  1. 数据:优先用与推理一致的 chat prompt-completion;system prompt 训练/评估共用。
  2. Loss:指令任务开 completion_only_lossassistant_only_loss,别让模型背 user 前缀。
  3. 配置:LoRA 用较高 lr + gradient_checkpointing;OOM 先动 batch 与 max_length
  4. 验收:看任务指标而不只看 loss;保存的是 adapter 时,部署需 PeftModel 或合并权重。

下一步可结合具体任务阅读 AMD ROCm LoRA 实操;其他工具见 03-00 框架对比;前置:01-01 Chat Template02-01 LoRA;验收见 04 评估指标系列

-------------本文结束感谢您的阅读-------------