大模型PPO训练原理与Minimind源码实战:从公式到工程实现

发布时间:2026/9/7 15:01:41
大模型PPO训练原理与Minimind源码实战:从公式到工程实现 刚开始读大模型相关的强化学习资料时我差点被一堆符号劝退策略梯度、重要性采样、优势函数、KL散度……每个概念单拿出来都能看懂但一放到训练脚本里就完全不知道它们长什么样。后来我找到 minimind 这个项目顺着它的源码把 PPO近端策略优化的完整训练流程理了一遍才算真正把这些理论串起来。minimind 是一个用纯 PyTorch 编写的开源大模型训练示例代码量不大却完整覆盖了从数据预处理、预训练、SFT 指令微调到 DPO、PPO 对齐的大模型训练闭环。最难得的是它的 PPO 实现没有依赖 TRL 这类封装好的强化学习库所有损失计算、优势估计、KL 惩罚、策略更新逻辑都以几乎“裸代码”的形式写在脚本里特别适合用来理解大模型 RLHF 的工程实现。这篇文章不适合只想抄配置跑训练的人更适合那些已经会跑 SFT但一打开 RLHF 相关代码就头晕的开发者。我会从 PPO 的核心原理讲起再到 minimind 源码里的具体实现最后分享一些我实际调试时踩过的坑。1. 为什么用 minimind 学大模型 PPO1.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也就是“这一步生成比预期好多少”。如果优势为正策略更新时就会鼓励类似的行为如果为负就会抑制。在实践中用得最多的优势估计方法是 GAEGeneralized 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 ModelActor 和 Reference 输出各个 token 的 log-probReward 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-probref_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写成了每次更新都会重新计算的变量导致概率比几乎恒为 1PPO 也就失去了意义。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 惩罚加在哪”这些问题会有一种真正落地了的底气。