ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

DPO偏好优化:如何用二分类替代强化学习实现大模型对齐

DPO偏好优化:如何用二分类替代强化学习实现大模型对齐 NeurIPS 2023 上有一篇让我印象特别深的工作就是这份 Direct Preference Optimization副标题叫“Your Language Model is Secretly a Reward Model”。我第一次读的时候还以为是标题党毕竟大家默认做 RLHF 就得老老实实分开训练奖励模型和策略模型怎么可能绕开整个强化学习管线。但读完之后我直接改了实验方案以前在 PPO 上调到想骂人的那套对齐流程被压缩成了一个普通的分类训练稳定且便宜。如果你是做大模型微调、用偏好数据优化模型效果、或者被 PPO 训练折腾到没脾气的研究员和工程师这篇工作很值得仔细读我会把论文的核心推导、复现心得、踩坑记录一并拆开讲。1. 先捋清 DPO 想解决的是哪一坨痛点1.1 大模型预训练之后的“对齐问题”到底是什么很多人一上来就谈 RLHF但不太清楚对齐到底对不齐的是什么。预训练阶段的语言模型核心目标就是给定上文预测下一个 token 的概率这个目标只要求模型在统计上像训练集里的文本并不关心它说出来的内容是不是用户想要的、有没有礼貌、会不会编造事实。你会发现预训练模型最大的问题不是“不会说话”而是“它不知道什么时候该闭嘴、什么时候该反驳、什么时候该给用户一个能直接用的答案”。SFT 能解决一部分问题比如让模型学会指令的对话格式、学会在问句后面跟答案但 SFT 本质是在模仿人类写出的参考回答它没有机制去区分“好的回答”和“差的回答”之间的相对偏好。同一个提示词人类写了两版回复一版清晰简洁另一版啰嗦且跑题SFT 对两者的损失可能在数值上差不多模型学不到“前者优于后者”这种排序信息。这时候就需要用偏好数据做对齐让模型学会把自己的输出分布往“人类更喜欢的回答”方向推。过去做对齐的主流方案是 RLHF流程大家应该很熟先拿人类偏好数据训练一个奖励模型再用这个奖励模型给策略模型的输出打分数最后通过强化学习去最大化期望奖励。听上去顺理成章但落地的时候你会遇到一堆麻烦。1.2 传统 RLHF 管线里最让人崩溃的四个环节我在实际工程中跑过标准的 PPO 版 RLHF体感就是“每一步单独都没那么难放在一起就是灾难”。首先是模型数量多一套完整的 PPO 管线至少要维护四个模型策略模型、参考模型、奖励模型、价值模型。价值模型还不一定是独立初始化得从奖励模型或策略模型里拉出来再调显存和显存的沟通成本非常高分布式训练时的通信量也上去了。其次是训练稳定性差PPO 对超参数极其敏感包括 KL 惩罚系数、GAE 的 lambda、clip 范围、价值损失系数甚至学习率调得不合适策略模型过两三个 step 就可能输出乱码或者开始反复说同样的句子。你很难判断训练过程中奖励分数的上升到底是模型真的变好了还是 PPO 在钻奖励模型的空子。第三个问题是奖励模型的泛化问题奖励模型本身也是用人类偏好数据训练出来的它只见过给定的一批比较换到策略模型采样出来的分布外样本时打分能力可能崩塌。你会发现策略模型经过强化学习后生成的句子逐渐偏向奖励模型的高分区域但这些句子在人类眼里未必更好这就是典型的 reward hacking。最后是工程复杂度光是写 PPO 那个多条序列并行计算 logprobs、处理优势函数、把采样的 prompt 和 response 重新打包成训练 batch 的代码就够一个工程组忙活几周。更别提在训练策略模型的时候还要用同策略采样去持续更新数据这在学术实验里跑还好放到产品迭代上速度和成本都很难接受。DPO 这篇论文最吸引我的地方就是它把这些痛点一次性抹平了。不需要奖励模型、不需要强化学习循环、不需要四个模型一起跑。它只要一份偏好数据集加上一个常规的交叉熵目标就能干成 PPO 能做到的事。2. DPO 的核心洞察把奖励函数反过来塞进策略里2.1 “Secretly a Reward Model”到底是什么意思论文标题里的副词 Secretly 是全篇最重要的关键词。作者指出在 Bradley-Terry 偏好模型这套假设下最优奖励函数其实和最优策略之间存在一一对应的闭式解关系。换句话说你根本不需要额外训练一个奖励模型来输出分数策略模型里已经隐式包含了一个奖励模型你要做的只是把这种隐含关系显式地解出来。这个思路的源头可以追溯到带 KL 约束的 RLHF 目标函数。传统做法是在给定奖励函数的前提下优化策略模型使奖励期望最大化但要约束它不要偏离原始参考模型太远也就是加一个 KL 散度惩罚项。这个约束项的系数 β 控制着“追求奖励”和“保持语言流畅度”之间的平衡。在数学上这个带约束的最优化问题存在一个解析解最优策略正比于参考策略乘以奖励的指数形式比例系数是配分函数 Z(x) 的倒数。作者看到这个闭式解之后转了半圈想既然最优策略和奖励函数的数学关系是双向的那我其实可以把奖励函数单独解出来表示成参考策略和最优策略的对数概率比值。把这个表达式代回到 Bradley-Terry 偏好概率公式里配分函数在成对比较的减法里会被约掉。于是损失函数里只剩策略模型的概率和参考模型的概率连奖励模型的影子都没有了。这一手重参数化的含金量在于偏好对齐的整个任务变成了一个标准的二分类问题。给定一个提示词和两个响应模型只需要把被偏好的响应概率推高把不被偏好的响应概率压低并且程度由两个响应在参考模型下的相对概率差来校准。路线图瞬间从“强化学习”降维成“监督学习”所有对训练稳定性的担忧都少了一大半。2.2 DPO 损失函数的数学直觉与工程意义DPO 的最终损失写出来并不长对每个偏好对计算胜出响应在策略模型和参考模型下的对数概率差再减去落败响应对应的概率差乘上 β过一层 sigmoid 再取负对数。最后形式就是-log(sigmoid(beta * (log(πθ(y_w|x)/πref(y_w|x)) - log(πθ(y_l|x)/πref(y_l|x)))))。这个式子其实非常贴近“相对奖励”的直觉。括号里面的内容衡量的是当前策略在多大程度上比参考模型更喜欢胜出响应、而不喜欢落败响应。整个结构类似逻辑回归参考模型充当了动态的基线。如果策略模型把胜出响应的概率压低了或者把落败响应抬高了损失就会变大梯度会推动概率分布回到正确方向。从工程角度讲这个损失函数还有一个很大的优点它不需要在训练过程中实时采样策略模型的输出只需要预先准备好静态的偏好数据集计算好参考模型在每条响应上的对数概率即可。也就是说训练之前就能把所有参考模型的 logprob 算出来缓存好训练时只更新策略模型的前向和反向。训练速度和 SFT 基本一样显存开销也没有额外负担。这对只有一两张卡的个人或小型团队来说简直是天上掉下来的好消息。2.3 一个直觉上的类比老师给学生改作文为了帮助不熟悉强化学习的读者理解这个结构我用改作文做个类比。传统 RLHF 像同时请了两位老师一位老师负责给作文打分另一位老师依据分数指导学生反复重写学生每次重写的时候还得确保自己没有丢掉原来会用的优美词句KL 约束。这个过程中老师有可能不一致学生也可能为了分数写出看起来华丽但实际跑题的内容。DPO 等于换了一种方式两位学生针对同一个题目各写了一版作文老师只需要告诉你“哪一版相对更好”学生的训练目标就是调整自己的写作水平让“写出好作文的概率”和“写出差作文的概率”之间的差距越来越大同时参考模型作为原来的自己确保不会改得面目全非。整个过程不再需要评分老师只需要对比反馈。现实中收集成对比偏好数据比收集精确的连续分数容易太多了这也是 DPO 在数据层面能被广泛应用的原因。3. DPO 的完整训练流程一份可直接落地的操作清单3.1 数据准备偏好对要从哪里来DPO 的输入数据是三元组结构包括提示词、被偏好的胜出响应和被抛弃的落败响应。公开数据集里最常用的有 Anthropic HH、UltraFeedback、OpenAssistant 以及斯坦福的 SHP 数据集。如果你在公司内部做对齐也可以用人工标注或者线上用户反馈来构造偏好对日志里“用户点了踩”的回复就是现成的负样本。预处理阶段最重要的一步是先用 SFT 模型把每条响应的对数概率提前算出来因为参考模型就是 SFT 模型的 frozen 副本。我建议直接在原始文本级别计算 logprob不要在 tokenization 之后做花式 padding因为 DPO 对响应长度的偏差很敏感。模板和特殊 token 的处理也要一致训练时怎么给模型拼接 prompt 和 response计算参考 logprob 时就怎么拼否则你会看到损失异常但找不出原因。3.2 模型初始化参考模型和策略模型的关系DPO 里的策略模型 πθ 和参考模型 πref 必须用同一个 SFT 模型作为初始化。这是论文原文明确要求的逻辑也很容易理解如果二者起点不一致那偏好损失本质上是在用参考模型和策略模型的固有分布差去拟合偏好信号噪声会非常大。更严格一点训练中参考模型要完整冻结只保留前向计算用来算 logprob不参与梯度更新。实践中还有一个容易忽略的细节就是参考模型最好使用和策略模型完全相同的参数副本不能是另一个阶段训练的模型。我在自己实验里试过用不同版本的同系列模型当参考模型结果 DPO 训练让策略模型的输出语感明显变差因为它在努力往一个“陌生参考模型”的方向修正自己的分布。这是个只会在细节处坑人的问题。3.3 核心训练循环一个最小实现示例DPO 的并行训练逻辑并不复杂主循环甚至和普通 SFT 一样。每轮从数据集里取 prompt、胜出响应和落败响应分别用策略模型和参考模型算出两组对数概率带上掩码对齐后求差过损失函数回传梯度更新策略模型。大致代码框架def dpo_loss(policy_chosen_logps, policy_rejected_logps, ref_chosen_logps, ref_rejected_logps, beta0.1): policy_log_ratios policy_chosen_logps - policy_rejected_logps ref_log_ratios ref_chosen_logps - ref_rejected_logps logits beta * (policy_log_ratios - ref_log_ratios) loss -torch.nn.functional.logsigmoid(logits).mean() return loss # 训练循环内 chosen_logps compute_logprobs(policy_model, prompt_ids, chosen_ids) rejected_logps compute_logprobs(policy_model, prompt_ids, rejected_ids) chosen_ref_logps compute_logprobs(ref_model, prompt_ids, chosen_ids) rejected_ref_logps compute_logprobs(ref_model, prompt_ids, rejected_ids) loss dpo_loss(chosen_logps, rejected_logps, chosen_ref_logps, rejected_ref_logps, betaconfig.beta) loss.backward() optimizer.step()这里 compute_logprobs 要注意需要把 prompt 部分排除在损失计算之外只累计 response 部分的对数概率。很多早期复现翻车都是因为在整段文本上算了 logprob把 prompt 的分布也纳入了更新目标导致生成质量不升反降。3.4 关键超参数 β 的选择逻辑β 在 DPO 里承担的是和 RLHF 中 KL 惩罚系数类似的功能控制模型向偏好方向偏离参考模型的程度。β 越小模型越倾向压制落败响应、抬高胜出响应训练更新的步子迈得越大但输出容易偏离原始风格β 越大模型越保守基本贴着参考模型的分布走偏好信号带来的改变很小。论文里给出的常用范围大致是 0.1 到 0.5 之间但具体取值要结合偏好数据的噪声水平来定。如果你的偏好数据来源比较杂标注质量也不高别把它踩得太低否则会把数据里的错误偏好强行刻进模型。我的做法是先在验证集上做小范围扫描用 0.05、0.1、0.3、0.5 几档对比以生成结果的 win rate 和人工抽样为准不要只看训练损失。4. 复现 DPO 时我要踩给各位看的几个坑4.1 偏好数据训太多轮会过拟合别把 DPO 当 SFT 死磕第一版复现的时候我沿用了 SFT 的多轮训练习惯把公开偏好集直接训了三个 epoch结果验证集上的 reward 确实一直在涨但人工看生成结果发现模型在重复训练数据里的措辞回答的覆盖面变窄了一旦给个没见过的提示词就容易说出套话。这个现象现在看很典型DPO 的目标是让胜出响应和落败响应的概率拉开当训练轮数过多时模型只需要记住哪些响应是好的就行根本不需要学会泛化。后来查论文里的实验设置发现大多数任务上作者都只训练了一个 epoch 甚至更少。原因是偏好数据集的规模本身不大几万条上下模型在海量预训练任务里已经具备充分的生成能力对齐只是微调排序偏好不需要大量重复。如果你用的是自己的高质量偏好对两三千条训一个 epoch 通常就能看到明显变化。多轮训练之后的提升大多是假象需要格外小心。4.2 参考模型 logprob 一定要缓存并且注意浮点数一致DPO 比 PPO 快的一个主要原因就是参考模型只参与静态计算。训练开始前把每一条胜出和落败响应的参考 logprob 预先算好保存下来训练时直接从内存或磁盘读取能省掉一半的前向开销。我第一次实现时偷懒在每轮实时算参考 logprob结果显存里要同时装两个全量模型batch size 被迫减半训练吞吐掉得特别厉害。缓存参考 logprob 的时候还要注意精度问题。我建议以半精度格式保存但计算时最好用高精度或至少保证 logprob 是从同一个 padding mask 下算出来的。如果训练脚本和数据预处理版本不一致导致 logprob 对应的 token 序列对不上DPO 损失的符号都会错误这种错误特别隐蔽损失曲线表面正常实际学到的却是反偏好。4.3 正负样本构造不当胜出响应未必“胜出”DPO 对偏好数据的质量极其敏感。很多入门者在构造偏好对时只用长度作为筛子比如把短的指为落败、长的指为胜出导致模型学到的是“说废话更安全”。我在实际标注里也发现人工标注的偏好对之间存在大量噪声同一个回复不同标注者可能给出相反判断。如果你的训练集规模不大我建议先做一轮一致性过滤只保留标注者一致同意的偏好对或者用更强的模型对公开数据集做一轮交叉验证剔除那些胜出响应明显劣于落败响应或者两者质量接近的数据。偏好差距越大DPO 训练的信号越强如果偏好差距太微妙模型很难从这种微弱信号中学到稳定的排序规律。4.4 别只看奖励分数必须做生成结果的手动检查任何一个“对齐方法”最终都要落到生成文本上而不是训练指标上。DPO 训练过程中你可能会看到偏好准确率接近 100%但实际生成样本在风格、多样性、安全性上表现很差。原因很简单偏好准确率只刻画了模型在训练数据里的排序能力并不直接测量“人类看到这段文本会觉得好”这个终极目标。我在每个实验节点都会固定抽 20 条 prompt不做随机采样从多个温度下生成结果人工过一遍再看有没有明显退化。采样温度为 0.7 到 1.0 时最容易暴露问题因为高温度下模型必须依赖真实分布而不是贪心路径。这个环节省不得也是我和很多“只看指标写报告”的团队最大的分歧。4.5 DPO 不是万能的它和 PPO 的适用边界DPO 的出现让没有强化学习背景的团队也能做对齐但它的局限性在论文实验里其实也有体现后续很多工作也提出了改进。最明显的一点是 DPO 只能利用静态的偏好数据它不像 PPO 那样可以在训练过程中不断让策略模型采样新样本并接受反馈。如果你的偏好数据分布和模型实际采样分布差异很大DPO 的优化效果就会打折扣。另一个值得注意的点是 DPO 对 KL 散度的控制比 PPO 温和因为它没有一个显式约束在优化每个 batch 前就在保证步长不要太大。当偏好信号非常强或者 β 挑得太小时模型依然会朝奖励区域过度优化变成“只会说好听话但丧失事实准确度”的模式。后续的 IPO、KTO、cDPO 等方法就是针对这些缺口进行的修补但 DPO 作为基准点的地位没有动摇。5. DPO 为什么能引起这么大反响以及它对训练范式的改变5.1 从“四模型强化学习”到“单模型二分类”的范式降级我在第一节说过传统 RLHF 要同时维护策略、参考、奖励、价值四套模型这还不算采样用的 rollout worker。DPO 的出现直接把奖励模型和价值模型从管线上删掉对齐训练回到只有策略和参考两条模型前向的简单代码。这在个人开发者和中小团队的环境里意义巨大意味着用一张消费级显卡就能在小模型上跑通对齐实验门槛低到可以放进教程。对于整个社区的影响最直接的变化是做 RLHF 相关研究的团队终于有了一线性价比极高的基线。以前你想验证一个新的偏好优化算法得先复现一套 RLHF 环境如今只需要在 DPO 的代码上改几十行损失函数就能开跑。这也解释了为什么一段时间里 Follow-up 工作井喷从改进正则化方式到解决偏置问题全部建立在 DPO 提供的简化框架上。5.2 我自己从 DPO 里得到的三条工程启示第一条是把复杂目标重参数化成简单目标的能力极度被低估。DPO 不改变“用偏好优化语言模型”这个任务它改变了这个任务的表达方式。很多工程问题看起来复杂是因为我们的现有表示方式太绕了一旦找到更直接的数学等价形式工程实现会瞬间变得清爽。第二条是参考点的重要性。DPO 里的参考模型不是可有可无的摆设它提供了训练中的“零点”。没有参考模型偏好概率就没有可比刻度有了参考模型模型知道自己是相对于原来的自己多喜欢胜出响应而不是绝对地无上限抬高某一条路径。这种相对更新模式在稳定性和可控性上优势明显值得在设计其他训练目标时借鉴。第三条是评估体系要跟上方法简化。方法简单了不代表“可以跑”就等于“效果好”。模型对齐领域最稀缺的还是高质量偏好数据和贴近真实使用场景的评测方式。DPO 把训练成本降下来了但如果你没有可靠的 evaluation pipeline模型到底有没有变好依然是一笔糊涂账。5.3 给准备上手 DPO 的同行一些工具建议如果你的场景是对话模型对齐建议从 7B 左右的模型开始试用公开的 UltraFeedback 或 Anthropic HH 数据集跑一个 epoch观察生成风格的变化。如果你的场景是离线偏好优化比如摘要、写作润色或者代码生成优先把偏好数据清洗好把正负样本的差距拉开这比调任何超参数都重要。在框架选型上TRL 库已经有比较成熟的 DPO Trainer 封装Hugging Face 生态里可以直接调用底层已经处理好了 logprob 掩码和缓存逻辑。如果你想深入理解实现细节强烈建议自己手写一遍 compute_logprobs搞清楚从输入文本到对数概率输出的全过程。哪怕你最后仍用封装好的库这个手写过程也会帮你排查掉大量隐蔽 bug。6. 一些想说在最后的大实话我见过不少团队拿着 DPO 跑了一轮训练看到 loss 下降就欢天喜地宣布对齐完成。但真正决定对齐质量的往往是那些最不性感的部分偏好数据干不干净、参考模型和策略模型是否真的严格一致、β 有没有针对任务调过、评测有没有人工抽查生成样本。这些细节我都在复现过程中交过学费写在这里是希望你能绕过这些弯路。DPO 真正教会我的是模型对齐本质上没有那么多高不可攀的强化学习门槛很多复杂管线里隐藏着可以被数学识破的冗余。如果你刚接触这个方向建议从手写一个能跑的 DPO 训练脚本入手跑两三个公开数据集对比训练前后模型在同一条 prompt 上的回答变化。这个过程的反馈非常及时比听任何讲座都更能帮你建立对偏好优化这件事的直观感觉。
返回列表