外观
RLHF训练流程与对齐替代方案详解
内容整理自学习笔记,仅供面试备考参考;不构成录用、培训或考试承诺。
1. LLM 经典预训练 Pipeline
目前基于 Transformer decoder 的 LLM(如 ChatGPT、LLaMA、Baichuan 等),通常有基于预训练的 base 模型和使用 RLHF 微调的 Chat 模型。Chat 模型的训练一般包括三个步骤:
| 步骤 | 名称 | 说明 |
|---|---|---|
| 1 | 预训练(Pre-training) | 从大量无标注文本数据中学习通用知识 |
| 2 | 有监督微调(SFT) | 优化模型以更好地遵守特定指令 |
| 3 | 对齐(Alignment) | 使 LLM 更有用且更安全地响应用户提示 |
2. 预训练(Pre-training)
利用数十亿到数万亿个 token 的庞大文本语料库对模型进行预训练,使模型能够根据提供的文本预测**「下一个单词」**。
3. 有监督微调(Supervised Fine-tuning)
3.1 什么是 SFT?
虽然 SFT 训练目标与预训练类似(也是预测下一个单词),但需要人工标注的指令数据集,其中模型输入是指令(根据任务不同也可能包含输入文本),输出为预期回复内容。
3.2 SFT 的训练数据格式
Instruction: "Write a limerick about a pelican."
指令:"写一首关于鹈鹕的打油诗。"
Output: "There once was a pelican so fine..."
输出:"从前有一只鹈鹕很好..."模型会把 "Write a limerick about a pelican" 作为输入,逐个 token 进行预测,输出 "There once was a pelican so fine..."
3.3 预训练 vs 有监督微调
| 对比维度 | 预训练(Pre-training) | 有监督微调(SFT) |
|---|---|---|
| 训练目标 | 预测下一个单词 | 预测下一个单词(相同) |
| 训练数据量 | 庞大(数十亿~数万亿 token) | 较小 |
| 数据格式 | 无标注文本 | 人工标注的指令数据 |
| 是否需人工标注 | 否 | 是 |
4. 对齐(Alignment)
通过微调的方式,将语言模型与人类的偏好、价值观进行对齐,这也是 RLHF 机制发挥的地方。
5. RLHF 流程详解
5.1 RLHF 三步流程
- 在预训练好的模型上进行有监督微调(SFT)
- 在有监督微调模型基础上创建一个 Reward Model(RM)
- 基于 RM 模型使用 PPO 算法微调 SFT 模型
5.2 如何进行有监督微调?
先收集一个 Prompts 集合,要求标注人员写出高质量回复,然后使用该数据集以监督的方式微调预训练的基础模型。
5.3 如何创建 RM 模型?
- 对于每个 Prompt,要求 SFT 后的 LLM 生成 4~9 个回复
- 由标注人员根据个人偏好对所有回复进行排序
- RM 来自 RLHF 第一步的 SFT 模型,SFT 的输出通过一个回归层(单个输出节点)转换为奖励分数
虽然排序过程很耗时,但工作量比第一步的有监督数据集构建要少一些。
5.4 如何用 PPO 微调 SFT 模型?
基于 RM 模型使用 Proximal Policy Optimization(PPO)算法微调 SFT 模型。
5.5 InstructGPT 的原理
InstructGPT 是基于强化学习的文本生成模型,核心原理涉及两个概念:
| 概念 | 说明 |
|---|---|
| RLHF | 先用人类生成的示例预训练模型,再通过与人类评估者交互收集评估结果,创建强化学习数据集 |
| Reward Shaping | 将人类评估者的反馈与模型生成的文本进行比较,计算差异度量作为奖励信号的一部分,引导模型训练 |
通过两者的结合,InstructGPT 能够通过人类评估者的反馈指导生成过程,逐步提升生成文本的质量和一致性。
6. LLaMA 2 的 RLHF
6.1 LLaMA 2 RLHF 概述
Llama-2-chat 在第一步 RLHF 微调上使用相同的指令数据,但在第二步使用了两个奖励模型:
- 侧重有用性(helpfulness)
- 侧重安全性(safety)
通过多个阶段不断进化,奖励模型也会根据 Llama-2-chat 模型出现的错误进行更新,并增加了拒绝采样(rejection sampling)步骤。
6.2 Margin Loss 实现逻辑
例如四个回复排序结果为 A < C < D < B,可得六个对比结果:A<C, A<D, A<B, C<D, C<B, D<B
| 对比项 | InstructGPT | LLaMA 2 |
|---|---|---|
| 回复对比数量 | 同一提示下 4~9 个输出排序 | 每次只看 2 个回复对比 |
| 对比标签 | 仅优劣 | 新增边际标签:「显著更好」和「好的不明显」 |
Margin Loss 中,m(r) 可调节两个回复之间的差值,如果对比结果为「显著更好」,则会增加梯度值,加快更新速度。
6.3 两个 RM 模型的实现逻辑
用于模型优化的最终奖励函数会将两个分数进行线性组合:
- 有用性 RM:评估回复对用户的有用程度
- 安全性 RM:评估回复的安全程度
6.4 拒绝采样逻辑
| 对比项 | 拒绝采样 | PPO |
|---|---|---|
| 更新方式 | 得到 K 个输出,使用最高奖励的输出更新梯度 | 每次基于单样本更新 |
| 训练阶段 | SFT 初始阶段后仅用拒绝采样,再结合 PPO | 与拒绝采样结合使用 |
Llama 2 使用训练流水线,迭代产生多个 RLHF 模型(从 RLHF-V1 到 RLHF-V5)。
7. RLHF 替代方案
7.1 为什么需要替代方案?
RLHF 在 InstructGPT 和 LLaMA 2 论文中被证明有效,但过程比较复杂。
7.2 替代方案总览
| 替代方案 | 论文 | 核心思路 |
|---|---|---|
| Constitutional AI | arxiv:2212.08073 | 基于人类提供规则列表的自我训练机制 |
| HIR | arxiv:2302.05206 | 基于重新标记的监督微调,将失败案例转化为有用训练数据 |
| DPO | arxiv:2305.18290 | 直接用交叉熵损失微调 LLM,无需单独 RM |
| ReST | arxiv:2308.08998 | 离线生成训练数据集,在质量越来越高的子集上迭代训练 |
| RLAIF | arxiv:2309.00267 | 用 LLM 生成奖励模型训练的评级,替代人类标注 |
Constitutional AI 中的"红队"(Red Team)指测试目标系统防御能力的过程,通过模拟现实世界攻击者的战术来挑战并改进系统。
8. RLHF 训练中如何选取最优 Checkpoint?
8.1 核心问题
RLHF 训练过程中,Reward Model 输出的只是近似奖励(Proxy Reward),不能完全相信训练过程中的 Reward 变化——"更高"的 Reward 不一定意味着"更好"的效果。
8.2 真实 Reward 与近似 Reward 的关系
| 指标 | 随 KL 增大的趋势 |
|---|---|
| 真实 Reward(实线) | 先逐步提升,到达峰值后逐渐减小 |
| 近似 Reward(虚线,RM 打分) | 一直稳步上升 |
最优模型出现在真实 Reward 曲线的最高点,但我们无法直接获得真实 Reward。
8.3 用 KL 估算真实 Reward
关键假设:真实 Reward 曲线与「当前模型和初始模型」之间的 KL 存在某种关系。
KL 可以被实时计算,OpenAI 找到了计算公式,公式中 3 个关键参数:
| 参数 | 说明 |
|---|---|
| α | 与 RM 大小和训练数据规模有关 |
| β | 与 RM 大小和训练数据规模有关 |
| d | 初始模型和当前模型的 KL 开根号 |
8.4 Scaling Law 实验结论
| 实验维度 | 结论 |
|---|---|
| RM 大小 | 相同训练数据下,RM 越大 actor 模型能获得更高的真实 reward |
| RM 大小 | RM 越大,模型能在更大的 KL 处才发生下降转折 |
| RM 数据集规模 | 数据集越大提升越大,但最少需超过 2000 条 |
| Policy Model 大小 | 模型越大,利用 RM 做提升的收益越小(初始分已较高) |
| Policy Model 大小 | 无论模型规模如何,最优 Reward 对应的 KL 值相同 |
论文中 2k 下限是在 3M~3B 大小模型下得出的结论,更大模型是否适用尚不确定。