ARTICLE DETAIL

资讯详情

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

最小化强化学习对齐复现:用 150 行代码实现基于偏好对的 DPO 核心损失

最小化强化学习对齐复现:用 150 行代码实现基于偏好对的 DPO 核心损失 最小化强化学习对齐复现用 150 行代码实现基于偏好对的 DPO 核心损失在大语言模型后训练Post-training的对齐阶段基于人类反馈的强化学习RLHF曾长期统治整个学术界。然而经典的近端策略优化PPO框架在工程落地中极其沉重它需要同时在显存中维护 Actor策略模型、Critic价值模型、Reward奖励模型以及 Reference参考基线模型共四个大型神经网络。四模型协同前向与反向传播不仅带来了极其恐怖的显存占用更让超参数调节变成了玄学调优策略梯度的极高方差经常导致模型在训练中途发生灾难性崩溃。斯坦福大学提出的直接偏好优化Direct Preference Optimization, DPO彻底改变了这一格局。DPO 通过精妙的数学变换证明了受 KL 散度约束的强化学习最优策略与隐式奖励函数之间存在精确的解析映射从而完全绕过了独立的奖励模型与强化学习循环将复杂的 RL 对齐降维为一个优雅、稳定的单阶段二元分类损失。本文从 Bradley-Terry 偏好模型的数学推导出发剥离所有冗余封装框架用 150 行纯粹的 PyTorch 原生代码实现 DPO 核心损失与训练步带你彻底看透大模型偏好对齐的数学底层。一、DPO 数学推导与隐式奖励重构理解 DPO 为何能干掉 PPO关键在于理解其如何将显式的 Reward Model 替换为模型自身的对数似然比。在受约束的 RL 目标函数中我们的目标是最大化期望奖励同时惩罚新策略 $\pi_\theta$ 偏离冻结参考策略 $\pi_{ref}$ 的 KL 散度$$\max_{\pi} \mathbb{E}{x \sim \mathcal{D}, y \sim \pi(y|x)} [r(x, y)] - \beta , \mathbb{D}{KL}(\pi(y|x) \parallel \pi_{ref}(y|x))$$对该拉格朗日目标求解变分极值可以推导出最优策略 $\pi^*$ 的封闭解析解$$\pi^*(y|x) \frac{1}{Z(x)} \pi_{ref}(y|x) \exp\left( \frac{1}{\beta} r(x, y) \right)$$其中 $Z(x) \sum_y \pi_{ref}(y|x) \exp\left( \frac{1}{\beta} r(x, y) \right)$ 为配分函数。两边取对数并移项我们赫然发现任意响应 $y$ 对应的真实奖励值 $r(x, y)$可以被严格表示为策略模型与参考模型输出概率的对数比值加上一个仅与输入 $x$ 相关的标量项$$r(x, y) \beta \log \frac{\pi^*(y|x)}{\pi_{ref}(y|x)} \beta \log Z(x)$$现在我们将这个隐式奖励函数代入人类偏好的 Bradley-Terry 模型中。对于同一个提示词 $x$人类更偏好获胜响应 $y_w$Winning而非失败响应 $y_l$Losing的后验概率为$$P(y_w \succ y_l \mid x) \sigma(r(x, y_w) - r(x, y_l))$$由于 $y_w$ 和 $y_l$ 面对的是相同的提示词 $x$配分函数项 $\beta \log Z(x)$ 在做差时被精准对消最终推导出纯粹依赖模型生成概率的 DPO 目标损失函数$$\mathcal{L}{DPO}(\theta) - \mathbb{E}{(x, y_w, y_l) \sim \mathcal{D}} \left[ \log \sigma \left( \beta \log \frac{\pi_\theta(y_w \mid x)}{\pi_{ref}(y_w \mid x)} - \beta \log \frac{\pi_\theta(y_l \mid x)}{\pi_{ref}(y_l \mid x)} \right) \right]$$对齐架构维度传统 PPO (RLHF)直接偏好优化 (DPO)在线驻留模型4 个 (Actor, Critic, Reward, Ref)2 个 (Actor, 冻结 Ref)优化本质策略梯度采样与时序差分拟合监督式的二元交叉熵分类显存与硬件门槛极高需要专用多卡流水线并行适中等同于常规 SFT 显存消耗训练稳定性极差极易发生梯度发散与崩溃极佳损失平滑收敛零模式崩溃二、DPO 损失的物理机制解构仔细审视 DPO 损失项内部的隐式差值 $\hat{r}\theta(x, y) \beta \log \frac{\pi\theta(y \mid x)}{\pi_{ref}(y \mid x)}$当策略模型 $\pi_\theta$ 提升了获胜文本 $y_w$ 的生成概率时第一项变大当它压低了失败文本 $y_l$ 的生成概率时第二项变小。两者之差通过 Sigmoid 函数映射后损失函数促使这一差值趋近于正无穷。超参数 $\beta$ 在其中扮演着“惩罚系数Temperature/KL Constraint”的角色$\beta$ 越大模型偏离参考策略 $\pi_{ref}$ 遭到的梯度反噬越大对齐倾向于保守$\beta$ 越小模型越敢于大刀阔斧地推高 $y_w$ 并打压 $y_l$但容易引发概率分布畸变与过拟合。在实际工程中$\beta$ 通常严格收敛在 0.05 到 0.2 的极窄区间内。三、150 行纯 PyTorch 核心实现以下代码实现了工业级 DPO 训练的核心逻辑包含带掩码的对数似然精确提取、隐式奖励差值构建、数值稳定的交叉熵损失计算以及梯度步迭代import torch import torch.nn as nn import torch.nn.functional as F from typing import Tuple, Dict def compute_sequence_logps( model: nn.Module, input_ids: torch.Tensor, attention_mask: torch.Tensor, labels: torch.Tensor ) - torch.Tensor: 计算模型在有效标签位置上的对数似然累加和 input_ids, labels: 形状为 (batch_size, seq_len) # 获取自回归 logits: (batch_size, seq_len, vocab_size) logits model(input_ids, attention_maskattention_mask).logits # 自回归移位第 t 个词的 logit 预测第 t1 个词的 label shift_logits logits[:, :-1, :].contiguous() shift_labels labels[:, 1:].contiguous() # 计算全量词表的对数概率 log_probs F.log_softmax(shift_logits, dim-1) # 提取真实标签位置处的对数概率 # loss_mask 用于屏蔽 prompt 部分以及 padding 补齐部分 (值为 -100) loss_mask (shift_labels ! -100) # 将 -100 临时替换为 0 避免 gather 越界 safe_labels shift_labels.clone() safe_labels[~loss_mask] 0 per_token_logps torch.gather( log_probs, dim-1, indexsafe_labels.unsqueeze(-1) ).squeeze(-1) # 仅累加有效 response 的 token 对数概率 sequence_logps (per_token_logps * loss_mask).sum(dim-1) return sequence_logps class DPOLoss(nn.Module): def __init__(self, beta: float 0.1, label_smoothing: float 0.0): super().__init__() self.beta beta self.label_smoothing label_smoothing def forward( self, policy_chosen_logps: torch.Tensor, policy_rejected_logps: torch.Tensor, reference_chosen_logps: torch.Tensor, reference_rejected_logps: torch.Tensor, ) - Tuple[torch.Tensor, Dict[str, float]]: 核心 DPO 损失与隐式奖励监控指标计算 # 计算策略模型与参考模型在 chosen 样本上的对数概率比 pi_logratios policy_chosen_logps - policy_rejected_logps ref_logratios reference_chosen_logps - reference_rejected_logps # 隐式奖励对齐项beta * (log(pi_w / ref_w) - log(pi_l / ref_l)) logits self.beta * (pi_logratios - ref_logratios) # 带平滑的二元交叉熵损失 if self.label_smoothing 0.0: losses ( -F.logsigmoid(logits) * (1 - self.label_smoothing) - F.logsigmoid(-logits) * self.label_smoothing ) else: losses -F.logsigmoid(logits) loss losses.mean() # 计算隐式奖励用于实时监控 chosen_rewards self.beta * (policy_chosen_logps - reference_chosen_logps).detach() rejected_rewards self.beta * (policy_rejected_logps - reference_rejected_logps).detach() reward_margin (chosen_rewards - rejected_rewards).mean().item() accuracy (logits 0).float().mean().item() metrics { loss: loss.item(), reward_margin: reward_margin, chosen_reward_mean: chosen_rewards.mean().item(), rejected_reward_mean: rejected_rewards.mean().item(), accuracy: accuracy, } return loss, metrics def train_dpo_step( policy_model: nn.Module, ref_model: nn.Module, optimizer: torch.optim.Optimizer, dpo_criterion: DPOLoss, batch: Dict[str, torch.Tensor], device: torch.device ) - Dict[str, float]: 单步 DPO 优化执行器 policy_model.train() ref_model.eval() # 1. 搬运数据至 GPU c_ids, c_mask, c_labels batch[chosen_input_ids].to(device), batch[chosen_mask].to(device), batch[chosen_labels].to(device) r_ids, r_mask, r_labels batch[rejected_input_ids].to(device), batch[rejected_mask].to(device), batch[rejected_labels].to(device) # 2. 计算策略模型梯度 optimizer.zero_grad() pol_chosen_logps compute_sequence_logps(policy_model, c_ids, c_mask, c_labels) pol_rejected_logps compute_sequence_logps(policy_model, r_ids, r_mask, r_labels) # 3. 冻结参考模型无梯度计算基线似然 with torch.no_grad(): ref_chosen_logps compute_sequence_logps(ref_model, c_ids, c_mask, c_labels) ref_rejected_logps compute_sequence_logps(ref_model, r_ids, r_mask, r_labels) # 4. 计算 DPO 损失与反向传播 loss, metrics dpo_criterion(pol_chosen_logps, pol_rejected_logps, ref_chosen_logps, ref_rejected_logps) loss.backward() # 梯度截断防止爆炸 torch.nn.utils.clip_grad_norm_(policy_model.parameters(), max_norm1.0) optimizer.step() return metrics四、复现实战中的三大关键暗坑在将上述核心代码接入大模型训练集群时以下三个工程细节直接决定了模型是否会发生退化1. 似然置换Likelihood Displacement隐患仔细观察 DPO 损失项它关注的是两者的相对差值。这意味着模型存在一种投机取巧的捷径不提高获胜样本 $y_w$ 的概率而是拼命降低失败样本 $y_l$ 的概率。在极端情况下策略模型会把 $y_l$ 中常见词汇的概率压制到接近零从而拉大差值。其宏观表现是DPO 训练准确率不断上升但模型的困惑度Perplexity剧烈恶化生成文本出现语法退化。在实际训练中必须密切监控chosen_reward_mean是否保持在正值区间。若发现chosen_reward_mean持续下跌且全靠rejected_reward_mean下坠支撑应当立即在损失中增加微弱的 SFT 负对数似然正则项NLL Loss。2. 参考模型显存节约技巧参考模型在整个训练过程中参数完全冻结且只参与前向传播。如果直接加载两个完整的 FP16 模型显存消耗直接翻倍。显存优化策略使用 BitsAndBytes 将参考模型以 NF44-bit NormalFloat量化载入显存或者在预先离线遍历数据集时将所有偏好对在 $\pi_{ref}$ 下的ref_chosen_logps和ref_rejected_logps预先计算并固化到 Parquet 缓存文件中。这样在正式训练时显存中只需驻留一个唯一的策略模型算力与显存开销直接直降 50%。3. 超参数 $\beta$ 的衰减陷阱许多初学者将 $\beta$ 设为 0.5 甚至 1.0导致隐式奖励项的梯度极其微弱对齐进展极其缓慢而若直接降到 0.01模型仅需经过 50 个 Step 就会对偏好集发生严重过拟合。实测表明在 7B 到 32B 规模模型上以 $\beta 0.1$ 起步配合余弦学习率衰减峰值学习率控制在 $5 \times 10^{-7}$ 到 $1 \times 10^{-6}$是保障对齐稳定性的黄金组合。
返回列表