SFT(三):使用TRL SFTTrainer与LoRA微调对话模型
前两篇已经完成了两层理解:
- SFT 原理篇:SFT 仍然是 next-token prediction,关键在于数据、Chat Template 和 Loss Mask。
- Transformers 手写篇:亲手构造
input_ids / attention_mask / labels,再交给原生Trainer。
现在可以进入更接近真实项目的组合:
Hugging Face Datasets
+ Transformers
+ TRL SFTTrainer
+ PEFT LoRA
本篇不会把 SFTTrainer 描述成“一行代码训练大模型”。它确实减少了很多样板代码,但仍然需要对下面这些问题负责:
- 数据格式是否正确。
- Chat Template 是否与模型匹配。
- Assistant-only loss 是否真的生效。
- 截断是否删除了回答。
- LoRA 是否插入了预期模块。
- 训练与验证是否存在数据泄漏。
- 保存的是完整模型还是 Adapter。
- 推理时是否加载了正确的 Base Model。
1. TRL 是什么
TRL 是 Hugging Face 生态中面向语言模型后训练的库。
它提供多种 Trainer,例如:
SFTTrainer:监督微调。DPOTrainer:偏好优化。RewardTrainer:奖励模型训练。PPOTrainer:基于 PPO 的强化学习训练。- 其他后训练算法 Trainer。
本篇只讨论 SFTTrainer。
它建立在 Transformers Trainer 之上,额外处理:
SFT 常见数据格式
Chat Template
Tokenization
Completion / Assistant Loss Mask
Packing
PEFT 集成
2. SFTTrainer 自动做了什么
上一份代码中,我们手动完成:
format_prompt
-> tokenize Prompt
-> tokenize Answer
-> 拼接 input_ids
-> 构造 labels=-100
-> 动态 Padding
现在可以直接提供:
{
"messages": [
{"role": "user", "content": "什么是过拟合?"},
{"role": "assistant", "content": "过拟合是……"},
]
}
SFTTrainer 会根据 Tokenizer 与配置完成中间处理。
但“自动”不等于“无需理解”。如果模板、角色或 EOS 配置错了,Trainer 只会稳定地训练错误数据。
3. 安装依赖
pip install -U torch transformers datasets accelerate trl peft
记录版本:
python -c "import torch, transformers, datasets, trl, peft; print(torch.__version__); print(transformers.__version__); print(datasets.__version__); print(trl.__version__); print(peft.__version__)"
TRL 迭代较快,尤其是:
- 数据格式支持。
SFTConfig参数。- Chat Template 处理。
- Completion-only 与 Assistant-only loss。
- Packing 和 Loss 实现。
运行项目时应固定一组经过验证的依赖版本,并保留 pip freeze。
4. 本篇模型与训练目标
教学代码使用:
MODEL_ID = "Qwen/Qwen3-0.6B-Base"
目标是:
Qwen3 Base Causal LM
+ 中文机器学习对话示范
+ Chat Template
+ Assistant-only loss
+ LoRA
-> 一个教学用指令微调 Adapter
选择小模型只是为了降低学习门槛。少量示例不能构建可靠的通用助手。
使用任何模型前都要检查 Model Card、许可证和最低 Transformers 版本。
5. 为什么使用 Base Model
Base Model 的主要训练阶段是预训练。
它可能已经具备语言和知识能力,但并不保证稳定遵循对话指令。
从 Base Model 开始能更清楚地观察:
Chat Template + SFT Data
如何让“文本续写”呈现为“用户与助手对话”
如果目标是快速适配业务,通常也可以从 Instruct Model 开始;此时应沿用原模型自带的 Chat Template,避免改变角色协议。
6. TRL 支持哪些 SFT 数据格式
当前 SFTTrainer 主要支持两种数据类型,每种又有标准与对话形式。
6.1 Standard language modeling
{"text": "The sky is blue."}
6.2 Conversational language modeling
{
"messages": [
{"role": "user", "content": "What color is the sky?"},
{"role": "assistant", "content": "It is blue."},
]
}
6.3 Standard prompt-completion
{
"prompt": "The sky is",
"completion": " blue.",
}
6.4 Conversational prompt-completion
{
"prompt": [
{"role": "user", "content": "What color is the sky?"},
],
"completion": [
{"role": "assistant", "content": "It is blue."},
],
}
本篇使用 messages,因为它最接近聊天模型的真实输入结构。
7. 构造 Conversational Dataset
from datasets import Dataset, DatasetDict
TRAIN_CONVERSATIONS = [
{
"messages": [
{
"role": "system",
"content": "你是一名耐心、准确的机器学习助教。",
},
{
"role": "user",
"content": "请用一句话解释什么是过拟合。",
},
{
"role": "assistant",
"content": "过拟合是模型过度记忆训练数据,导致在未见数据上泛化较差的现象。",
},
]
},
]
dataset = DatasetDict(
{
"train": Dataset.from_list(TRAIN_CONVERSATIONS),
"validation": Dataset.from_list(VALIDATION_CONVERSATIONS),
}
)
7.1 为什么保留结构化 messages
结构化数据比提前拼接字符串更容易:
- 检查角色顺序。
- 替换 system prompt。
- 支持多轮对话。
- 应用不同模型的 Chat Template。
- 区分 assistant token。
- 扩展工具调用字段。
8. 数据校验不能省
ALLOWED_ROLES = {"system", "user", "assistant"}
def validate_conversation(example):
messages = example["messages"]
assert messages, "messages 不能为空"
assistant_count = 0
for message in messages:
assert message["role"] in ALLOWED_ROLES
assert isinstance(message["content"], str)
assert message["content"].strip()
assistant_count += message["role"] == "assistant"
assert assistant_count > 0, "每条训练样本至少需要一个 assistant 回答"
还应检查:
- 是否连续出现多个相同角色。
- system 是否只出现在允许位置。
- assistant 是否复制了 Prompt。
- 是否存在空回答。
- 是否存在超长样本。
- 训练集与验证集是否重复。
9. Chat Template 是训练协议
messages 只是数据结构,模型真正看到的仍是一串 token。
Chat Template 负责把:
[
{"role": "user", "content": "什么是SFT?"},
{"role": "assistant", "content": "SFT是监督微调。"},
]
变成类似:
<|user|>
什么是SFT?<|end|>
<|assistant|>
SFT是监督微调。<|end|>
模板定义:
- 角色起始 token。
- 消息结束 token。
- system prompt 放置方式。
- assistant 回答边界。
- EOS 对齐。
- 可选工具调用格式。
10. 内置模板与 chat_template_path
本文完整脚本使用 Qwen3 Tokenizer 已有的 Chat Template 和控制 token。当前 TRL 能为已知的 Qwen3 模板补充 Assistant-only loss 所需的生成区域标记,因此核心配置不必克隆另一套角色协议。
如果模型没有模板,SFTConfig 也可以指定一个模板来源:
SFTConfig(
chat_template_path="HuggingFaceTB/SmolLM3-3B",
)
TRL 会使用对应模板处理 conversational dataset,并处理所需特殊 token。
有些 Base Model Tokenizer 已经包含合适 Chat Template,这时不一定需要克隆其他模板。
判断原则:
- 优先检查模型自己的 Tokenizer 配置。
- 确认模板是否适合训练目标。
- 确认 EOS 与消息结束 token 对齐。
- 确认 Assistant-only loss 所需标记是否受支持。
- 保存最终 Tokenizer,不能只保存 Adapter。
11. assistant_only_loss=True 到底做什么
SFTConfig(
assistant_only_loss=True,
)
它让 system 和 user token 对应 labels 被忽略:
system token -> labels = -100
user token -> labels = -100
assistant token -> labels = token id
Padding token -> labels = -100
这与上一篇手写的:
labels = [-100] * len(prompt_ids) + answer_ids
属于同一原理,只是 TRL 根据 Chat Template 生成 assistant mask。
11.1 不是所有模板都能自动识别 Assistant
Assistant-only loss 需要模板能够标记 assistant 生成区域。
当前 TRL 文档使用模板中的 generation 区域标记来获得 assistant mask,并对部分已知模型族提供适配。
因此开启参数后仍应检查一条处理后的样本,而不是只相信配置名。
12. completion_only_loss 与 assistant_only_loss 的区别
12.1 completion_only_loss
适合 prompt / completion 数据:
prompt -> -100
completion -> 参与 loss
12.2 assistant_only_loss
适合 conversational 数据:
system / user -> -100
assistant -> 参与 loss
多轮对话中可以有多个 assistant 区域。
12.3 可以同时使用吗
对于 conversational prompt-completion 数据,两种范围可以组合:先只考虑 completion 部分,再只训练其中 assistant 消息。
具体行为与 TRL 版本和数据格式有关,应根据当前官方文档与处理结果确认。
13. SFTConfig 的核心参数
from trl import SFTConfig
training_args = SFTConfig(
output_dir="outputs/sft-trl-lora/checkpoints",
max_length=512,
eos_token="<|im_end|>",
assistant_only_loss=True,
packing=False,
loss_type="nll",
num_train_epochs=3,
per_device_train_batch_size=1,
per_device_eval_batch_size=1,
gradient_accumulation_steps=8,
learning_rate=1e-4,
warmup_ratio=0.1,
lr_scheduler_type="cosine",
logging_steps=1,
eval_strategy="epoch",
save_strategy="epoch",
load_best_model_at_end=True,
metric_for_best_model="eval_loss",
greater_is_better=False,
gradient_checkpointing=True,
report_to="none",
seed=42,
model_init_kwargs={"dtype": dtype},
)
13.1 max_length
它限制处理后的序列长度。
不能只根据模型最大上下文设置。还要结合:
- 数据长度分布。
- GPU 显存。
- Batch Size。
- Attention 的二次复杂度。
- 截断后回答是否保留。
13.2 packing
packing=True 会把多个短样本打包,提高 token 利用率。
学习阶段先设为 False,更容易逐条检查。
13.3 eos_token
Qwen3 Chat Template 使用 <|im_end|> 结束一条消息,因此显式把训练结束 token 与模板对齐:
eos_token="<|im_end|>"
13.4 learning_rate
LoRA 只训练新增参数,常使用比 Full Fine-tuning 更高的学习率。
1e-4 是一个实验起点,不是通用最优值。
13.5 model_init_kwargs
当 model 以 Hub ID 字符串传给 SFTTrainer 时,可以通过它传递 from_pretrained() 参数:
model_init_kwargs={"dtype": torch.bfloat16}
14. 为什么本篇显式设置 loss_type=”nll”
当前 TRL 还提供降低大词表 logits 峰值显存的 chunked loss 实现。
但具体 loss 实现与 PEFT、Liger 等功能之间可能存在版本相关的兼容限制。
本篇目标是讲清 LoRA SFT,因此显式使用标准负对数似然:
loss_type="nll"
数学目标仍然是:
$$
\mathcal{L}
=-
\frac{1}{N}
\sum_{t: label_t\neq -100}
\log p_{\theta}(x_t\mid x_{<t})
$$
引入新的显存优化前,应先阅读正在使用的 TRL 版本文档并单独验证兼容性。
15. LoRA 的数学原理
原线性层:
$$
y=Wx
$$
LoRA 冻结 $W$,学习低秩增量:
$$
y=(W+\Delta W)x
$$
$$
\Delta W=\frac{\alpha}{r}BA
$$
其中:
$$
A\in\mathbb{R}^{r\times d_{in}}
$$
$$
B\in\mathbb{R}^{d_{out}\times r}
$$
如果原矩阵参数量为:
$$
d_{out}d_{in}
$$
LoRA 新增参数量为:
$$
r(d_{in}+d_{out})
$$
当 $r$ 很小时,新增参数远少于完整矩阵。
16. 配置 LoraConfig
from peft import LoraConfig
peft_config = LoraConfig(
r=16,
lora_alpha=32,
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
target_modules="all-linear",
)
16.1 r
r 是低秩维度。
更大的 r:
- 表达能力更高。
- 可训练参数更多。
- 显存和计算增加。
- 不一定带来更好泛化。
16.2 lora_alpha
它控制 LoRA 更新的缩放。
原始形式常写为:
$$
\frac{\alpha}{r}
$$
16.3 lora_dropout
只在 LoRA 分支上应用 Dropout,可缓解小数据过拟合。
16.4 target_modules=”all-linear”
PEFT 会选择模型中的线性层,同时对 PreTrainedModel 排除最终输出层。
优点是避免硬编码不同架构的模块名。
如果只想训练特定投影,也可以显式配置:
target_modules=["q_proj", "k_proj", "v_proj", "o_proj"]
前提是目标模型确实使用这些名称。
17. 创建 SFTTrainer
from trl import SFTTrainer
trainer = SFTTrainer(
model=MODEL_ID,
args=training_args,
train_dataset=dataset["train"],
eval_dataset=dataset["validation"],
peft_config=peft_config,
)
这几行背后发生:
根据 MODEL_ID 加载 Causal LM
-> 加载 Tokenizer / Processing Class
-> 准备 Chat Template
-> 处理 messages
-> Tokenize
-> 构造 Assistant Mask 与 labels
-> 用 PEFT 包装模型
-> 创建 Trainer 训练循环
18. 确认可训练参数
trainer.model.print_trainable_parameters()
预期看到:
trainable params: 较小数值
all params: Base Model + Adapter
trainable%: 很小比例
还可以手动统计:
trainable = sum(
parameter.numel()
for parameter in trainer.model.parameters()
if parameter.requires_grad
)
total = sum(
parameter.numel()
for parameter in trainer.model.parameters()
)
print(trainable, total, trainable / total)
如果几乎所有参数都可训练,说明 PEFT 配置没有按预期生效。
19. LoRA 为什么能减少训练显存
Full Fine-tuning 需要为大量参数保存:
- 梯度。
- AdamW 一阶矩。
- AdamW 二阶矩。
LoRA 冻结主干后,这些状态只为 Adapter 参数创建。
但 LoRA 不会消除所有显存:
Base Model 权重仍要加载
前向激活仍然存在
Attention 中间结果仍然存在
部分 logits 仍然存在
因此序列过长或 Batch 太大时,LoRA 仍可能 OOM。
20. Gradient Checkpointing
gradient_checkpointing=True
它不保存全部前向激活,而是在反向传播时重新计算一部分中间结果。
交换关系是:
更少显存
<- 代价 ->
更多计算时间
它和 LoRA 解决不同问题:
- LoRA 减少可训练参数与优化器状态。
- Gradient Checkpointing 减少前向激活保存。
21. BF16 与 FP16
def choose_dtype():
if not torch.cuda.is_available():
return torch.float32
if torch.cuda.is_bf16_supported():
return torch.bfloat16
return torch.float16
一般而言:
- 支持 BF16 的设备优先尝试 BF16。
- 较旧 GPU 可能只能高效使用 FP16。
- CPU 教学检查使用 FP32 更稳妥,但训练很慢。
混合精度不是简单把所有张量永久变成低精度;训练框架会对不同计算使用合适策略。
22. Packing 何时开启
先统计长度:
P50 = 120 token
P90 = 220 token
max_length = 1024
如果大多数样本很短,Packing 可以显著降低 Padding 比例。
SFTConfig(
packing=True,
)
但第一次训练建议关闭:
packing=False
先确认:
- 不 Packing 时 loss 正常下降。
- Assistant Mask 正确。
- 单条样本能被正确 decode。
- 验证集行为符合预期。
再开启 Packing 做速度对比。
23. 启动训练
trainer.train()
建议重点观察:
loss。eval_loss。learning_rate。grad_norm。- 有效 token 数。
- 每秒处理 token 数。
- GPU 峰值显存。
训练 loss 突然为 0 或 NaN 时,先检查数据与 labels,而不是立即调整优化器。
24. 保存 LoRA Adapter
ADAPTER_DIR = "outputs/sft-trl-lora/final-adapter"
trainer.save_model(ADAPTER_DIR)
trainer.processing_class.save_pretrained(ADAPTER_DIR)
Adapter 目录通常包含:
final-adapter/
adapter_config.json
adapter_model.safetensors
tokenizer / chat template assets
training_args.bin
它通常不包含完整 Base Model 权重。
部署时需要:
原 Base Model
+ LoRA Adapter
+ 训练时对应 Tokenizer / Chat Template
25. 重新加载 Adapter
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(ADAPTER_DIR)
base_model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
dtype=dtype,
)
model = PeftModel.from_pretrained(
base_model,
ADAPTER_DIR,
)
model.eval()
必须保证:
MODEL_ID与训练 Adapter 时的 Base Model 一致。- Tokenizer 使用训练后保存的版本。
- 特殊 token 数量与 embedding 对齐。
- Chat Template 与训练时一致。
26. 使用 Chat Template 推理
推理数据不包含标准答案:
messages = [
{
"role": "system",
"content": "你是一名耐心、准确的机器学习助教。",
},
{
"role": "user",
"content": "请解释训练集和验证集的区别。",
},
]
构造模型输入:
inputs = tokenizer.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
return_dict=True,
return_tensors="pt",
).to(model.device)
训练时使用 add_generation_prompt=False,因为标准 assistant 答案已经存在;推理时使用 True,提示模型开始 assistant turn。
27. GenerationConfig 与生成
from transformers import GenerationConfig
generation_config = GenerationConfig(
max_new_tokens=128,
do_sample=True,
temperature=0.8,
top_p=0.9,
pad_token_id=tokenizer.pad_token_id,
eos_token_id=tokenizer.eos_token_id,
)
生成:
with torch.no_grad():
output_ids = model.generate(
**inputs,
generation_config=generation_config,
)
prompt_length = inputs["input_ids"].shape[1]
new_tokens = output_ids[0, prompt_length:]
answer = tokenizer.decode(new_tokens, skip_special_tokens=True)
print(answer)
SFT 更新了模型参数;GenerationConfig 决定如何从新 logits 分布中选择 token。两者仍然属于不同层次。
28. 可选:合并 Adapter
如果推理框架不方便分别加载 Base Model 和 Adapter,可以合并:
merged_model = model.merge_and_unload()
merged_model.save_pretrained(
"outputs/sft-trl-lora/merged-model",
safe_serialization=True,
)
tokenizer.save_pretrained(
"outputs/sft-trl-lora/merged-model",
)
合并后:
W_merged = W_base + Delta_W_lora
优点:
- 推理时只加载一个模型目录。
- 不需要额外 Adapter 路由。
代价:
- 失去 Adapter 小体积优势。
- 每个任务可能都要保存一份完整模型。
- 不再方便动态切换多个 Adapter。
因此完整脚本默认不自动合并,必须显式开启开关。
29. QLoRA 与本篇 LoRA 有什么区别
本篇:
低精度或普通精度 Base Model
+ LoRA Adapter
QLoRA 通常进一步:
4-bit 量化 Base Model
+ 可训练 LoRA Adapter
它进一步降低 Base Model 权重显存,但会增加:
- 量化配置。
- bitsandbytes 等依赖。
- 设备兼容问题。
- 量化 dtype 与计算 dtype 的选择。
- 保存、加载和合并复杂度。
在 LoRA 流程完全正确前,不建议一开始就叠加 QLoRA。
30. 完整代码
完整代码位于:
source/code/nlp-blog/sft_trl_lora.py
运行:
python source/code/nlp-blog/sft_trl_lora.py
脚本会执行:
创建 messages Dataset
-> 校验角色与内容
-> 创建 SFTConfig
-> 创建 LoraConfig
-> SFTTrainer 自动应用 Chat Template
-> Assistant-only loss
-> 训练与验证
-> 保存 Adapter 与 Tokenizer
-> 重新加载 Base + Adapter
-> apply_chat_template
-> GenerationConfig
-> generate
-> 可选 merge_and_unload
31. 怎样检查 Assistant Mask 是否真的正确
不同版本的内部处理接口可能变化,因此不要依赖私有属性。
可以采用三层检查。
31.1 数据层
打印原始 messages,确认角色和内容。
31.2 模板层
formatted = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=False,
)
print(formatted)
确认 user、assistant 和 turn-end token 都在正确位置。
31.3 训练行为层
构造最小数据集并比较:
assistant_only_loss=False。assistant_only_loss=True。
观察有效 token 数、loss 和显存是否符合预期。
如果需要逐 token 审计,可以预先使用 Chat Template 生成 input_ids 与 assistant mask,构造预 Tokenized Dataset,再把 labels=-100 的位置 decode 出来检查。
32. 常见错误
32.1 把 prompt-only 数据传给 SFTTrainer
SFT 需要目标文本。只有 Prompt 没有 Completion,无法构造标准监督答案。
32.2 使用了错误 Chat Template
模型可能继续用户消息、输出错误控制 token 或无法停止。
32.3 训练时 add_generation_prompt=True
标准 assistant 答案已经在数据中,额外的生成起始标记可能制造错误尾部。
32.4 assistant_only_loss=True 但模板不支持 Assistant Mask
参数名存在不代表模板能够正确标记区域。检查当前官方支持条件。
32.5 EOS 与模板结束 token 不一致
模型可能无法在预期位置停止。
32.6 LoRA target_modules 不匹配
可能没有插入 Adapter,或插入范围与预期不同。打印模型和可训练参数。
32.7 误以为 LoRA 不需要 Base Model
未合并的 Adapter 只保存增量参数。
32.8 保存 Adapter 却遗漏 Tokenizer
Chat Template 和特殊 token 可能随 Tokenizer 保存,必须一起保留。
32.9 eval_loss 下降就认为回答变好
仍需生成评估、任务指标和人工检查。
32.10 小数据训练太多 Epoch
模型可能记住答案,同时破坏原有通用能力。
32.11 随意开启所有显存优化
Packing、Gradient Checkpointing、量化、Liger、Flash Attention 和特殊 Loss 实现可能存在组合限制。一次只增加一个变量并验证。
32.12 用测试集调整超参数
验证集用于选择配置,测试集只用于最终评估。
33. 一个可靠的 SFT 实验记录应包含什么
模型
model_id
revision / commit hash
architecture
parameter count
tokenizer revision
chat template
数据
数据来源与许可证
清洗规则
去重规则
训练/验证/测试数量
token 长度分布
截断比例
有效 assistant token 比例
训练
Full / LoRA / QLoRA
LoRA r / alpha / target modules
learning rate
batch size
gradient accumulation
max_length
packing
dtype
训练步数
随机种子
评估
eval_loss
任务指标
固定生成 Prompt
GenerationConfig
人工错误分类
Base Model 对照
34. 从 SFT 继续学习什么
完成本系列后,可以继续:
数据去重与质量评分
-> Chat Template 自定义
-> LoRA target modules 对比
-> QLoRA
-> Packing 与长度分桶
-> 生成式评估
-> 多轮对话评估
-> DPO 偏好优化
-> Reward Model
-> PPO / GRPO 等在线或强化学习方法
-> 分布式训练与推理部署
其中 DPO、PPO 或 GRPO 不会替代 SFT 的基础作用。多数后训练流程仍需要一个质量足够的 SFT 模型作为起点。
35. 总结
这一篇把手写 SFT 数据管线替换成了 TRL 与 PEFT:
messages
-> Chat Template
-> Tokenizer
-> Assistant Mask
-> labels=-100 / token ids
-> SFTTrainer
-> Causal LM Loss
-> LoRA 参数更新
-> Adapter + Tokenizer 保存
-> Base + Adapter 重新加载
-> GenerationConfig + generate
我现在可以把 SFTTrainer 理解成:
一个知道 SFT 常见数据格式和处理规则、并与 Transformers Trainer 和 PEFT 集成的训练器,而不是一种新的模型或新的语言模型数学目标。
系列三篇对应三层能力:
- 原理层:知道为什么是 next-token loss,知道三个 Mask 的区别。
- 张量层:能自己构造
input_ids / attention_mask / labels并手动 forward。 - 工程层:能使用
SFTTrainer + LoRA,也知道自动化边界和排错方法。
真正开始项目时,最先投入精力的往往不是继续增加 Trainer 参数,而是数据质量、模板一致性和可靠评估。
参考资料
- Hugging Face TRL:SFT Trainer
- Hugging Face TRL:Dataset formats
- Hugging Face TRL:Reducing memory usage
- Hugging Face Transformers:Chat templates
- Hugging Face PEFT:LoRA conceptual guide
- Hugging Face PEFT:LoraConfig
- Qwen3-0.6B-Base Model Card
- LoRA: Low-Rank Adaptation of Large Language Models
- QLoRA: Efficient Finetuning of Quantized LLMs
- 上一篇:SFT(二)用Hugging Face Transformers手写监督微调
