DPOTrainer 详解:TRL 直接偏好优化

05-01 偏好对齐选型 选定 DPO 后,DPOTrainer 是 TRL 中的实现入口。本篇覆盖数据格式、配置、与 SFT adapter 的衔接;损失推导见 Math-05/20 RLHF 与 KL

段末注释:DPO(Direct Preference Optimization,直接偏好优化)通过增大 chosen 相对 rejected 的对数似然差来对齐偏好,无需单独训练奖励模型。

系列索引:微调技术路线导读 | 数据:01-03 偏好数据


一、在流水线中的位置

1
2
3
4
5
SFT adapter (π_θ 初值, π_ref 冻结)

└─ DPOTrainer(chosen/rejected)

└─ adapter_dpo → 评估 → 部署

二、依赖与最小示例

1
pip install "trl>=0.15" transformers peft datasets accelerate
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
from datasets import load_dataset
from trl import DPOConfig, DPOTrainer
from peft import LoraConfig, PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer

BASE_ID = "./sft-merged-or-base"
SFT_ADAPTER = "./lora-sft"

tokenizer = AutoTokenizer.from_pretrained(BASE_ID)
policy = AutoModelForCausalLM.from_pretrained(BASE_ID, torch_dtype="bfloat16", device_map="auto")
policy = PeftModel.from_pretrained(policy, SFT_ADAPTER, is_trainable=True)

ref_model = AutoModelForCausalLM.from_pretrained(BASE_ID, torch_dtype="bfloat16", device_map="auto")
ref_model = PeftModel.from_pretrained(ref_model, SFT_ADAPTER)
ref_model.eval()
for p in ref_model.parameters():
p.requires_grad = False

train_ds = load_dataset("trl-lib/hh-rlhf-helpful-base", split="train[:5%]")

training_args = DPOConfig(
output_dir="./dpo-out",
beta=0.1,
learning_rate=5e-5,
per_device_train_batch_size=2,
gradient_accumulation_steps=4,
num_train_epochs=1,
max_length=1024,
max_prompt_length=512,
bf16=True,
gradient_checkpointing=True,
)

trainer = DPOTrainer(
model=policy,
ref_model=ref_model,
args=training_args,
train_dataset=train_ds,
processing_class=tokenizer,
)
trainer.train()
trainer.save_model("./dpo-out")

三、数据格式

01-03 一致:

1
2
3
4
5
{
"prompt": [{"role": "user", "content": "..."}],
"chosen": [{"role": "assistant", "content": "..."}],
"rejected": [{"role": "assistant", "content": "..."}],
}

TRL 也支持 prompt/chosen/rejected 为字符串列(视版本);推荐 messages 与 SFT 一致。


四、DPOConfig 关键参数

参数 典型值 说明
beta 0.1–0.5 KL 强度;越大越贴近 ref
learning_rate 5e-6–5e-5 通常 低于 LoRA SFT
max_length 512–2048 prompt+chosen 总长上限
max_prompt_length 略小于 max_length 防截断丢 chosen
loss_type sigmoid(默认) 变体见 TRL 文档

五、ref_model 要点

规则 说明
ref = SFT 后策略 对齐起点与 KL 锚点
ref 冻结 不参与反向传播
policy 可 train 从 SFT adapter 继续训
勿重复 peft_config policy 已 PeftModel 则不再传

无 ref 时 TRL 可能用 policy 副本或禁用 KL 项(依版本);显式传入 ref_model 最稳妥


六、LoRA 两阶段参数建议

阶段 lr beta epoch
SFT (LoRA) 1e-4 1–3
DPO (LoRA) 5e-5 0.1 1

rank 可与 SFT 相同;DPO 不必增大 rank。


七、评估

监控 说明
loss DPO 损失下降
rewards/margins chosen/rejected reward 差
人工 win rate 并排抽检
下游任务 勿只看 DPO loss

共性局限:04 评估导读 §五


八、常见踩坑

现象 对策
变短/复读 降 beta;检查 chosen 长度
无效对齐 提高偏好对区分度 01-03
OOM 降 batch;gradient_checkpointing
ref 未冻结 显式 requires_grad=False

九、小结

DPOTrainer = SFT checkpoint + 冻结 ref + chosen/rejected 数据 + β/lr。前置 01-03 数据02-04 续训

规划:见 06-02 GRPOTrainer;选型见 05-01

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