RAFT:抛开 PPO 的四模型同驻,best-of-K 采样加微调就赢下对齐

RAFT: Reward rAnked FineTuning for Generative Foundation Model Alignment

Hanze Dong, Wei Xiong, Deepanshu Goyal, Yihan Zhang, Winnie Chow, Rui Pan, Shizhe Diao, Jipeng Zhang, Kashun Shum, Tong Zhang

cs.LG, cs.AI, cs.CL, cs.CV, stat.ML

2023-04-14

对齐不走 PPO:每条 prompt 采 K 个回答、挑奖励最高的做监督微调,循环到收敛。LLaMA-7B 上奖励 2.294 超过 PPO 的 2.077,训练只装 1 个模型。

这篇在解决什么

用人类反馈做强化学习(RLHF)来对齐大模型,工业上几乎成了标配,但它背的包袱很重。最常用的 PPO(Proximal Policy Optimization)是 trial-and-error 式的学习,训练不稳、效率低。更扎眼的是显存:PPO 要同时把四个模型塞进显存,正在训练的策略模型、用作参考的旧模型、奖励模型、价值评估的 critic 模型各占一份,一张卡上能塞下的参数因此被砍掉一大半。再加上奖励模型不完美,容易被策略钻空子刷高分(reward hacking),只盯着奖励优化还会让生成质量下降(alignment tax)。

作者想绕开这些。核心直觉是:best-of-K,也就是每条 prompt 采 K 个回答、挑奖励最高的那个,本身就能逼近 RLHF 的效果,但只花推理成本、不训练。能不能把这个挑出来的好样本反过来当监督信号,让模型迭代地学好?

方法

RAFT(Reward rAnked FineTuning)就三步,循环到奖励收敛:

更新后的模型再回到第一步。关键设计是把生成和训练解耦:排序时才用奖励模型,微调时只装策略模型自己一个,峰值显存里不再同时挤着四个模型。这也是它能直接套到扩散模型上的原因,只要你有能打分的奖励函数,这套循环就成立。

几个旋钮:K 越大越偏向高奖励(代价是采样量);温度 λ 越高生成越多样,太高会出乱码,把可用范围压窄;KL 惩罚系数 β 可选,用来拉住模型别离初始分布太远、防止过优化。

作者给了 best-of-K 的理论上界:E[max r] ≤ E[r] + B√(log K / 2),意思是 K 带来的边际收益递减。这解释了为什么迭代式地「采-选-学」比一次性堆大 K 更划算。

结果

主实验在 LLaMA-7B 上,用 Anthropic 的 HH-RLHF(人类偏好对话数据):

模型奖励困惑度平均长度
LLaMA-7B(base)-0.4354.781119.9
LLaMA-7B-SFT0.7723.781145.4
LLaMA-7B + PPO2.0774.156127.8
RAFT-K32-λ1.02.2944.031156.2

RAFT 拿到最高奖励 2.294,超过 PPO 的 2.077,困惑度也保持在合理区间。在 GPT-4 和人工两两对比里,RAFT-K32 对 PPO 多数打赢:对 PPO-β0.1,GPT-4 判 65 胜 32 负,人工判 66 胜 14 负。

K 的消融:K 从 8 涨到 32,奖励从 2.180 升到 2.329,印证「K 越大奖励越高」;K=32 大约 10-12 轮收敛,K=8 要 15-18 轮。计算上,K∈{8,16,32} 分别花 5、6.05、7.05 小时,PPO 最快配置约 8.7 小时,RAFT 更快。

扩散模型上的差距更悬殊。在 Stable Diffusion v1.5、256×256 分辨率:

指标预训练DDPORAFT
域内美学分4.636.046.14
域外美学分4.645.766.07
训练时间(单卡 A40)N/A415 分钟8.4 分钟

RAFT 在美学分上和 DDPO 持平甚至更好,训练快了约 50 倍。它还能做蒸馏:用 LLaMA-7B 当老师采样,把 GPT-Neo-2.7B 学上去,奖励从 -1.23 拉到 0.739。

为什么重要

它给「对齐」换了个更省的工程实现:不碰 PPO 那套 critic、reference、reward 同驻的复杂度,用「采样-挑选-监督微调」的循环就追平甚至超过 RLHF。对显存吃紧、想快速迭代的团队,这是条直接能落地的路;对扩散模型这类没有现成 RL 训练栈的生成模型,RAFT 几乎即插即用。

放到今天的语境里看,这篇 2023 年的工作被近期 RLVR(可验证奖励的强化学习)论文重新引用不是偶然。「训练时做 best-of-K 采样」这个想法,正是 GRPO 这类方法去掉 critic、靠采样估计基线的思想前身。读 RAFT 有助于看清当前简化 RL 训练的来龙去脉。

局限与存疑

术语

原文与代码

社区讨论

相关论文

全部论文解读