第 11 章 TRL 与模型对齐
学习目标
- 理解预训练 → 监督微调(SFT)→ 偏好对齐(DPO/GRPO)的链路
- 用
SFTConfig+SFTTrainer做监督微调 - 用
DPOConfig+DPOTrainer做偏好对齐 - 了解 GRPO 与 RLHF 的关系
11.1 从预训练到对齐
大语言模型的训练分阶段:
预训练(海量文本,学语言) → 监督微调 SFT(指令数据,学对话)
→ 偏好对齐 DPO/RLHF(人类偏好,学「好回答」)TRL(Transformer Reinforcement Learning) 库覆盖后两个阶段,本书只讲 Python 侧的训练接口。
| 方法 | 数据形态 | 学什么 |
|---|---|---|
| SFT | text(指令+回答) | 模仿高质量回答 |
| DPO | prompt + chosen + rejected | 偏好优答、远离劣答 |
| GRPO | prompt + 奖励函数 | 用组内相对奖励优化策略 |
11.2 监督微调:SFT
数据是一行行 {"text": "指令\n回答"}。用 SFTConfig 定义超参,SFTTrainer 训练:
from datasets import Dataset
from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed
from trl import SFTConfig, SFTTrainer
set_seed(42)
model = AutoModelForCausalLM.from_pretrained("hf-internal-testing/tiny-random-gpt2")
tokenizer = AutoTokenizer.from_pretrained("hf-internal-testing/tiny-random-gpt2")
tokenizer.pad_token = tokenizer.eos_token
dataset = Dataset.from_list([
{"text": "User: 什么是模型微调?\nAssistant: 模型微调是在预训练模型基础上,用下游任务数据继续训练。"},
{"text": "User: 解释 LoRA。\nAssistant: LoRA 是一种参数高效微调方法,只训练低秩矩阵。"},
{"text": "User: 什么是分词器?\nAssistant: 分词器把文本切成模型能处理的 token。"},
{"text": "User: 什么是 Trainer?\nAssistant: Trainer 是 transformers 提供的高层训练封装。"},
])
sft_args = SFTConfig(
output_dir="./results-sft",
max_steps=3,
per_device_train_batch_size=2,
logging_steps=1,
report_to=[],
disable_tqdm=True,
fp16=True,
)
trainer = SFTTrainer(
model=model,
args=sft_args,
train_dataset=dataset,
processing_class=tokenizer,
)
result = trainer.train()
print("steps:", result.global_step, "| loss:", round(result.training_loss, 4))训练日志(真实运行,预处理进度条省略):
{'loss': '6.836', 'grad_norm': '2.113', 'learning_rate': '2e-05', 'entropy': '6.901', 'num_tokens': '190', 'mean_token_accuracy': '0', 'epoch': '0.3333'}
{'loss': '6.807', 'grad_norm': '1.957', 'learning_rate': '1.333e-05', 'entropy': '6.901', 'num_tokens': '356', 'mean_token_accuracy': '0', 'epoch': '0.6667'}
{'loss': '6.821', 'grad_norm': '1.639', 'learning_rate': '6.667e-06', 'entropy': '6.901', 'num_tokens': '532', 'epoch': '1'}
steps: 3 | loss: 6.8213TRL 预处理时会自动完成:加 EOS、分词、建 labels、截断、丢弃全掩码样本——你只需要提供原始文本。
示例模型
tiny-random-gpt2是随机初始化的测试模型,loss 高、输出无意义是正常的;换真实模型(如openai-community/gpt2或更大的开源模型)后,loss 会随训练下降。
11.3 偏好对齐:DPO
DPO 数据是三元组:prompt(指令)、chosen(好回答)、rejected(差回答):
from transformers import AutoModelForCausalLM, AutoTokenizer, set_seed
from datasets import Dataset
from trl import DPOConfig, DPOTrainer
set_seed(42)
model = AutoModelForCausalLM.from_pretrained("hf-internal-testing/tiny-random-gpt2")
ref_model = AutoModelForCausalLM.from_pretrained("hf-internal-testing/tiny-random-gpt2")
tokenizer = AutoTokenizer.from_pretrained("hf-internal-testing/tiny-random-gpt2")
tokenizer.pad_token = tokenizer.eos_token
pref_data = Dataset.from_list([
{"prompt": "解释什么是 LoRA。",
"chosen": "LoRA 只训练低秩矩阵,参数量极小。",
"rejected": "LoRA 需要训练全部参数。"},
{"prompt": "什么是微调?",
"chosen": "用下游数据继续训练预训练模型。",
"rejected": "微调就是重新预训练。"},
{"prompt": "如何加载模型?",
"chosen": "用 from_pretrained 从 Hub 加载。",
"rejected": "手动复制权重文件。"},
])
dpo_args = DPOConfig(
output_dir="./results-dpo",
max_steps=1,
per_device_train_batch_size=1,
report_to=[],
disable_tqdm=True,
fp16=True,
)
dpo = DPOTrainer(
model=model,
ref_model=ref_model,
args=dpo_args,
train_dataset=pref_data,
processing_class=tokenizer,
)
result = dpo.train()
print("steps:", result.global_step, "| loss:", round(result.training_loss, 4))训练日志(真实运行):
{'train_runtime': '0.8046', 'train_samples_per_second': '1.243', 'train_steps_per_second': '1.243',
'train_loss': '0.6931', 'entropy': '6.901', 'num_tokens': '125',
'logits/chosen': '0.000737', 'logits/rejected': '0.001088',
'rewards/chosen': '0', 'rewards/rejected': '0', 'rewards/accuracies': '0', 'rewards/margins': '0',
'logps/chosen': '-331.6', 'logps/rejected': '-229.2', 'epoch': '0.3333'}
steps: 1 | loss: 0.6931关键字段:
logits/chosen、logits/rejected:模型对优答/劣答的打分。rewards/accuracies:chosen 得分是否高于 rejected 的比例(1 步时还是 0,训练充分后应接近 1)。logps/chosen、logps/rejected:chosen/rejected 的对数概率,DPO 损失的核心输入。
ref_model 是「对齐前的参考模型」,DPO 防止模型偏离它太远。常见做法是直接用 SFT 后的模型同时作为 policy 与 reference;示例里为了可复现,显式加载了两个实例。
11.4 GRPO 与 RLHF
GRPO(Group Relative Policy Optimization) 是当前开源界主流的策略优化方法:对每个 prompt 采样一组回答,用奖励函数打分,再以「组内相对优势」更新策略。TRL 提供 GRPOConfig 与 GRPOTrainer:
from trl import GRPOConfig, GRPOTrainer
grpo_args = GRPOConfig(
output_dir="./results-grpo",
max_steps=10,
per_device_train_batch_size=1,
report_to=[],
)
grpo = GRPOTrainer(
model=model,
args=grpo_args,
train_dataset=prompt_dataset,
processing_class=tokenizer,
reward_funcs=[my_reward_function], # 自定义奖励函数
)
grpo.train()GRPO 需要采样生成,计算量远大于 SFT/DPO,通常配合 vLLM 等推理加速器使用;本书不做完整运行示例。概念上记住一句话:RLHF 把「人类偏好」变成奖励信号,GRPO 用组内相对比较稳定地优化它。
动手实践
- 把 11.2 的
max_steps改成 1 和 10,观察 loss 变化;换openai-community/gpt2重跑(注意下载体积)。 - 自己构造 5 条中文偏好数据(注意:chosen 必须明显优于 rejected),跑 2 步 DPO,观察
rewards/accuracies。 - 对比 SFT 与 DPO 训练日志中的字段,列出各自独有的指标。
常见错误
错误 1:ref_model=None 在某些模型上失败。
TRL 会用「模型 id + config」重建参考模型;若仓库的 config.json 缺 architectures,会报 TypeError: 'NoneType' object is not subscriptable。显式传 ref_model 实例最稳妥。
错误 2:GPT-2 没有 pad token 直接训练。
ValueError: ... pad_token must be set ...先 tokenizer.pad_token = tokenizer.eos_token。
错误 3:把 SFT 的 text 数据直接喂给 DPO。
DPO 必须有三列(prompt/chosen/rejected),列名不可省略;TRL 预处理按列名识别。
章末练习
基础
- 写出 SFT 数据与 DPO 数据的列结构差异。
- 在 11.2 代码中,
processing_class参数传的是什么对象?作用是什么?
提高
- 用
tiny-random-gpt2跑 10 步 SFT,记录 loss;再用训练后的模型做一次generate,对比训练前后的输出。 - 解释
logps/chosen与logps/rejected的差异如何影响 DPO 损失。
挑战
- 给 DPO 数据设计一个简单的启发式奖励(如回答长度、关键词命中),写一个
reward_funcs并在 GRPO 配置里接入,跑 3 步,记录 reward 曲线。
章末自测
- SFT 的英文全称与中文含义是什么?
- DPO 数据需要哪三列?
- DPO 中
ref_model的作用是什么? - GRPO 中的「R」代表什么?
- 判断:SFTTrainer 的输入数据必须已经完成分词。
- 训练日志中
rewards/accuracies理想情况下应趋近多少?
