数据截至 (上游 commit 67dfbe211a07)
03 · DPO 系:偏好优化的损失推导与实现
这一章讲什么: DPO 是 TRL 里最能体现「读论文公式 → 读代码」对应关系的 Trainer。先做损失推导、建立直觉,再看实现里的三个工程亮点(拼批、参考模型、损失动物园),最后顺带看只需二元标注的 KTO 和被挪进 experimental 的 ORPO/CPO。
1. 它要解决的小问题
经典 RLHF 的对齐流程很重:先训一个奖励模型,再用 PPO 在线采样、边生成边优化。这套流程有四个模型同时在场、调参困难。
DPO(Direct Preference Optimization,2305.18290)的观察是:既然奖励模型的最终用途只是给 PPO 提供信号,能不能跳过它,直接在偏好数据上优化策略? 答案是可以,而且 loss 简单到像监督学习。
输入数据也随之变轻:每条样本是一个三元组 prompt / chosen / rejected(同一个 prompt 的两个回答,标注哪个更好)。
2. 思路:从 Bradley-Terry 到一行损失
推导链(直觉版,三步):
- 偏好建模:用 Bradley-Terry 模型描述「人更偏好 y_w 而非 y_l」的概率——两回答奖励之差的 sigmoid。
- 闭式最优解:带 KL 约束的 RLHF 目标有闭式最优策略,反解出「奖励 = β·log(π/π_ref) + 常数」。也就是说,策略模型自己和参考模型的对数概率比,就是一个隐式奖励,不需要单独训奖励模型。
- 代入即得损失:把这个隐式奖励代回 Bradley-Terry,对偏好数据做最大似然:
L = -log σ( β · [ (log πθ(y_w|x) - log π_ref(y_w|x)) - (log πθ(y_l|x) - log π_ref(y_l|x)) ] )
└──────── chosen 的对数比 ────────┘ └──────── rejected 的对数比 ───────┘
白话读法:chosen 回答相对参考模型「变多」的程度,要大于 rejected「变多」的程度;差距越大损失越小。β 控制允许偏离参考模型多远。
原理演示
# 示意,非源码:DPO 损失的全部逻辑
def dpo_loss(policy_chosen_logp, policy_rejected_logp,
ref_chosen_logp, ref_rejected_logp, beta):
chosen_logratio = policy_chosen_logp - ref_chosen_logp # 变好多少
rejected_logratio = policy_rejected_logp - ref_rejected_logp
delta = chosen_logratio - rejected_logratio # 净偏好差距
return -F.logsigmoid(beta * delta)
重点看:不需要奖励模型、不需要在线采样,两对 logp 就够。下面看真实代码里它们怎么来。
3. 真实实现:一次前向算两边
3.1 拼批:chosen 在前半,rejected 在后半
最直接的做法是对 chosen 和 rejected 各跑一次前向。TRL 不这么做——DataCollatorForPreference.torch_call 把两者沿 batch 维拼成一个 batch(trl/trainer/dpo_trainer.py:176-178):
# 摘自 trl/trainer/dpo_trainer.py:176-178
input_ids = prompt_chosen_ids + prompt_rejected_ids # list 拼接 = batch 维拼接
attention_mask = chosen_attention_mask + rejected_attention_mask
completion_mask = chosen_mask + rejected_mask
于是一次 model(**inputs) 同时算出两边的 logits。completion_mask 把 prompt 段标 0、回答段标 1。
3.2 从 logits 到 log 比
_compute_loss(trl/trainer/dpo_trainer.py:1359)的核心四步:
# 摘自 trl/trainer/dpo_trainer.py:1374-1394
shift_logits = outputs.logits[..., :-1, :] # 因果 LM 错位
shift_labels = input_ids[..., 1:]
per_token_logps = selective_log_softmax(shift_logits, shift_labels) # 只取目标 token 的 logp
per_token_logps[shift_completion_mask == 0] = 0.0 # prompt 段清零
...
logps = per_token_logps.sum(dim=1) # 序列级 logp
chosen_logps, rejected_logps = logps.chunk(2, dim=0) # 拼批在这里切开
selective_log_softmax(trl/trainer/utils.py:480)避免物化log_softmax的全词表中间张量——和第 2 章的分块 CE 同一类省显存手法。chunk(2)依赖 §3.1 的拼批约定:前半 chosen、后半 rejected。改 collator 而不改这里的切分,是最经典的静默 bug 来源。
3.3 损失本体
对数比算完后,先过一个可选的 f-散度变换(trl/trainer/dpo_trainer.py:1430-1460:reverse_kl 是标准 DPO,另有 forward_kl/js_divergence/alpha_divergence,对应广义偏好优化论文的变换 f′),再进入 loss 分支:
# 摘自 trl/trainer/dpo_trainer.py:1464-1465(标准 DPO)
if loss_type == "sigmoid":
per_sequence_loss = -F.logsigmoid(self.beta * delta_score)
和 §2 的演示逐字对应。delta_score 就是 chosen 与 rejected 两个「散度变换后的对数比」之差。
4. 参考模型的三种获得方式
DPO 需要一个冻结的参考模型算 π_ref。TRL 按你的训练方式给三条路(trl/trainer/dpo_trainer.py:919-962):
| 情形 | 参考模型怎么来 | 代价 |
|---|---|---|
| 全量微调 | 用同一模型 id 再加载一份副本(create_model_from_path,dpo_trainer.py:934) | 一份完整模型显存 |
| PEFT/LoRA | 不加载:前向时临时禁用 adapter,基座模型即参考模型(dpo_trainer.py:1404-1410) | 零额外显存 |
| 任何情形 | precompute_ref_log_probs=True:训练前把整个数据集的 ref logp 预存进数据集列(_precompute_ref_logps,dpo_trainer.py:1133) | 一次性离线计算 |
PEFT 那条路最值得记住,它是全库通用的技巧(GRPO 也在用),核心是 use_adapter 上下文管理器(trl/trainer/utils.py:1225):
# 摘自 trl/trainer/dpo_trainer.py:1406-1410
model = self.accelerator.unwrap_model(model)
with use_adapter(model, adapter_name="ref" if "ref" in model.peft_config else None):
ref_outputs = self.model(**ref_model_kwargs)
- 新训 adapter 时
adapter_name=None:PEFT 的disable_adapter上下文里,前向走的就是纯基座 = 参考模型。 - 继续训一个已有 adapter 时:初始化会把初始权重存成名为
"ref"的额外 adapter,这里切过去——保证参考点是「训练开始前」而非「裸基座」。
5. 一个参数装下的论文动物园
DPOConfig.loss_type 接受 15 种值(trl/trainer/dpo_config.py:79-81),每种对应一篇后续论文对 sigmoid 损失的改造:
| loss_type | 论文思路一句话 |
|---|---|
sigmoid(默认) | 原始 DPO |
hinge | 把 log-sigmoid 换成 hinge,SLiM-HF |
ipo | 防过拟合:直接回归偏好差距到 1/(2β),带长度归一化 |
nca_pair / bco_pair / exo_pair | 换个偏好似然形式 |
robust | 假设标注含噪(label smoothing 进损失) |
apo_zero / apo_down | 分「模型本就偏差 chosen / 偏坏 chosen」两种锚点 |
sppo_hard / aot / aot_unpaired / discopop / sigmoid_norm | 各自的 margin/排序/归一化改造 |
sft | 退化为对 chosen 做 SFT(RPO 的组成项) |
两个实现层面的看点:
- 多损失组合:
loss_type是列表,loss_weights给权重,加权求和(trl/trainer/dpo_trainer.py:1462-1463的循环)。MPO 论文就是["sigmoid", "bco_pair", "sft"]这样的组合。 - IPO 的长度归一化是代码里的隐藏知识:
trl/trainer/dpo_trainer.py:1470-1476的注释明说——论文没写,但 token 求和会让平方损失随回答长度缩放,所以 TRL 按回答长度归一,并且「与 IPO 作者确认过,论文结果对应归一化形式」。读代码才能捡到的实现细节。
另一个工程点是 ld_alpha(trl/trainer/dpo_trainer.py:1381-1393):chosen/rejected 共享前缀长度之外的「尾巴」按 α 打折,实现 LDD(length-desensitization)——防止长度差主导对数比。