
1. 从一次训练崩溃说起为什么RL Infra值得单独聊第一次跑RLHF的PPO训练我盯着终端里跳动的loss曲线以为一切顺利。结果到第300步左右KL散度突然炸了reward曲线像断了线的风筝一样往下掉GPU显存直接OOM。那会儿我才意识到强化学习这套东西跟监督学习完全是两个物种——它不是把数据喂进去等收敛就完事了而是一个采样、评估、更新、再采样的闭环系统任何一个环节的时序错位或者资源没对齐整个训练就会崩。这就是我想写这个系列的原因。标题里的RL Infra不是指某个具体框架而是指支撑强化学习训练跑起来的那套工程底座rollout怎么调度、actor和critic怎么共享显存、reference model什么时候前向、advantage怎么算、KL惩罚加在哪一步。这些东西在论文里往往一笔带过但真正落地的时候它们决定了你是跑通还是跑崩。这篇是系列第一篇聚焦最基础也最容易被忽视的部分RL训练的loop结构以及RLHF场景下PPO的完整数据流。我会把loop拆成几个阶段讲清楚每个阶段在干什么、为什么这么设计、代码里对应哪几行。适合两类人看一是刚接触RLHF、想知道PPO到底怎么跑起来的同学二是已经在跑训练、但被各种时序bug和显存问题折磨的工程师。读完你应该能自己搭一个最小可用的PPO训练循环并且知道每个环节的坑在哪。2. RL训练loop的本质一个带反馈的采样-更新闭环2.1 为什么RL不能像监督学习那样喂数据等收敛监督学习的范式很直接你有一批固定的数据模型前向算loss反向传播更新参数重复若干epoch。数据是静态的梯度是无偏的收敛性有理论保证。但强化学习不一样它的数据是策略自己生成的——你当前的策略决定了下一次采样看到什么状态、采取什么动作、拿到什么奖励。策略一变数据分布就变这就是所谓的非平稳分布问题。这个差异带来的直接后果是你不能把RL当成一个普通的训练任务来调度。监督学习里dataloader可以预取、可以shuffle、可以多epoch复用RL里rollout必须用当前策略跑跑完的数据用完就得扔或者进replay buffer但要注意off-policy的程度。更麻烦的是RLHF里还有四个模型同时在显存里actor、critic、reward model、reference model。它们的前向时机、是否需要梯度、显存占用都不一样调度错了就是OOM或者数值不稳定。我见过不少团队一开始想用纯监督学习的思路套RL结果卡在为什么loss不降或者为什么训练到一半崩了上。根子就在于没理解loop的时序约束。2.2 loop的四个阶段采样、评估、优势估计、策略更新一个标准的on-policy RL loop拆开来看是四个阶段顺序不能乱第一阶段是采样Rollout。用当前策略π_θ在环境里跑若干步收集轨迹trajectory。在RLHF里环境就是prompt数据集加上reward model策略是语言模型动作是生成的token。这一步只做前向不需要梯度所以可以用torch.no_grad()包起来显存占用相对小。第二阶段是评估Evaluation。对采样出来的每个token算出它的log概率logprob、reward、以及reference model下的logprob。这里有个关键点reward通常只在序列末尾给一个标量比如reward model对完整回复打分但PPO需要的是每个token的reward信号所以要做reward的广播和折扣累积。同时reference model的logprob用来算KL惩罚防止策略跑偏太远。第三阶段是优势估计Advantage Estimation。用critic网络或者GAE算出每个token的advantage。advantage的含义是这个动作比平均水平好多少它是策略梯度的核心信号。GAEGeneralized Advantage Estimation通过引入λ参数在偏差和方差之间做权衡这个后面细讲。第四阶段是策略更新Policy Update。用advantage构造PPO的clipped objective反向传播更新actor和critic。这一步需要梯度显存占用最大。更新完之后回到第一阶段用新策略重新采样。这四个阶段构成一个完整的iteration。注意采样和更新是严格交替的不能像监督学习那样先采一大批数据然后更新很多轮。PPO虽然允许在同一个batch上做多个epoch的更新这是它的proximal特性但epoch数不能太多否则策略偏离采样时的分布太远importance sampling的ratio会爆炸。2.3 on-policy与off-policy的取舍PPO为什么卡在中间严格来说PPO是近似on-policy的算法。它用旧策略采的数据来更新新策略但通过clip机制限制新旧策略的差异。这个设计很巧妙纯on-policy比如A2C每次更新后数据就作废样本效率极低纯off-policy比如DQN可以用replay buffer但训练不稳定尤其在连续动作空间和高维输出比如语言模型上很难调。PPO的clip操作本质上是在说我允许你用旧数据更新几次但每次更新后新旧策略的ratio不能超过1±ε。这个ε通常取0.1或0.2。超过这个范围梯度就被截断了相当于告诉模型这一步走太远了别走了。在RLHF里这个特性特别重要。因为语言模型的输出空间是离散的、巨大的词表几万维策略的微小变化可能导致生成结果的巨大差异。如果没有clip一次更新就可能让模型输出完全崩坏。我实测下来ε0.2是个比较稳的起点太小了学不动太大了容易崩。3. RLHF-PPO的数据流拆解四个模型怎么协同3.1 四个模型的角色与显存占用分析RLHF-PPO里同时存在四个模型各自的职责和资源需求完全不同模型职责是否需要梯度显存占用7B模型fp16前向时机Actor生成回复被优化是参数14G 梯度14G 优化器状态56G ≈ 84G采样更新Critic估计状态价值是同actor ≈ 84G评估更新Reward Model给完整回复打分否参数14G评估阶段Reference Model算KL基准否参数14G评估阶段四个模型加起来7B规模下光参数和优化器状态就要接近200G显存。这就是为什么RLHF训练通常需要多卡甚至多机——单卡80G根本放不下。实际工程里常见的做法是actor和critic用FSDP或ZeRO-3分片reward和reference用fp16推理模式加载甚至可以把reference model的logprob预先算好缓存起来。这里有个容易踩的坑reward model和reference model的前向必须在torch.no_grad()下做否则PyTorch会默认构建计算图显存直接翻倍。我见过有人忘了加这个结果训练到一半OOM排查了半天才发现是reference model在偷偷存梯度。3.2 采样阶段prompt怎么变成训练数据采样阶段的输入是一批prompt输出是完整的轨迹数据。具体流程是这样的从prompt数据集里取一个batch比如batch_size64。Actor模型对每个prompt做自回归生成直到遇到EOS或达到max_length。记录每个token的logprob用log_softmax算、生成的token id、以及attention mask。把生成的完整序列拼上原始prompt送给reward model打分。同时送给reference model算每个token的logprob。这里的关键细节是logprob的计算方式。生成的时候模型输出的是logits你需要对它做log_softmax再gather对应token的值。注意要在正确的维度上操作语言模型的logits形状是[batch, seq_len, vocab_size]gather的时候要用生成的token id在最后一维索引。import torch import torch.nn.functional as F def compute_logprobs(logits, labels): # logits: [batch, seq_len, vocab_size] # labels: [batch, seq_len] log_probs F.log_softmax(logits, dim-1) # gather对应token的logprob selected log_probs.gather(dim-1, indexlabels.unsqueeze(-1)) return selected.squeeze(-1) # [batch, seq_len]这段代码看起来简单但有个隐藏的坑padding token的logprob要mask掉。如果不mask这些无效位置的logprob会进入loss计算导致梯度信号被稀释。正确做法是用attention mask或者response mask把padding位置置零。3.3 评估阶段reward、KL、value三路信号的汇合评估阶段是整个loop里信息密度最高的地方三路信号在这里汇合第一路是reward信号。Reward model对完整回复打一个标量分数比如2.3。但这个分数是给整个序列的PPO需要的是每个token的reward。常见的做法是只在最后一个token位置放这个分数前面的token reward为0。然后做折扣累积discount factor γ通常取1.0因为语言生成没有明显的时间折扣。第二路是KL惩罚。KL散度衡量当前策略和reference策略的差异公式是KL(π_θ || π_ref) log(π_θ) - log(π_ref)。这个值逐token计算然后加到reward上作为惩罚项。KL系数通常叫β或kl_coef控制惩罚强度典型值在0.01到0.1之间。β太小模型会跑偏到reward hackingβ太大模型学不动输出和reference几乎一样。第三路是value估计。Critic模型对每个token位置输出一个value表示从这个位置开始预期能拿到多少累积reward。这个value用来算advantage。三路信号汇合后每个token位置就有了reward_with_kl、value、old_logprob。这些就是下一步算advantage的输入。注意KL惩罚的符号很容易搞反。正确的做法是reward - β * KL因为我们要惩罚策略偏离reference。如果写成reward β * KL模型会主动去最大化KL直接跑飞。3.4 优势估计GAE的λ参数到底在调什么GAE的核心思想是用TD残差的指数加权平均来估计advantage。公式是这样的δ_t r_t γ * V(s_{t1}) - V(s_t) A_t δ_t (γλ) * δ_{t1} (γλ)^2 * δ_{t2} ...其中λ控制偏差和方差的权衡。λ0时A_t δ_t这是纯TD估计偏差大但方差小λ1时A_t等于蒙特卡洛回报减去baseline方差大但偏差小。实践中λ通常取0.95这是个经验值在语言模型任务上表现比较稳。实现上GAE可以从后往前递推计算def compute_gae(rewards, values, gamma1.0, lam0.95): # rewards: [batch, seq_len] # values: [batch, seq_len] advantages torch.zeros_like(rewards) last_gae 0 for t in reversed(range(rewards.size(1))): if t rewards.size(1) - 1: next_value 0 else: next_value values[:, t 1] delta rewards[:, t] gamma * next_value - values[:, t] last_gae delta gamma * lam * last_gae advantages[:, t] last_gae returns advantages values return advantages, returns这段代码里有个细节最后一个token的next_value设为0因为序列结束了没有后续reward。另外returns advantages values这是critic的训练目标。GAE算完之后通常还要做advantage normalization减均值除标准差这能显著提升训练稳定性。我试过不做normalization训练前期loss波动特别大。4. PPO的clip机制与loss构造为什么它能稳住训练4.1 importance sampling ratio的物理含义PPO的核心是importance sampling ratio定义为ratio π_new(a|s) / π_old(a|s)这个ratio的物理含义是新策略选择这个动作的概率是旧策略的多少倍。如果ratio1说明新旧策略对这个动作的偏好一样ratio2说明新策略更倾向于选这个动作ratio0.5说明新策略不太想选它了。在策略梯度里我们用旧策略采的数据来估计新策略的梯度所以需要乘上这个ratio做修正。但ratio不能太大否则方差爆炸。PPO的clip就是给ratio设了个上下界。4.2 clip的两种写法与梯度截断效果PPO的clipped objective有两种等价写法我习惯用第一种def ppo_loss(new_logprobs, old_logprobs, advantages, clip_eps0.2): ratio torch.exp(new_logprobs - old_logprobs) surr1 ratio * advantages surr2 torch.clamp(ratio, 1 - clip_eps, 1 clip_eps) * advantages policy_loss -torch.min(surr1, surr2).mean() return policy_loss这里的逻辑是当advantage为正这个动作好我们希望增大ratio但最多增到1ε当advantage为负这个动作差我们希望减小ratio但最多减到1-ε。torch.min保证了无论哪种情况梯度都不会超过clip边界。实测下来这个clip机制是PPO能稳住训练的关键。我做过对比实验去掉clip训练到200步左右ratio就会飙到5以上loss直接NaN加上clipratio基本稳定在0.8到1.2之间训练能跑到几千步。4.3 value loss与entropy bonus的配比PPO的总loss由三部分组成total_loss policy_loss vf_coef * value_loss - ent_coef * entropyvalue_loss是critic的训练目标通常用MSE(returns - values)^2。vf_coef控制它在总loss里的权重典型值0.5或1.0。entropy bonus鼓励策略保持探索性防止过早收敛到确定性策略。ent_coef通常取0.01左右。在RLHF里entropy bonus要小心用因为语言模型的输出空间太大entropy本身就不低加太多会让模型输出变得随机、不连贯。我踩过的坑是一开始把ent_coef设成0.05结果模型生成的回复开始出现重复和无意义token。后来降到0.01输出质量明显改善。这个参数跟任务强相关建议从小值开始试。5. 实操搭一个最小可用的PPO训练循环5.1 环境准备与依赖版本先列一下我用的环境版本不匹配是新手最容易卡住的地方python3.10 torch2.1.0 transformers4.36.0 accelerate0.25.0 peft0.7.0 trl0.7.4trl库封装了PPOTrainer但我建议先手写一遍loop理解每个环节在干什么再用封装好的工具。手写一遍之后你会发现那些框架帮你处理的细节比如mask、padding、梯度累积才是真正容易出错的地方。5.2 数据准备prompt dataset的构造要点RLHF的prompt数据集不需要label只需要输入文本。但有几个细节要注意prompt长度要控制。太短的prompt比如几个词生成空间太大reward信号稀疏太长的prompt占显存而且模型可能直接复述。我一般控制在50到200 token之间。多样性要够。如果prompt全是同一类问题模型会过拟合到那个分布KL散度很快爆炸。要去重。重复的prompt会导致advantage估计有偏因为同一个输入被采样多次reward的方差被低估。from datasets import Dataset prompts [ 解释一下什么是梯度下降, 写一段Python代码实现快速排序, 用一句话总结相对论, # ... 更多prompt ] dataset Dataset.from_dict({prompt: prompts})5.3 训练主循环从rollout到update的完整代码下面是一个简化版的PPO训练循环去掉了分布式和混合精度的部分但保留了核心逻辑for iteration in range(num_iterations): # 1. 采样 batch next(dataloader) prompts batch[prompt] with torch.no_grad(): responses actor.generate( prompts, max_new_tokens128, do_sampleTrue, temperature0.7, ) old_logprobs compute_logprobs(actor(responses).logits, responses) ref_logprobs compute_logprobs(reference(responses).logits, responses) rewards reward_model(responses) values critic(responses) # 2. 算KL和调整后的reward kl old_logprobs - ref_logprobs rewards_with_kl rewards - kl_coef * kl # 3. 算advantage advantages, returns compute_gae(rewards_with_kl, values) advantages (advantages - advantages.mean()) / (advantages.std() 1e-8) # 4. 多epoch更新 for epoch in range(ppo_epochs): new_logprobs compute_logprobs(actor(responses).logits, responses) new_values critic(responses) policy_loss ppo_loss(new_logprobs, old_logprobs, advantages) value_loss F.mse_loss(new_values, returns) entropy compute_entropy(actor(responses).logits) total_loss policy_loss 0.5 * value_loss - 0.01 * entropy optimizer.zero_grad() total_loss.backward() torch.nn.utils.clip_grad_norm_(actor.parameters(), 1.0) optimizer.step()这段代码里梯度裁剪clip_grad_norm_是必须的。RL的梯度方差比监督学习大得多不裁剪的话很容易出现梯度爆炸。我一般设max_norm1.0实测比较稳。5.4 关键参数速查表与调参顺序调参这件事顺序比数值更重要。我推荐的调参顺序是优先级参数推荐范围作用调参信号1kl_coef0.01-0.1控制策略偏离KL10说明太小输出无变化说明太大2clip_eps0.1-0.2限制策略更新幅度ratio频繁触边说明太小3lr1e-6到5e-6学习率loss震荡说明太大4ppo_epochs2-4每批数据更新次数超过4容易过拟合5batch_size32-128采样批量显存允许下越大越稳6ent_coef0.001-0.01探索强度输出重复说明太大先调kl_coef因为它直接决定训练是否稳定。然后调clip_eps和lr这两个影响收敛速度。最后微调其他参数。6. 常见问题与排查技巧实录6.1 KL散度爆炸的三种典型原因KL爆炸是RLHF训练最常见的崩溃模式表现是KL值从个位数突然跳到几百甚至几千。我遇到过三种原因第一种是reward hacking。Reward model有缺陷模型找到了钻空子的方式比如生成重复的、讨好的、或者特定格式的文本reward很高但质量很差。这时候KL会飙升因为模型在拼命偏离reference去追reward。解决办法是提高kl_coef或者给reward加长度惩罚、重复惩罚。第二种是学习率太大。策略更新步子太大一次更新就跑到分布外。表现是KL在训练早期就爆炸。解决办法是把lr降到1e-6甚至5e-7。第三种是reference model加载错误。如果reference model和actor初始化不一致比如reference加载的是base模型actor加载的是SFT模型KL从一开始就会很大。这个坑很隐蔽建议训练前先跑一步检查初始KL是否接近0。6.2 reward不涨但loss在降的排查思路这个现象说明模型在优化某个东西但不是你想要的reward。排查步骤检查reward model的输出分布。如果所有reward都差不多方差很小说明reward model区分度不够模型学不到有效信号。检查advantage的分布。如果advantage全是正的或全是负的说明value估计有偏GAE没算对。检查KL惩罚是否过大。如果kl_coef设成1.0reward信号会被KL淹没模型只顾着不偏离reference。检查mask是否正确。如果padding位置的loss没mask掉梯度会被无效token稀释。我遇到过一次排查了两天才发现是reward model的tokenizer和actor不一致导致reward分数全是噪声。这种问题只能靠打印中间变量来定位。6.3 显存不够时的五种优化手段显存是RLHF训练的硬约束我按性价比排序梯度检查点gradient checkpointing用时间换显存能省30%到50%的激活显存。代价是训练速度慢20%左右。FSDP或ZeRO-3把参数、梯度、优化器状态分片到多卡。7B模型用4卡FSDP基本能跑起来。reference model预计算reference的logprob在训练中是不变的因为reference不更新可以提前算好存到磁盘训练时直接读。这能省一个模型的显存。降低max_length序列长度对显存是平方级影响。从512降到256显存能省一半以上。混合精度bf16比fp32省一半显存而且数值稳定性比fp16好。现在基本是标配。提示reference model预计算这个技巧特别实用但要注意prompt和response的对应关系。如果训练时重新采样了response预计算的logprob就对不上了。所以这个技巧只适用于response固定的场景比如offline RLHF。6.4 训练不稳定的信号识别速查表信号可能原因快速验证解决方向KL突然飙升reward hacking / lr太大打印reward和KL曲线提高kl_coef / 降lrratio频繁触边clip_eps太小统计ratio1.2的比例增大clip_epsvalue loss不降critic学习率不匹配单独看value loss曲线调vf_coef或critic lr输出重复entropy太低看生成样本增大ent_coef梯度norm爆炸梯度未裁剪打印grad_norm加clip_grad_normreward方差为0reward model失效看reward分布换reward model这张表是我踩坑踩出来的基本覆盖了80%的训练异常。建议训练时把这些指标都打到TensorBoard上出问题第一时间能定位。7. 一些工程上的个人体会跑RLHF训练这段时间最大的感受是算法本身不难难的是工程细节。PPO的公式就那么几行但要让它在四个模型、多卡、混合精度的环境下稳定跑起来需要处理的边界情况非常多。我现在养成的习惯是每次改完代码先跑一个mini batchbatch_size2max_length64把每个中间变量的形状、均值、方差都打印一遍确认无误再上全量。这个习惯帮我省了无数次的训练到一半崩了。另外KL系数的调度也值得试试。固定kl_coef在训练后期往往太松模型会慢慢跑偏。我试过用线性warmup加cosine衰减前期严格约束、后期逐渐放开效果比固定值好一些。这个没有标准答案得根据任务调。最后说一个容易被忽视的点reward model的质量决定了RLHF的上限。如果reward model本身有偏PPO再稳也训不出好模型。我建议在跑PPO之前先拿一批样本人工评估一下reward model的打分确认它和人类偏好的一致性。这一步花的时间比后面调参省的时间多得多。