DPO 从零开始 本节摘要:奖励模型加 PPO 是经典 RLHF 栈。直接偏好优化(Direct Preference Optimization, DPO)把这个栈塌缩成单一监督损失,直接在偏好对上拟合策略。本节从奖励差恒等式推导 DPO 损失,交付能工作的「参考模型 + 策略模型」对,逐 token 算对数概率,在偏好夹具(chosen/rejected 补全)上训微型 transformer,用测试钉住损失数学与梯度方向。损失是一个 sigmoid,作用在每例一个标量上,该标量由四个对数概率算出:无独立奖励模型、无 PPO、无显式 KL 项(KL 约束已烤进闭式推导)。训练把 chosen 的对数概率推上、把 rejected 的推下,参考模型冻住不动——参考是你的安全网。
本节摘要:奖励模型加 PPO 是经典 RLHF 栈。直接偏好优化(Direct Preference Optimization, DPO)把这个栈塌缩成单一监督损失,直接在偏好对上拟合策略。本节从奖励差恒等式推导 DPO 损失,交付能工作的「参考模型 + 策略模型」对,逐 token 算对数概率,在偏好夹具(chosen/rejected 补全)上训微型 transformer,用测试钉住损失数学与梯度方向。损失是一个 sigmoid,作用在每例一个标量上,该标量由四个对数概率算出:无独立奖励模型、无 PPO、无显式 KL 项(KL 约束已烤进闭式推导)。训练把 chosen 的对数概率推上、把 rejected 的推下,参考模型冻住不动——参考是你的安全网。
对应原课程:Phase 19 · Lesson 40 ·
dpo-from-scratch(原英文phases/19-capstone-projects/40-dpo-from-scratch/docs/en.md)。本节属「从零构建 GPT」赛道第十一节。
阅读完本节,你应当能够:
(prompt, chosen, rejected) 三元组上训策略,看着 chosen 对数概率相对 rejected 上升。你有个 SFT 模型,能听指令,但输出参差:有的补全清晰,有的啰嗦或错。你还有一个小偏好数据集:同一 prompt,人类标一个补全为 chosen、另一个为 rejected。
经典 RLHF 答案是两阶段流水线:先在偏好上训奖励模型,再用 PPO 对奖励优化策略。这有效但贵:PPO 时两个模型同时在内存、KL 控制把策略钉在参考附近、奖励模型脆时奖励黑客。
从 Bradley-Terry 模型开始。给定 prompt x 与两个补全 y_w(chosen)、y_l(rejected),人类偏好 y_w 的概率是 P(y_w > y_l | x) = sigmoid(r(x,y_w) - r(x,y_l)),其中 r 是某隐奖励函数。RLHF 先从偏好拟合 r,再训策略 pi 最大化 r 带 KL 锚:max_pi E[r(x,y)] - beta*KL(pi||pi_ref)。
DPO 推导观察到:该目标下的最优策略 pi* 有关于 r 的闭式:pi*(y|x) = (1/Z(x)) * pi_ref(y|x) * exp(r(x,y)/beta)。反解 r 得 r(x,y) = beta*(log pi*(y|x) - log pi_ref(y|x)) + beta*log Z(x)。代入偏好差,Z(x) 项抵消:
r(x,y_w) - r(x,y_l) = beta * (log pi_theta(y_w|x) - log pi_ref(y_w|x) - log pi_theta(y_l|x) + log pi_ref(y_l|x))
代回 Bradley-Terry sigmoid,对偏好对取负对数似然:
L_DPO(theta) = -E[ log sigmoid( beta * (log pi_theta(y_w|x) - log pi_ref(y_w|x) - log pi_theta(y_l|x) + log pi_ref(y_l|x)) ) ]
这就是损失——一个 sigmoid,作用在每例一个标量上,由四个对数概率算出。无独立奖励模型、无 PPO、损失里无 KL 项(KL 已烤进闭式)。
对 log pi_theta(y_w|x) 求梯度:d L_DPO / d log pi_theta(y_w|x) = -beta*(1 - sigmoid(z)),对所有 z 为负——提升策略对 chosen 的对数概率降损失。对称地,对 log pi_theta(y_l|x) 的梯度为正——提升 rejected 对数概率升损失。训练推 chosen 上、rejected 下,参考冻住不动。
数据是 12 个偏好三元组 (prompt, chosen, rejected),chosen 短而精,rejected 啰嗦、跑题或错,覆盖与第 37 节同族任务(首都、算术、列表),使从 SFT 基线出发的策略有合理起点。夹具刻意小:DPO 生产用几万对,这里重点是损失数学与循环在小数据上端到端跑通、chosen-rejected 对数概率差可见增长。
参考不变性三性质必须成立:参考参数永不收梯度、参考对数概率跨 epoch 不变、策略从与参考相同的权重起步(最优 theta 是参考加一个学到的更新,把策略初始化成参考的拷贝是良定义起点)。实现上:参考前向包在 torch.no_grad()、参考每参数 requires_grad=False、策略用 policy.load_state_dict(reference.state_dict()) 构建。
def sequence_log_prob(model, prompt_ids, completion_ids, no_grad=False): ctx = torch.no_grad() if no_grad else contextlib.nullcontext() with ctx: ids = torch.cat([prompt_ids, completion_ids], dim=-1) logits = model(ids) # (1, T, V) # 只在 completion 位算下一 token 对数概率 start = prompt_ids.shape[-1] logp = F.log_softmax(logits, dim=-1) comp_logp = logp[:, start-1:-1, :] # 预测 completion 各位 target = completion_ids return comp_logp.gather(-1, target.unsqueeze(-1)).sum() def dpo_loss(lp_w, lp_l, lr_w, lr_l, beta): z = beta * ((lp_w - lr_w) - (lp_l - lr_l)) # 四对数概率 -> 标量 return -F.logsigmoid(z) # 闭式,无 KL 项
训练循环每 epoch 算 policy 与 reference 下 chosen/rejected 的对数概率,施加损失,步 Adam:
def train_dpo(policy, reference, loader, beta=0.1, steps=30): opt = torch.optim.Adam(policy.parameters(), lr=5e-4) for s in range(steps): prompt, chosen, rejected = next(loader) lpw = sequence_log_prob(policy, prompt, chosen) lpl = sequence_log_prob(policy, prompt, rejected) lrw = sequence_log_prob(reference, prompt, chosen, no_grad=True) lrl = sequence_log_prob(reference, prompt, rejected, no_grad=True) loss = dpo_loss(lpw, lpl, lrw, lrl, beta).mean() loss.backward(); opt.step(); opt.zero_grad() print("margin", (lpw - lpl).item()) # chosen-rejected 差应增长
设计要点:参考前向一定要
no_grad——否则 autograd 把参考参数也建进图、反向时报错或悄悄更新参考。策略与参考共享架构但权重分开:策略从参考拷贝起步,训练时漂移、参考不动。evaluate_margins返回任意时刻策略下 chosen-rejected 对数概率差的均值,训练时它应单调上升。
HuggingFace trl 的 DPOTrainer 把参考管理、对数概率计算、损失打包:DPOConfig(beta=0.1) + trainer.train()。本节手写让你看清四对数概率怎么算、sigmoid 闭式怎么落、参考不变性怎么守。与 RLHF(PPO + 奖励模型)相比:DPO 同解(在 Bradley-Terry 下)、代码少一个数量级、内存少一个模型,但牺牲了对「在线采样新补全」的灵活性。变体众多:IPO 用 (z-1)^2 替 sigmoid+log 更稳;SimPO 去掉参考模型、用长度归一;KTO 用非偏好单点信号。本节的 dpo_loss 与 sequence_log_prob 是这些变体的共同底座。
main.py + 测试:钉住损失数学(给定四对数概率算出预期标量)、梯度符号(chosen 梯度负、rejected 正)、参考不变性(跨 epoch 参考对数概率不变)。demo 从小预热预训建参考与策略、拷权重、训 30 步、打每步损失与 margin,成功退零。sequence_log_prob、dpo_loss、train_dpo 可复用——换到任何因果 LM 上,只要模型有「输入 + 错位目标」接口,DPO 循环原样工作。
(z-1)^2,对比在夹具上的收敛。beta 从 0.01 到 1.0,绘最终 margin 与训练稳定性。Z(x) 在差里抵消。beta*((lpw-lrw)-(lpl-lrl)) 标量上,由四个 log-prob 算出。pi_ref 使对数比变大、sigmoid 饱和、梯度受阻。下一节,我们做「评估流水线」——把困惑度、精确匹配、token F1、判官四个异质评估聚合成单一加权报告。