跳到主要内容

数据截至 (上游 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 到一行损失

推导链(直觉版,三步):

  1. 偏好建模:用 Bradley-Terry 模型描述「人更偏好 y_w 而非 y_l」的概率——两回答奖励之差的 sigmoid。
  2. 闭式最优解:带 KL 约束的 RLHF 目标有闭式最优策略,反解出「奖励 = β·log(π/π_ref) + 常数」。也就是说,策略模型自己和参考模型的对数概率比,就是一个隐式奖励,不需要单独训奖励模型。
  3. 代入即得损失:把这个隐式奖励代回 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 维拼成一个 batchtrl/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_losstrl/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_softmaxtrl/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-1460reverse_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_pathdpo_trainer.py:934一份完整模型显存
PEFT/LoRA不加载:前向时临时禁用 adapter,基座模型即参考模型(dpo_trainer.py:1404-1410零额外显存
任何情形precompute_ref_log_probs=True:训练前把整个数据集的 ref logp 预存进数据集列(_precompute_ref_logpsdpo_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_alphatrl/trainer/dpo_trainer.py:1381-1393):chosen/rejected 共享前缀长度之外的「尾巴」按 α 打折,实现 LDD(length-desensitization)——防止长度差主导对数比。


6. KTO:没有「成对偏好」,只有「好/坏」

6.1 小问题

DPO 要「同一 prompt 两个回答分胜负」。现实中更多数据长这样:一个回答 + 一个二元标签(点赞/点踩)。KTO(Kahneman-Tversky Optimization,2402.01306)用前景理论的价值函数替代偏好对。

6.2 实现怎么对应

  • 数据prompt / completion / label(bool),collator 是 DataCollatorForUnpairedPreferencetrl/trainer/kto_trainer.py:90)。
  • KL 项的估计很取巧:论文需要一个参考点 KL(πθ‖π_ref)。TRL 在 collator 里把 batch 内的 completion 错位循环配对(cycling),拿「别人的 completion」当错配样本算 KL(trl/trainer/kto_trainer.py:223 的说明)。
  • 损失非对称trl/trainer/kto_trainer.py:1577:1594):
# 摘自 trl/trainer/kto_trainer.py:1577, 1594(loss_type="kto")
chosen_losses = 1 - F.sigmoid(self.beta * (chosen_logratios - kl)) # 好样本: 超过 KL 基准越多越好
rejected_losses = 1 - F.sigmoid(self.beta * (kl - rejected_logratios)) # 坏样本: 低于 KL 基准越多越好

好样本和坏样本的损失以 kl 为锚点方向相反,最后按 desirable_weight/undesirable_weight 加权拼接(trl/trainer/kto_trainer.py:1603-1606)——两类样本数量悬殊时靠这两个权重平衡。

6.3 同族在 experimental

ORPO(把 SFT 和偏好惩罚合进一个无参考模型的损失,trl/experimental/orpo/orpo_trainer.py:95)、CPO(trl/experimental/cpo/cpo_trainer.py:83)、BCO 都在 experimental 层——它们都试图进一步砍掉参考模型,但接口和实现还在迭代,故不在稳定层。


7. 关键细节与坑

  • 拼批约定是隐式的。 collator 拼、loss 里 chunk(2) 切,两处必须同步改。自定义 collator 时这是第一个会踩的雷。
  • 前缀 token 化不匹配只打 warningtrl/trainer/dpo_trainer.py:1073-1085):如果 tokenizer 对 promptprompt+chosen 的切分不一致,completion_mask 会悄悄错位。看到这个 warning 要停下来查。
  • modelref_model 不能是同一个对象dpo_trainer.py:589-591 直接报错)——那样对数比恒为 0,loss 永远是 log 2,训练假装在跑。
  • sync_ref_model 与 PEFT、预计算互斥dpo_trainer.py:964-979):同步式参考模型(EMA 更新)只支持全量微调。
  • β 是最敏感的旋钮。 它直接缩放任一方向的漂移;loss_type="ipo" 时它变成正则系数 τ,语义不同,换 loss 时别把旧 β 带过去。

8. 代码地图(本章)

主题文件路径符号名
DPO 损失主流程trl/trainer/dpo_trainer.pyDPOTrainer._compute_loss
偏好数据拼批trl/trainer/dpo_trainer.pyDataCollatorForPreference.torch_call
偏好数据 tokenizetrl/trainer/dpo_trainer.pytokenize_fn_prepare_dataset 内)
参考模型加载/同步trl/trainer/dpo_trainer.pyDPOTrainer.__init__SyncRefModelCallback
预计算 ref logptrl/trainer/dpo_trainer.pyDPOTrainer._precompute_ref_logps
PEFT 参考模型技巧trl/trainer/utils.pyuse_adapter
损失变体清单trl/trainer/dpo_config.pyDPOConfig.loss_type
f-散度变换trl/trainer/dpo_trainer.py_compute_lossf_divergence_type 分支
KTOtrl/trainer/kto_trainer.pyKTOTrainer.forwardDataCollatorForUnpairedPreference
ORPO / CPOtrl/experimental/orpo/orpo_trainer.pytrl/experimental/cpo/cpo_trainer.pyORPOTrainerCPOTrainer