刚开始读大模型相关的强化学习资料时,我差点被一堆符号劝退:策略梯度、重要性采样、优势函数、KL散度……每个概念单拿出来都能看懂,但一放到训练脚本里,就完全不知道它们长什么样。后来我找到 minimind 这个项目,顺着它的源码把 PPO(近端策略优化)的完整训练流程理了一遍,才算真正把这些理论串起来。
minimind 是一个用纯 PyTorch 编写的开源大模型训练示例,代码量不大,却完整覆盖了从数据预处理、预训练、SFT 指令微调,到 DPO、PPO 对齐的大模型训练闭环。最难得的是,它的 PPO 实现没有依赖 TRL 这类封装好的强化学习库,所有损失计算、优势估计、KL 惩罚、策略更新逻辑都以几乎“裸代码”的形式写在脚本里,特别适合用来理解大模型 RLHF 的工程实现。
这篇文章不适合只想抄配置跑训练的人,更适合那些已经会跑 SFT,但一打开 RLHF 相关代码就头晕的开发者。我会从 PPO 的核心原理讲起,再到 minimind 源码里的具体实现,最后分享一些我实际调试时踩过的坑。
1. 为什么用 minimind 学大模型 PPO
1.1 小项目里藏着一个完整的 RLHF 闭环
minimind 这个项目最打动我的点,不是它的效果多惊艳,而是它把“训练一个小型语言模型”这件事完整走了一遍。它从开源中文语料里清洗数据,实现了简化版 LLaMA 结构,然后依次完成预训练、SFT、偏好对齐等阶段。在偏好对齐阶段,它同时提供了 DPO 和 PPO 两条路线,而 PPO 部分正是本文要重点关注的内容。
大模型的 RLHF 在工业界往往被拆成多个独立服务:策略模型一个 GPU 集群,奖励模型一个服务,参考模型又一个服务,中间还有分布式通信、日志采集、人工标注流动等复杂环节。对于想学习原理的人来说,这种工程复杂度反而会掩盖算法本身。minimind 恰恰相反,它把所有模型加载在同一个脚本里,用单卡甚至普通开发机就能跑起来,PPO 的核心逻辑就摆在眼前,没有任何一层“封装烟雾弹”。
我印象很深的一点是,它的代码风格很适合“受教育”。主训练循环里没有奇怪的抽象,就是常见的 for 循环 + batch 处理,每一个中间变量都保留着名字,比如log_probs、old_log_probs、ref_log_probs、rewards、advantages。你几乎可以照着源码,把论文里的公式逐一对应上去。
1.2 PPO 在大模型训练里的真实定位
先说清楚大模型对齐的整体流程。SFT 让模型学会按指令输出内容,能说人话了,但“说人话”和“说得好”是两码事。为了把人类偏好注入模型,通常要先训练一个奖励模型(Reward Model),它学习人类对回复质量的打分。之后用强化学习让策略模型去最大化这个奖励分数。PPO 就是这里最常用的强化学习算法,负责把奖励模型的反馈转化为模型参数的更新信号。
在语言生成任务里,我们需要重新定义强化学习的几个要素。状态(State)是已经生成的上下文,动作(Action)是下一步要生成的 token,奖励(Reward)是完整回复结束后得到的整体评分。PPO 要做的事,就是不断调整生成 token 的概率分布,让那些“更容易拿高分”的回复路径有更高的出现概率。
但这里有一个天然的工程难题:如果每更新一步策略,就要重新采样一批数据,样本效率会非常低。PPO 的解决办法是用旧策略采样一批轨迹,然后通过重要性采样去估计新策略下的期望收益,同时用力裁剪目标限制单次更新幅度。这个概念放在代码里,就是概率比ratio和torch.clamp,后面我会结合源码细讲。
2. 读代码前先读懂 PPO 的核心逻辑
2.1 目标函数里那个 min 和 clip 到底在防什么
PPO 论文里的核心目标函数,写出来是:
L(θ) = E[ min( r_t(θ) * A_t, clip(r_t(θ), 1-ε, 1+ε) * A_t ) ]
其中 r_t(θ) = π_θ(a_t|s_t) / π_old(a_t|s_t),也就是新策略和旧策略在某个动作上的概率比。A_t 是优势函数,衡量这个动作比平均水平好多少。
“概率比”这个概念很多人容易绕晕。简单说,我们用旧策略采了一批样本,现在新策略参数变了,同一个动作在新策略下的概率可能变了。这个比值如果大于 1,说明新策略觉得这个动作比以前更重要;如果小于 1,说明新策略对这个动作的兴趣下降了。
为什么要加裁剪?因为如果新策略对某个动作的概率比涨到 3 倍,而那个动作的优势恰好为正,那么梯度会剧烈放大这个动作的概率,导致一步更新过大,策略直接失控。裁剪的作用就是限制这个比值只能落在 [1-ε, 1+ε] 范围内,超出部分不再继续提供梯度收益。ε 通常取 0.2,也就是单次更新最多让某个动作概率比变化约 20%。
在 minimind 的代码里,这一逻辑最终浓缩为surr1 = ratio * advantages、surr2 = torch.clamp(ratio, 1 - clip_eps, 1 + clip_eps) * advantages、policy_loss = -torch.min(surr1, surr2).mean()。这可能是全篇源码里最值得抄到笔记本上的三行代码。
2.2 为什么大模型 PPO 少不了一个参考模型
如果只优化奖励模型给出的分数,策略模型会抓住一切漏洞猛冲。比如奖励模型偏好“回答中包含更多礼貌用语”,模型就可能在回复里反复堆叠“谢谢您的提问,这是一个很好的问题”。从奖励分数看,训练似乎成功了,但实际生成质量一塌糊涂。
为了防止模型在优化奖励时偏离原来的语言能力,大模型 PPO 会额外引入一个参考模型(Reference Model)。参考模型通常就是 SFT 阶段得到的模型,在 PPO 训练过程中参数完全冻结,不参与梯度更新。它的作用只有一个:提供一个“不偏离太远”的锚点。
具体做法是,对同一批生成回复,分别计算当前 Actor 模型和参考模型的 log-prob,然后两者相减。这个差值就是对每个 token 级别 KL 散度的一种近似。最终奖励变成:
final_reward = reward_model_score - kl_coef * (log_probs - ref_log_probs)
这样一来,模型虽然可以去争取更高奖励,但每偏离参考模型一步,都要付出代价。偏离越大,代价越大。kl_coef 就是这一约束的强度系数,调大模型更保守,调小模型更容易放飞自我。
2.3 优势函数:GAE 在语言生成中的含义
PPO 里除了 Actor 和 Reference,还需要一个 Critic(价值)网络,用来估计状态的价值。它的输出用于计算优势函数 A_t,也就是“这一步生成比预期好多少”。如果优势为正,策略更新时就会鼓励类似的行为;如果为负,就会抑制。
在实践中用得最多的优势估计方法是 GAE(Generalized Advantage Estimation)。它不只看单步奖励,还会通过一个递归公式把未来多步的信息逐步累计进来,好处是方差更低、训练更稳。代码实现其实非常短,通常就是一层循环:
advantages = torch.zeros_like(rewards) last_gae = 0 for t in reversed(range(seq_len)): next_value = values[t + 1] if t + 1 < seq_len else 0 delta = rewards[t] + gamma * next_value - values[t] advantages[t] = last_gae = delta + gamma * lam * last_gae在语言生成场景里,gamma 经常设置为 1.0 或接近 1.0。因为一条回复的长度是有限的,不存在真正意义上的“长期折扣”,每一 token 的重要性并没有随时间推移而衰减。lam 则用来控制 GAE 对多步信息的依赖程度,常见取值是 0.95。
3. 模块定位:如何高效翻阅 minimind 源码
3.1 建议的阅读顺序与文件定位
很多人拿到源码就直接打开核心模型文件,从 attention 开始读,结果还没看到 PPO 相关内容,人已经累了。我的建议是先看 README,再打开配置文件,最后再进训练脚本。
minimind 的仓库结构不算复杂,预训练、SFT、RL 相关脚本通常分开存放。PPO 的训练入口一般在train_rl.py这类命名里。打开之后,先别盯着代码逐行看,先搜索以下几类关键字:
clip_eps或clip_epsilon,定位裁剪相关的超参数;kl_coef,定位 KL 惩罚强度;advantages,定位优势计算;policy_loss,定位策略更新。
这四个关键字找齐之后,整个 PPO 训练脚本的大致骨架就出来了。阅读顺序应该是:配置文件里的超参数 -> 主训练循环 -> 数据采样和奖励计算 -> 损失计算。模型定义放到最后再看,因为它只是工具,PPO 真正的灵魂在数据流和损失函数里。
3.2 用 no_grad 快速画出模型分工图
在 minimind 里,同时存在多个“大模型”,刚接触时最容易搞混的是:哪个是 Actor?哪个是 Reference?哪个是 Reward Model?它们之间什么关系?
有一个非常实用的技巧:在源码里搜索torch.no_grad或者requires_grad_(False)。凡是包裹在这些语句里运行的模型,基本都是不参与梯度更新的辅助模型。这样能帮你快速画出模型分工图:
| 模型 | 是否更新参数 | 主要作用 |
|---|---|---|
| Actor | 是 | 生成回复,是被训练的“主模型” |
| Reference | 否 | 提供旧策略或 SFT 基准,用于计算 KL 惩罚 |
| Reward Model | 否 | 给生成回复打分,输出一个标量或序列化奖励 |
| Critic | 是 | 估计状态价值,用于计算优势函数 |
这张图画清楚之后,后面读代码时会非常省力。比如看到with torch.no_grad(): ref_log_probs = ref_model(...)就能立刻反应出这行代码在计算什么。
3.3 一条生成样本变成训练数据的完整链路
把 PPO 训练的一个 step 拆成五步,会清晰很多。
第一步,从数据集里采样一批 prompt。这些 prompt 可以是问题、指令,或者对话的上半部分。第二步,用当前 Actor 模型以一定的 temperature 和 top_p 生成回复,生成过程不需要梯度,只做推理。第三步,将完整序列分别送入 Actor、Reference 和 Reward Model:Actor 和 Reference 输出各个 token 的 log-prob,Reward Model 输出这条回复的奖励分数。第四步,计算最终奖励,也就是在奖励分数基础上减去 KL 惩罚项。第五步,将一批样本累积起来,通过优势函数计算 advantage,再用 PPO 裁剪损失更新 Actor,用价值损失更新 Critic。
在实际源码里,这五步的顺序不一定完全按我列出的来,有些实现会把 rollout 和 update 分成两个循环。但只要抓住这个链路,你就能从大段代码里准确识别出“现在在哪一步”。
4. 关键代码段精读:从 log-prob 到策略更新
4.1 先算对 log-prob,后面才不会白忙
PPO 里的很多计算都依赖 log-prob,它表示模型给某个真实生成的 token 分配的概率取对数。计算方式非常直接,只需要把模型输出的 logits 做 log_softmax,再用 gather 取出真实 token 位置对应的值。
logits = model(input_ids)[0] # shape: [batch, seq_len, vocab_size] log_probs = logits.log_softmax(dim=-1) token_log_probs = torch.gather(log_probs, -1, labels.unsqueeze(-1)).squeeze(-1)有一个细节需要特别注意:生成回复时,不同样本的回复长度可能不同,所以 padding 位置一定要在后续计算中 mask 掉。否则,padding 部分也会参与 log-prob 平均,导致数值偏差。mini思维里通常会有对应的 mask 逻辑,阅读时留意一下即可。
4.2 奖励构建:KL 惩罚是怎么叠加上去的
奖励计算是 PPO 实现中容易被忽略的一步,但它的设计直接决定训练稳定性。在 minimind 里,最终权重大概率是这样加出来的:
reward = reward_model_score - kl_coef * (log_probs - ref_log_probs)这段代码里,log_probs是当前 Actor 在一条回复上的 log-prob,ref_log_probs是冻结的参考模型在同样回复上的 log-prob。两者的差就是“策略偏离参考模型的程度”,通常按 token 维度计算,再在序列长度上平均或者累计。
值得注意的是,这里的 log-prob 必须对应同一批生成回复、同一系列 token 位置,否则计算出来的 KL 完全失真。我建议阅读时顺手验证一下维度相对齐。
如果你发现 reward model 的分数本身波动很大,可以在计算最终 reward 之前,先对同一批次的 reward 做标准化:
reward_score = (reward_score - reward_score.mean()) / (reward_score.std() + 1e-8)这种处理能显著提升训练稳定性,即使原始 reward 的绝对值范围很怪,标准化后模型面对的每个 batch 奖励分布都相对一致。
4.3 GAE 优势估计的极简实现
前文给出过 GAE 的标准循环实现。这里我想补充一个工程细节:values 来自 Critic,但在代码里必须有明确的detach(),否则价值网络自身的梯度会串到策略更新里。
语言模型场景里,Critic 的输入通常是同一个序列的 hidden state,或者干脆用 Actor 的 logits 作为输入特征。minimind 里可能没有过度复杂的价值网络结构,更可能是一个线性层或者小型 MLP 头部。阅读时不用纠结它的结构,只需要清楚:values 的 shape 应该与 rewards 一致。
用循环实现 GAE 在序列很长时效率一般,但可读性是最好的。如果看到向量化实现,也不用慌,它的原理完全一样。
4.4 裁剪损失、Value 损失与参数更新
策略更新的核心代码就是之前反复提到的那几行:
ratio = torch.exp(log_probs - old_log_probs) surr1 = ratio * advantages surr2 = torch.clamp(ratio, 1 - clip_eps, 1 + clip_eps) * advantages policy_loss = -torch.min(surr1, surr2).mean()需要特别留意的是old_log_probs和log_probs的区别。old_log_probs是在 rollout 阶段缓存下来的旧策略 log-prob,在后续多次更新中保持固定;log_probs则是每次 inner update 重新用当前策略计算的最新 log-prob。两者差距越大,ratio 越偏离 1,裁剪就越容易生效。
价值损失通常用均方误差:
value_loss = F.mse_loss(values, returns)这里的returns = advantages + values,也被称为折扣回报。总损失一般是策略损失、价值损失和可选熵正则项的组合。熵正则项的作用是保留一定的随机性,防止策略过早坍缩,但大模型 PPO 里是否使用,取决于具体实现。阅读时看一下 loss 最终组合就能知道作者的取舍。
5. 实操中的常见坑和调试经验
5.1 训练不稳定:先查 KL 再查学习率
我遇到过的最典型现象是:训练刚开始几步奖励确实上升了,但某个 step 之后 KL 突然飙到几十,生成质量急剧下降。排查原因时,往往发现 KL 系数设置过小,或者学习率过大。
建议的做法是:训练时同时记录奖励和 KL 散度两条曲线。如果 KL 上升速度过快,先尝试把学习率降到原来的 0.1 倍观察;如果还压不住,再调大 kl_coef。先定学习率,再定 KL 系数,逐个变量去试,不要同时改一堆参数。
一个小技巧:如果 KL 始终无法控制,检查一下old_log_probs是否在更新循环中被意外覆盖。我确实踩过这种低级错误——把old_log_probs写成了每次更新都会重新计算的变量,导致概率比几乎恒为 1,PPO 也就失去了意义。
5.2 奖励模型被攻击:样本评审比曲线更重要
奖励模型并不是绝对可靠的。PPO 的优化能力很强,模型很快会找到奖励模型打分逻辑里的“捷径”。典型的例子是:模型学会输出格式非常结构化的废话,比如堆砌大量小标题、重复使用特定句式,奖励分数一路走高,但实际读起来空洞无物。
我在调试时习惯每隔若干 step 就保存一批生成样本,肉眼检查输出质量。不要只看 loss 曲线或者奖励曲线,唯一能确认训练方向正确的,是看实际生成的文本是否符合人类直觉。如果发现模型进入了“高分废稿”状态,多半需要回退到更早的 checkpoint,再重新调整奖励构造或 KL 系数。
5.3 显存不足时的几个实用补救措施
minimind 的模型规模已经很小,但 RL 训练需要同时加载多个模型,显存压力还是比 SFT 大不少。如果直接跑崩,按优先级尝试下面几种方法:
- 让 Actor、Reference、Reward Model 共享同一份基座权重,只保留不同的 head 或 LoRA 参数,能省下大量显存。
- 生成阶段和奖励计算阶段全部用
torch.no_grad(),避免中间激活值占用缓存。 - 把 Critic 网络做得尽量小,不一定要和 Actor 同一规模。
- 用梯度累积模拟更大 batch,而不是直接增大 batch size。
这些方法不改变算法本质,只是工程层面的取舍。对于学习目的来说,跑通一个小规模的训练循环比追求完美效果更重要。
5.4 给奖励加个标准化,稳定性会好很多
很多人复现 PPO 时,发现策略更新一步之后 loss 波动特别大,大概率是 reward 分布太不稳定。reward model 的输出可能在 -3 到 +8 之间乱跳,不同 batch 的分布也完全不同。这种情况下,直接在原始 reward 上做 KL 惩罚,梯度方向会很混乱。
我尝试过一种非常有效的做法:在构造最终 reward 前,对同一个 batch 内的 reward_model_score 减均值除标准差,再做 KL 惩罚叠加,效果会稳定很多。这也是目前很多开源 RL 框架里默认的实现方式。不要担心“不够原始”,工程上稳定优先。
写在后面
我在跑 minimind 的 PPO 之前,一直觉得自己理解 PPO 公式,真正跑完一遍才知道,理论到实现之间隔着一整条数据流。比如old_log_probs必须从 rollout 阶段缓存下来,比如 KL 惩罚需要挂在奖励上而不是单独做 loss,这些细节在论文里完全不写,但在代码里少一个都不行。
如果你正在准备入门大模型 RLHF,我不建议一上来就啃大型分布式框架。先把 minimind 里这个最小闭环读懂、跑通,再去看更深层的封装实现,会顺畅很多。读完这套代码之后,你对“谁是 Actor、谁是 Critic、KL 惩罚加在哪”这些问题,会有一种真正落地了的底气。