27届大模型岗面试准备(七):RLHF 三阶段全解——从 PPO 到 DPO 的原理、对比与手写损失

前面两篇我们把预训练和 SFT 讲完了。走到这一步,模型已经"会说人话",但还谈不上"说得让人满意"。SFT 只教了模型模仿标注数据的分布,没有教它区分"好回答"和"更好的回答"。真正把 ChatGPT 和普通指令模型拉开差距的,是对齐阶段——也就是面试里被问烂了但依然年年必考的 RLHF。

这一篇按面试深度要求展开:RLHF 三阶段的完整链路、PPO 在 LLM 场景下的具体形态、DPO 为什么能把 RL 干掉、两者怎么选,最后手写一个可运行的 DPO 损失来验证理解。这条线在27届面试里的出现频率极高,尤其是"DPO 的损失函数怎么推出来的"这种题,答不出推导思路基本就凉一半。

一、为什么 SFT 之后还需要 RLHF

先回答一个面试官爱用的开场题:SFT 已经用高质量数据微调过了,为什么还要 RLHF?

三个层面的原因:

第一,SFT 的监督信号是"逐 token 模仿",不是"整体偏好"。 SFT 的交叉熵损失对每个 token 一视同仁,模型学到的是"标注员在这个位置写了什么词",而不是"这个回答整体上好在哪"。两个回答可能 token 级别差异很小,但一个有事实错误、一个没有——SFT 损失几乎无法区分。

第二,"写出好回答"比"判断哪个回答好"难得多。 让标注员从零写一个完美回答,成本高且上限受限于标注员水平;但给两个回答让人选哪个更好,又快又准。RLHF 的本质就是把这种"判别式的人类偏好"转化为训练信号,突破 SFT 的示范上限。

第三,负反馈的缺失。 SFT 只有正样本,模型从没被告知"什么不该说"。有害内容、幻觉、废话连篇,这些都需要负向信号来抑制,而偏好数据天然携带负样本(被拒绝的那个回答)。

把这三点说清楚,比背"对齐人类价值观"这种空话强十倍。

二、RLHF 三阶段完整链路

经典 RLHF(InstructGPT 论文范式)分三个阶段,面试要求能画出数据流:

阶段一:SFT。 用高质量指令数据监督微调,得到 π_SFT。这是后续所有阶段的初始化点,上一篇讲过,不展开。

阶段二:训练奖励模型(Reward Model, RM)。 采一批 prompt,用 SFT 模型对每个 prompt 生成多个回答(InstructGPT 是 4~9 个),标注员对回答排序。把排序拆成两两偏好对 (chosen, rejected),用 Bradley-Terry 模型建模偏好概率:

P(chosen ≻ rejected) = σ(r(x, y_c) − r(x, y_r))

RM 的损失就是最大化这个概率的负对数:

L_RM = −E[ log σ(r(x, y_c) − r(x, y_r)) ]

RM 通常从 SFT 模型初始化,把 LM head 换成一个输出标量的 value head,取最后一个 token 位置的输出作为整句的分数。面试追问点:为什么用排序而不是打绝对分? 因为人类打绝对分的一致性很差(同一个回答不同人打 6 分和 8 分很常见),但两两比较的一致性高得多,排序信号更干净。

阶段三:PPO 强化学习。 把语言模型生成视为一个 RL 问题:状态是"prompt + 已生成的 token 序列",动作是"生成下一个 token",奖励由 RM 在句末给出。用 PPO 优化以下目标:

maximize E[ r(x, y) − β·KL(π_θ(y|x) ‖ π_ref(y|x)) ]

其中 KL 惩罚项防止策略跑得离参考模型(通常是 SFT 模型)太远。这个 KL 项是面试必考细节:没有它,模型会迅速学会欺骗 RM——生成一些 RM 打高分但人类看来是乱码的文本,这就是 reward hacking。

三、PPO 在 LLM 场景下的工程形态

很多候选人能背出 PPO 的 clip 公式,却答不上"PPO 训练时显存里要放几个模型"。这才是区分度所在。

PPO 训练需要同时维护四个模型:

模型 作用 是否更新参数 显存占用
Actor(策略模型) 生成回答,被优化的对象 权重+梯度+优化器状态
Critic(价值模型) 估计每个 token 位置的期望回报,算优势函数 权重+梯度+优化器状态
Reward Model 对完整回答打分 否(冻结) 仅权重
Reference Model 计算 KL 惩罚的基准 否(冻结) 仅权重

以 7B 模型为例,四个模型全上,再算上 Adam 优化器状态(每个可训练参数额外 8 字节),不做任何优化的话显存需求轻松超过 300GB。这就是为什么 PPO 训练必须配合 ZeRO 分片、Actor/Critic 共享底座、RM/Ref 用 LoRA 挂载等手段——也是为什么业界拼命想找 PPO 的替代品。

PPO 的核心更新公式(能写出来是加分项):

L_PPO = −E[ min( ρ_t·A_t, clip(ρ_t, 1−ε, 1+ε)·A_t ) ]

其中 ρ_t = π_θ(a_t|s_t) / π_old(a_t|s_t) 是重要性采样比,A_t 是 GAE 估计的优势。clip 的作用是限制单步更新幅度,防止策略崩溃。

面试高频追问:PPO 训练 LLM 有哪些不稳定因素?

  1. reward hacking:RM 是个不完美的代理,策略会钻它的空子;
  2. KL 系数 β 难调:太大学不动,太小跑飞;
  3. Critic 预热问题:价值估计不准时优势函数噪声大;
  4. 生成长度漂移:模型学会用长度骗 RM 分数(RM 往往偏好长回答)。

四、DPO:把 RL 从 RLHF 里拿掉

2023 年的 DPO(Direct Preference Optimization)是对齐领域最重要的简化。面试官问 DPO,核心考察点只有一个:你能不能讲清它是怎么从 RLHF 目标推导出来的——因为这决定了你是"用过"还是"理解"。

推导链路(面试口头版):

第一步:RLHF 的优化目标 max E[r(x,y)] − β·KL(π‖π_ref) 存在闭式最优解

π*(y|x) = (1/Z(x)) · π_ref(y|x) · exp(r(x,y)/β)

第二步:把这个式子反解,用最优策略表示奖励函数:

r(x,y) = β·log( π*(y|x) / π_ref(y|x) ) + β·log Z(x)

第三步:把这个 r 代回 Bradley-Terry 偏好模型。关键在于配分函数 Z(x) 只和 x 有关,在 chosen 和 rejected 的分差里恰好消掉

L_DPO = −E[ log σ( β·log(π_θ(y_c|x)/π_ref(y_c|x)) − β·log(π_θ(y_r|x)/π_ref(y_r|x)) ) ]

一句话总结给面试官:"你的语言模型本身就是一个隐式奖励模型——DPO 把'训练 RM + PPO 采样优化'两步压缩成了一个直接在偏好数据上的分类损失。"

DPO 训练时只需要两个模型(π_θ 和冻结的 π_ref),不需要采样、不需要 Critic、不需要 RM,显存和工程复杂度断崖式下降。

五、PPO vs DPO vs 后续变体:怎么选

维度 PPO (RLHF) DPO 变体动向
需要的模型数 4 个(Actor/Critic/RM/Ref) 2 个(Policy/Ref) ORPO 只要 1 个
是否在线采样 是(on-policy) 否(离线偏好对) Online DPO 补采样
数据形态 prompt 即可,回答自己采 必须有成对偏好数据 KTO 只需单边标签
训练稳定性 差,超参敏感 好,接近 SFT IPO 修过拟合
效果上限 高(探索出分布外好回答) 受限于偏好数据覆盖
工程成本 极高
典型使用方 OpenAI、字节(豆包) Zephyr、Llama 3(配合迭代式采样)

答题口径:资源有限、偏好数据现成、要快速见效——选 DPO;有充足算力、追求上限、有在线数据飞轮——PPO(或 GRPO 这类简化版)仍是天花板更高的选择。Llama 3 的做法值得引用:迭代式 DPO——用当前模型采样生成新回答、RM 排序构造新偏好对、再做下一轮 DPO,兼顾了离线训练的稳定和在线采样的探索。

另外一定要知道 GRPO(DeepSeek 使用):它去掉了 Critic,用同一 prompt 采样一组回答、以组内平均奖励作为 baseline 来估计优势,显存需求大幅下降,是 2025 年以来推理模型 RL 训练的主流选择。面试提到 GRPO 并能说出"组内相对优势替代 Critic"这一点,属于明显加分。

六、手写 DPO 损失(可运行验证)

背公式不如写代码。下面用纯 PyTorch 实现 DPO 损失,并用一个玩具模型验证"训练后 chosen 的隐式奖励上升、rejected 下降":

import torch
import torch.nn.functional as F

def dpo_loss(policy_chosen_logps, policy_rejected_logps,
             ref_chosen_logps, ref_rejected_logps, beta=0.1):
    """DPO 损失。输入均为 (batch,) 的序列级 log prob 之和。"""
    chosen_rewards = beta * (policy_chosen_logps - ref_chosen_logps)
    rejected_rewards = beta * (policy_rejected_logps - ref_rejected_logps)
    logits = chosen_rewards - rejected_rewards          # 隐式奖励差
    loss = -F.logsigmoid(logits).mean()
    acc = (logits > 0).float().mean()                   # 隐式奖励准确率
    return loss, chosen_rewards.mean(), rejected_rewards.mean(), acc

def seq_logprob(logits, labels):
    """把 (B,T,V) 的 logits 和 (B,T) 的 labels 变成序列级 logp 之和。"""
    logp = torch.log_softmax(logits, dim=-1)
    token_logp = torch.gather(logp, 2, labels.unsqueeze(-1)).squeeze(-1)
    return token_logp.sum(-1)

# ---- 玩具实验:2 层 MLP 当"语言模型",词表 50,序列长 8 ----
torch.manual_seed(0)
V, T, B = 50, 8, 16
policy = torch.nn.Sequential(torch.nn.Embedding(V, 64),
                             torch.nn.Flatten(0, 1),
                             torch.nn.Linear(64, V))
ref = torch.nn.Sequential(torch.nn.Embedding(V, 64),
                          torch.nn.Flatten(0, 1),
                          torch.nn.Linear(64, V))
ref.load_state_dict(policy.state_dict())   # ref 初始化 = policy
for p in ref.parameters():
    p.requires_grad_(False)

chosen = torch.randint(0, V, (B, T))       # 模拟偏好对
rejected = torch.randint(0, V, (B, T))
opt = torch.optim.Adam(policy.parameters(), lr=1e-3)

def forward(model, ids):
    logits = model(ids).view(B, T, V)
    return seq_logprob(logits, ids)

for step in range(200):
    pc, pr = forward(policy, chosen), forward(policy, rejected)
    with torch.no_grad():
        rc, rr = forward(ref, chosen), forward(ref, rejected)
    loss, cr, rr_, acc = dpo_loss(pc, pr, rc, rr)
    opt.zero_grad(); loss.backward(); opt.step()
    if step % 50 == 0:
        print(f"step {step:3d} | loss {loss.item():.4f} | "
              f"chosen_r {cr.item():+.3f} | rejected_r {rr_.item():+.3f} | acc {acc.item():.2f}")

运行输出会看到:loss 从 0.693(=ln2,随机初始时 chosen/rejected 无差异)持续下降,chosen 的隐式奖励转正、rejected 转负,准确率升到 1.0。面试可以主动讲的细节:初始 loss 恰为 ln2 是一个 sanity check——因为 policy 与 ref 相同,隐式奖励差为 0,σ(0)=0.5。这个检查思路和上一篇 SFT 初始 loss≈ln(V) 一脉相承。

再提两个实现层面的考点:一是 log prob 必须只对回答部分求和(prompt 部分要 mask 掉),否则 prompt 的概率变化会污染梯度;二是 β 越小对 ref 的约束越弱、模型偏移越大,实践常取 0.1~0.5。

七、高频面试题清单

  1. RM 为什么用 pairwise 排序损失而不是回归打分?(人类打分一致性差、排序信号稳定)
  2. PPO 里 KL 惩罚去掉会发生什么?(reward hacking,生成 RM 高分乱码)
  3. DPO 推导中配分函数 Z(x) 是怎么消掉的?(只依赖 x,chosen/rejected 相减抵消)
  4. DPO 有什么局限?(离线数据覆盖有限、容易过拟合偏好对、chosen 概率可能同时下降——即 likelihood displacement 问题)
  5. GRPO 相比 PPO 改了什么?(去 Critic,组内平均奖励做 baseline)
  6. 你的项目里如果要做对齐,选哪个方案,为什么?(结合资源与数据实际情况作答,展示工程判断)

下一篇讲参数高效微调 PEFT:LoRA/QLoRA 的原理、显存账和注入代码。

更多推荐