跳到主要内容

[LLM 3/10] RLHF 与 PPO:用一个无法求导的奖励去训练模型

· 阅读需 10 分钟
Kobkrit Viriyayudhakorn
CEO, iApp Technology

第 2 章我们靠"逐 token 模仿标准答案"来教模型。但真正让 AI 助手可用的那些性质—— 答得对、有礼貌、不胡编、不滑回英语——既没有标准答案可模仿,也写不成一个直接的 loss function。 这一章是那个问题最正统的答案:RLHF(Reinforcement Learning from Human Feedback)配 PPO。 我们会从泰语偏好对训练一个真正的 reward model,然后从零手写约 120 行 PPO 循环, 最后做一个我在整个系列里最喜欢的实验:解开 KL 牵引绳,现场看模型作弊刷分。 这是全系列刻意安排的最重的一章,因为第 4 章(DPO)和第 5 章(GRPO) 都是从本章的公式出发,各自选择"删掉"其中一块零件。

Open in Colab03_rlhf_ppo.ipynb

1. 问题(Problem statement)

第 2 章的 SFT 藏着一个前提假设:必须有标准答案可以模仿。 但想想我们真正想要的东西,比如"把数学题做对,并用读得通的泰语解释清楚"—— 这句话没有唯一答案,好的回答可以有一百种写法,而"读得通"三个字根本写不成公式。

一旦试图直接 optimize 这些目标,我们总会撞上两堵墙:

第一堵墙——质量写不成 loss。 "更好"无法定义成函数,但人类比较的能力极强: 给两个回答让人指出更喜欢哪个,立刻能答,而且相互间还算一致。 所以现实中能收集到的数据是三样东西:prompt xx、被选中的回答 ywy_w、被拒绝的回答 yly_l

第二堵墙——就算有了分数,也无法 backprop。 假设存在一个魔法函数 r(x,y)r(x,y) 能给每个回答打分,你依然没法做 supervised 训练: 因为回答 yy 是一个 token 一个 token 采样出来的,分数在采样结束之后才到, 而导数无法逆着采样往回走——从 rr 回到权重 θ\theta 的路径恰好在那里断掉了。

想走的路撞上哪堵墙
直接给"好回答"写 loss"好"无法定义成公式,只有比较
让人打分,然后 backprop分数在 token 采样之后——梯度过不了采样这一关
让人在训练过程中实时打分人类连 rollout 的一个零头都打不过来

这就是本章标题的由来:我们要 optimize 的是一个无法求导的奖励。 能做到这件事的工具,名叫 reinforcement learning。

2. 我们要做什么(Solution)

RLHF 用两步走同时拆掉两堵墙:

  • Stage A——Reward Model: 训练一个模型 rϕ(x,y)r_\phi(x,y) 去模仿人类在偏好对上的比较(拆第一堵墙,并顶替打分打不过来的人类)
  • Stage B——PPO: 用 policy-gradient RL 把 policy 推向 rϕr_\phi 的高分方向,全程不需要对采样求导(拆第二堵墙),同时系上KL 牵引绳,不让它跑离初始模型

付出的代价是复杂度:训练期间有四个模型同时待在显存里—— policy πθ\pi_\theta(被训练的那个)、reference πref\pi_{\text{ref}}(被冻结的初始模型)、 reward model rϕr_\phi,以及一个还没登场的 value network(价值网络)VψV_\psi(3.4 节见)。

本章的核心观点

RLHF 是隔着一个不完美的代理(reward model)去 optimize 一个你无法求导的奖励。 而毫无约束地猛压一个代理指标,注定按 Goodhart's law 的方式坏掉: 当指标变成目标,它就不再是好指标。

所以公式 3.2 里那个 KL 项不是求个心安的 regularizer—— 它是你和 reward hacking(奖励欺骗)之间唯一的屏障。 第 8 节会用把它拆掉的方式,亲眼验证这句话。

还有一句话请揣在兜里带完全系列:本章的 objective 是整个系列后半程的母公式。 第 4 章(DPO)把它解成闭式,让 reward model 和 RL 循环相互抵消。 第 5 章(GRPO)换掉 advantage(优势)的估计方式,让价值网络消失。 把这一章弄懂,后面两章就成了一眼能读懂的"零件删减"。

3. 公式(Equation)

3.1 Reward model:Bradley–Terry

Stage A 用一个短短的 loss 训练 rϕr_\phi

LRM(ϕ)=E(x,yw,yl)D[logσ(rϕ(x,yw)rϕ(x,yl))]\mathcal{L}_{\text{RM}}(\phi) = -\mathbb{E}_{(x,y_w,y_l)\sim\mathcal{D}}\Big[\log\sigma\big(r_\phi(x,y_w) - r_\phi(x,y_l)\big)\Big]
  • rϕ(x,y)r_\phi(x,y) = 每段文本一个标量分数——实践中就是把语言模型的头换成单层 linear(num_labels=1
  • σ\sigma = sigmoid,把分数差变成人类会选 ywy_w 的概率(Bradley–Terry 模型)
  • 差值 rϕ(x,yw)rϕ(x,yl)r_\phi(x,y_w) - r_\phi(x,y_l) 拉得越开,loss 越低

一个常被忽视、事后吃亏的点:这个 loss 只看得见分数的差值。 把 rϕr_\phi 换成 rϕ+cr_\phi + ccc 是任意常数)——loss 纹丝不动。 也就是说 reward model 的绝对尺度没有意义,也不被训练所确定。 跑两次可能得到平均分 3.7 和 −12.4,排序却完全一致。 这就是进 PPO 之前必须先 standardize reward(减 mean 除 std)的原因——记住这一点,它会在第 7 节和第 9 节回来。

3.2 母公式:RLHF objective

如果全系列只背一个公式,就背这个:

maxθ ExD,yπθ(x)[rϕ(x,y)]    βDKL(πθ(x)πref(x))\max_\theta\ \mathbb{E}_{x\sim\mathcal{D},\,y\sim\pi_\theta(\cdot|x)}\big[r_\phi(x,y)\big] \;-\; \beta\,\mathbb{D}_{\text{KL}}\big(\pi_\theta(\cdot|x)\,\|\,\pi_{\text{ref}}(\cdot|x)\big)

用人话读一遍:"把 reward 分数拿到最多,但每离初始模型远一步,都要交罚款。"

  • πθ\pi_\theta = policy,正在训练的模型——注意 yyπθ\pi_\theta 自己采样出来的,这是它与从静态文件学习的 SFT 之间的结构性区别
  • πref\pi_{\text{ref}} = reference,初始模型(第 2 章 SFT 之后的模型),训练全程冻结
  • β\beta = 每偏离一个 nat 的价格——牵引绳的紧度
  • DKL\mathbb{D}_{\text{KL}} = policy 与 reference 之间的分布距离
为什么这是系列后半程的母公式

第 4 章(DPO)将证明这个公式存在闭式解,然后把它反过来写,让 rϕr_\phi 和 RL 循环双双消失。 第 5 章(GRPO)保留 RL 骨架,但换掉 advantage 的算法,让 VψV_\psi 消失。 两章都没有提出新的 objective——它们只是用不同的工具,解同一个公式

3.3 PPO 裁剪代理目标:Stage B 的发动机

原始的 policy gradient(REINFORCE)一批 rollout 只能更新一次就得扔,非常昂贵——因为 generate 才是瓶颈。 PPO 想把同一批 rollout 榨上好几个 epoch,就需要一个校正系数(importance sampling ratio):

ρt=πθ(atst)πθold(atst)\rho_t = \frac{\pi_\theta(a_t \mid s_t)}{\pi_{\theta_{\text{old}}}(a_t \mid s_t)}
  • sts_t = 位置 tt 的状态,即 prompt 加上已经采样出的全部 token
  • ata_t = "动作",即 rollout 时已经采样出的下一个 token
  • πθold\pi_{\theta_{\text{old}}} = rollout 那一刻的 policy snapshot——只算一次,然后冻结

再用 clip 把 ρt\rho_t 夹住:

LCLIP(θ)=Et[min(ρtA^t, clip(ρt,1ϵ,1+ϵ)A^t)]\mathcal{L}^{\text{CLIP}}(\theta) = \mathbb{E}_t\Big[\min\big(\rho_t\,\hat A_t,\ \text{clip}(\rho_t,\,1-\epsilon,\,1+\epsilon)\,\hat A_t\big)\Big]
  • A^t\hat A_t = 优势(advantage),"这个 token 比预期好多少"(下一小节定义)
  • ϵ\epsilon = trust region 的宽度(标准值 0.2)

核心在于 min + clip 合起来构成一种刻意的悲观: 如果 A^t\hat A_t 为正(好 token),把 ρt\rho_t 往上推的收益被封顶1+ϵ1+\epsilon——推过头没有任何额外收益,梯度为零。 但如果 A^t\hat A_t 为负(坏 token),min 永远会选更差的那一支——罚款没有上限。 一句话总结:收益有限,损失无限。policy 因此只会在原地附近迈小步。

别混淆:这里有两个"旧模型",而且不是同一个

公式 3.2 的 πref\pi_{\text{ref}}整个训练期间冻结,担任 KL 牵引绳。 公式 3.3 的 πθold\pi_{\theta_{\text{old}}}最近一次 rollout 的 snapshot,每轮都换,担任 trust region。 自己手写 PPO 的头号高频 bug,就是把这两个装进了同一个变量。

3.4 GAE:怎么算 advantage 才不会淹死在噪声里

advantage 由价值网络 VψV_\psi 的 TD error 构建:

δt=rt+γVψ(st+1)Vψ(st)\delta_t = r_t + \gamma V_\psi(s_{t+1}) - V_\psi(s_t) A^t=l=0(γλ)lδt+l\hat A_t = \sum_{l=0}^{\infty} (\gamma\lambda)^l\,\delta_{t+l}
  • Vψ(st)V_\psi(s_t) = 价值网络,预测"从这里走到结束,还能收多少 reward"——这就是第四个模型
  • rtr_t = 每个 token 的 reward(在我们的任务里:每个位置的 KL 罚款,外加最后一个 token 上的任务得分)
  • γ\gamma = discount factor(LLM 任务通常取 1.0)
  • λ\lambda = bias–variance 旋钮:λ=0\lambda = 0 全盘信任 VψV_\psiVψV_\psi 预测跑偏时 bias 大),λ=1\lambda = 1 完全不信、等着看真实结局(背上整条链的噪声,variance 大),常用值 0.95
把这个 V_ψ 记牢——它就是 GRPO 要干掉的那个

VψV_\psi 是一个和 policy 差不多大的模型,要用它自己的 loss 同步训练。 VψV_\psi 预测乱来,advantage 就乱来,policy 学到的就是乱来的信号——PPO 的经典崩法。 第 5 章会回答这个问题:"如果用同一个 prompt 采样出的一组回答的平均分来代替 VψV_\psi 呢?" 那就是 GRPO 的全部——用一个平均值删掉第四个模型。

3.5 PPO 的完整 loss:三项,两个模型

把所有零件拼成 optimizer 真正看到的那一个 loss(按 minimize 的写法):

LPPO=LCLIP  +  c1Et[(Vψ(st)R^t)2]    c2Et[H[πθ(st)]]\mathcal{L}_{\text{PPO}} = -\mathcal{L}^{\text{CLIP}} \;+\; c_1\,\mathbb{E}_t\Big[\big(V_\psi(s_t) - \hat R_t\big)^2\Big] \;-\; c_2\,\mathbb{E}_t\Big[\mathcal{H}\big[\pi_\theta(\cdot \mid s_t)\big]\Big]
  • 第一项 = 3.3 节的裁剪代理目标(加负号,因为我们要 maximize)
  • 第二项 = value loss,教 VψV_\psi 贴近真实 return R^t\hat R_tc1c_1 通常取 0.5
  • 第三项 = entropy bonus H\mathcal{H},防止分布过早塌缩,c2c_2 通常取 0.01
  • 至于公式 3.2 的 KL 牵引绳,实践中习惯把它塞进逐 token 的 reward:rtrtβ(logπθlogπref)r_t \leftarrow r_t - \beta\,(\log\pi_\theta - \log\pi_{\text{ref}})——第 7 节用的正是这个写法

数一数需要调的玩具:4 个模型,加上 ϵ,β,γ,λ,c1,c2\epsilon, \beta, \gamma, \lambda, c_1, c_2,再加两套 learning rate。 这就是 PPO "换个 seed 跑两遍,结果完全两回事"名声的来源, 也是第 4 章整章存在的理由。

完整内容在课程中

这篇文章大约是本章的前 30%。其余部分——环境准备、数据准备、核心代码、实测结果与总结——都在免费的 LLM Finetuning 课程中,使用 Google 登录即可阅读。

在课程中阅读完整章节 →