1. 从一次失败的PPO训练说起为什么loop结构才是RL Infra的骨架我第一次跑RLHF的PPO训练时盯着loss曲线看了三个小时发现reward在涨、KL在飘、value loss像心电图一样抖。当时以为是超参没调好换了三组学习率、调了两轮batch size问题依旧。后来把整个训练循环打印出来逐行对才发现真正的问题出在loop的嵌套顺序上——我把reference model的forward和reward model的scoring放在了同一个micro-batch里交替执行导致显存峰值直接翻倍梯度累积的步数也被打乱了。这件事让我意识到一个很朴素但容易被忽略的事实RL训练的基础设施本质上就是一个精心设计的loop。PPO也好、GRPO也好、DPO也好算法论文里那些公式推导看着复杂落到工程实现上核心就是采样-评估-更新这三步怎么循环、循环里每一步的数据怎么流转、显存怎么复用、梯度怎么累积。你把loop结构搞对了后面调参才有意义loop结构错了再好的超参也救不回来。这篇内容适合两类人看一类是刚接触RLHF、想搞明白PPO训练代码到底在干什么的算法工程师另一类是在做RL Infra、需要把训练循环设计得既高效又不容易出bug的系统工程师。我会从最基础的loop结构讲起把RLHF-PPO的训练循环拆到每一行代码的粒度然后讲清楚每个环节为什么这么设计、常见的坑在哪里、怎么验证自己的实现是对的。代码部分用PyTorch风格伪代码重点是逻辑而不是框架细节你换成任何训练框架都能套进去。关键词里的RL、RLHF、PPO、loop、code这五个词基本就是这篇内容的全部骨架。我不会泛泛地讲PPO算法原理——那些东西论文里都有——我讲的是当你真正要写一个能跑起来的PPO训练循环时需要做哪些决策、这些决策背后的权衡是什么。2. RLHF-PPO训练循环的四层嵌套结构2.1 最外层epoch循环与数据消费顺序PPO训练的最外层是一个epoch循环但这个epoch和 supervised learning里的epoch含义不太一样。在SFT阶段一个epoch意味着把整个数据集过一遍在PPO阶段一个epoch通常指的是用当前policy采样一批prompt生成对应的response然后用这批数据做若干次梯度更新。这里第一个容易踩的坑是prompt的消费顺序。很多人习惯性地把prompt数据集shuffle之后按顺序取但PPO的采样成本远高于SFT——每一条prompt都要经过完整的自回归生成生成长度可能几百到几千token。如果你把长prompt和短prompt混在一起采样batch内的序列长度差异会非常大padding浪费严重GPU利用率可能只有30%不到。我的做法是在数据预处理阶段就按prompt长度分桶每个batch从同一个桶里采样。这样batch内长度接近padding开销小生成速度也稳定。代价是牺牲了一点随机性但PPO本身对prompt的多样性要求没有SFT那么高这个trade-off是划算的。# 按长度分桶的prompt sampler伪代码 class BucketedPromptSampler: def __init__(self, prompts, bucket_size50): self.buckets defaultdict(list) for p in prompts: bucket_id len(p) // bucket_size self.buckets[bucket_id].append(p) def sample_batch(self, batch_size): # 随机选一个桶从桶里采样 bucket_id random.choice(list(self.buckets.keys())) bucket self.buckets[bucket_id] if len(bucket) batch_size: # 桶不够大时跨桶补充 ... return random.sample(bucket, batch_size)2.2 第二层rollout循环与生成策略rollout循环是PPO训练里最耗时的部分。每一步你要用当前policy对batch里的prompt做自回归生成同时记录每个token的log probability。这里有几个关键决策生成时用不用KV Cache必须用。不用KV Cache的话生成长度L的序列需要O(L²)的attention计算实际训练中根本跑不动。但用了KV Cache之后你要注意生成时的log prob和训练时的log prob可能有数值差异——因为生成时是逐token算的训练时是整序列并行算的。这个差异在早期实现里经常导致ratio偏离1后来大家统一用训练时的forward重新算一遍old log prob才把这个坑填上。temperature和top-p怎么设我的经验是rollout阶段temperature设0.7到1.0之间top-p设0.9到0.95。temperature太低会导致生成多样性不足PPO的探索能力下降太高则生成质量差reward model给的分普遍偏低梯度信号弱。这个参数没有理论最优值得根据你的reward model和任务类型调。生成长度怎么控制设max_new_tokens是必须的但更重要的是设一个合理的值。太短了response不完整reward model没法准确打分太长了显存吃不消而且大部分token对reward的贡献很小。我一般会先统计SFT阶段response的长度分布取95分位数作为max_new_tokens。# rollout阶段的核心逻辑 def rollout(policy_model, prompts, max_new_tokens512, temperature0.8): with torch.no_grad(): # 生成response generated policy_model.generate( prompts, max_new_tokensmax_new_tokens, temperaturetemperature, do_sampleTrue, return_dict_in_generateTrue, output_scoresTrue ) # 关键用训练模式重新算一遍log prob # 而不是直接用generate返回的scores full_ids generated.sequences log_probs compute_log_probs(policy_model, full_ids) return full_ids, log_probs2.3 第三层advantage计算与GAE的数值稳定性拿到rollout数据之后下一步是算advantage。PPO用的是GAEGeneralized Advantage Estimation公式本身不复杂但实现的时候有几个数值稳定性的坑。第一个坑是reward的归一化。reward model输出的分数范围可能很大直接拿去做GAEadvantage的方差会非常大梯度更新不稳定。常见的做法是在batch内做reward normalization减均值除标准差。但要注意如果你用了KL penaltyKL项也要一起归一化否则两项量纲不一致。第二个坑是GAE的lambda和gamma参数。gamma是discount factor在RLHF里通常设1.0因为我们的episode就是整个response不存在跨episode的折扣。lambda控制bias-variance trade-off设0.95是常见选择。但这两个参数和reward normalization是耦合的——如果你把reward归一化了gamma和lambda的效果会发生变化需要重新调。第三个坑是value model的预测值。GAE需要value model对每个token的value估计如果value model训练得不好advantage的估计会有很大偏差。我的做法是在PPO训练之前先用一批数据把value model warm up一下让它对reward的量级有个基本概念。def compute_gae(rewards, values, gamma1.0, lam0.95): rewards: [batch, seq_len] values: [batch, seq_len] batch_size, seq_len rewards.shape advantages torch.zeros_like(rewards) last_gae 0 for t in reversed(range(seq_len)): if t seq_len - 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, returns2.4 最内层PPO update与ratio clipping最内层是PPO的梯度更新循环。这里核心是ratio clipping机制新policy和旧policy的概率比不能偏离太远否则梯度会被clip掉。实现上要注意的是ratio是在token级别算的但loss是在序列级别聚合的。常见做法是对每个token算ratio和clipped ratio取min之后乘上advantage然后对所有token求平均。但这里有个细节padding token不能算进去否则会拉低loss。所以需要一个mask把padding位置排除掉。另一个细节是KL penalty的计算。KL散度有两种估计方式一种是直接用log prob的差另一种是用近似公式。我推荐用log prob差的方式虽然方差大一点但无偏。KL penalty的系数需要调太小了policy会偏离reference model太远生成质量下降太大了policy学不动reward涨不上去。def ppo_loss(new_log_probs, old_log_probs, advantages, values, returns, mask, clip_ratio0.2, vf_coef0.5, kl_coef0.01): # 计算ratio ratio torch.exp(new_log_probs - old_log_probs) # clipped ratio clipped_ratio torch.clamp(ratio, 1 - clip_ratio, 1 clip_ratio) # policy loss policy_loss -torch.min(ratio * advantages, clipped_ratio * advantages) policy_loss (policy_loss * mask).sum() / mask.sum() # value loss value_loss ((values - returns) ** 2 * mask).sum() / mask.sum() # KL penalty kl (old_log_probs - new_log_probs) * mask kl kl.sum() / mask.sum() total_loss policy_loss vf_coef * value_loss kl_coef * kl return total_loss, policy_loss, value_loss, kl3. 显存与计算效率loop设计中的工程权衡3.1 四个模型同时驻留显存的现实问题RLHF-PPO训练和普通SFT最大的工程差异在于你同时需要四个模型——policy model、reference model、reward model、value model。如果每个模型都是7B参数用fp16存储光权重就要占56GB显存加上optimizer state、gradient、activation单卡根本放不下。常见的解决方案有三种方案一模型并行。把四个模型分散到不同GPU上policy和value放一起因为都要训练reference和reward放另外的卡上只做forward。这种方案的问题是通信开销大尤其是policy和reference之间要传log prob跨卡通信会成为瓶颈。方案二offload。把reference和reward model offload到CPU内存需要的时候再load到GPU。这种方案省显存但慢适合显存极度紧张的场景。方案三共享底座。如果policy和reference来自同一个SFT模型可以让它们共享底层参数只在顶层做区分。value model也可以和policy共享底座加一个value head。这样显存占用能降到原来的40%左右。这是目前最主流的做法。我自己的经验是如果显存允许优先用方案三如果显存不够方案二比方案一更实用因为PPO训练本身计算密度不高offload带来的延迟可以接受。3.2 micro-batch与gradient accumulation的配合PPO的batch size通常很大因为要保证advantage估计的稳定性但显存有限所以必须用gradient accumulation。这里的关键是rollout的batch size和update的batch size可以不一样。具体来说你可以用batch size 64做rollout采样64条数据然后把这64条数据分成4个micro-batch每个16条做4次forward-backward累积梯度之后再更新一次。这样显存占用按16条算但等效batch size是64。要注意的是PPO的ratio clipping是在update阶段算的如果你把batch拆成micro-batch每个micro-batch的ratio是独立算的这会导致clipping行为和理论上有偏差。实践中这个偏差可以接受但如果你的clip_ratio设得很小比如0.1偏差会比较明显。# gradient accumulation的PPO update def ppo_update_with_accumulation(policy, value_model, batch, micro_batch_size16, accumulation_steps4): policy.train() value_model.train() optimizer.zero_grad() total_loss 0 for i in range(accumulation_steps): micro_batch slice_batch(batch, i, micro_batch_size) loss, _, _, _ ppo_loss( new_log_probscompute_log_probs(policy, micro_batch), old_log_probsmicro_batch[old_log_probs], advantagesmicro_batch[advantages], valuesvalue_model(micro_batch[input_ids]), returnsmicro_batch[returns], maskmicro_batch[mask] ) loss loss / accumulation_steps loss.backward() total_loss loss.item() torch.nn.utils.clip_grad_norm_(policy.parameters(), 1.0) optimizer.step() return total_loss3.3 生成阶段的显存峰值控制rollout阶段的显存峰值往往比update阶段还高因为自回归生成需要维护KV Cache而且生成长度不确定。控制显存峰值有几个实用技巧动态batch。不要固定batch size而是根据当前显存余量动态调整。显存多的时候多采样几条显存紧张的时候少采样几条。这个逻辑可以用一个简单的反馈循环实现。提前终止。如果一条response已经生成了EOS token就不要再继续生成了。很多实现里为了对齐长度会继续生成padding token这是纯浪费。分块生成。如果max_new_tokens设得很大比如2048可以分块生成每生成256个token就检查一次是否所有序列都结束了结束了就提前退出。提示rollout阶段的显存峰值通常出现在生成最后几个token的时候因为此时KV Cache最长。如果你在训练初期就OOM先检查是不是max_new_tokens设太大了而不是急着加卡。4. 代码实现中的五个隐蔽陷阱4.1 log prob的计算方式不一致这是最常见也最隐蔽的bug。生成的时候模型返回的scores是每个位置的logits你对logits做log_softmax再gather对应token的log prob这是一套计算路径。训练的时候你对完整序列做一次forward得到logits再做log_softmax和gather这是另一套路径。两条路径在数值上可能有微小差异累积到ratio上就可能偏离1。解决方案很简单但容易被忽略rollout阶段生成完之后用训练时的forward重新算一遍old log prob。多花一次forward的时间但能保证ratio的初始值是准确的1.0。4.2 reward model的padding位置处理reward model通常只在最后一个token位置输出一个标量分数。但如果你把padding token也喂进去reward model可能会在padding位置也输出分数如果你不小心取了padding位置的分数reward就完全错了。正确的做法是用attention mask找到每个序列的最后一个有效token位置只取那个位置的reward。这个操作看起来简单但在batch里每个序列长度不同的时候很容易写错索引。def get_last_token_reward(reward_model, input_ids, attention_mask): outputs reward_model(input_ids, attention_maskattention_mask) # 找到每个序列最后一个有效token的位置 last_token_idx attention_mask.sum(dim1) - 1 # 取对应位置的reward rewards outputs[torch.arange(outputs.size(0)), last_token_idx] return rewards4.3 KL penalty的方向搞反KL散度是不对称的KL(p||q)不等于KL(q||p)。在PPO里我们通常用KL(policy||reference)作为penalty意思是惩罚policy偏离reference的程度。但有些实现里写成了KL(reference||policy)虽然数值上差异不大但在policy和reference差异很大的时候方向搞反会导致penalty的行为不符合预期。更隐蔽的是有些实现用log prob的差来估计KL但符号搞反了。正确的写法是kl old_log_probs - new_log_probs因为old policy是referencenew policy是当前policy。如果你写反了KL penalty会变成负的policy会越跑越偏。4.4 advantage的mask处理advantage是在token级别算的但padding位置的advantage是无效的。如果你不mask掉padding位置这些位置的advantage会参与loss计算引入噪声。更严重的是如果padding位置的advantage恰好很大会把梯度带偏。mask的处理要贯穿整个loss计算policy loss要mask、value loss要mask、KL penalty也要mask。我见过不少实现只mask了policy loss忘了value loss和KL结果训练不稳定。4.5 梯度裁剪的顺序梯度裁剪应该在所有micro-batch的梯度都累积完之后、optimizer.step()之前做。如果你在每个micro-batch之后都裁剪一次等效于对梯度做了多次缩放训练动态会完全变样。另外policy和value model的梯度要分开裁剪。如果它们共享底座裁剪的时候要注意不要互相影响。我的做法是给policy和value分别设clip值policy用1.0value用0.5因为value loss通常比policy loss大需要更强的裁剪。5. 怎么验证你的PPO loop写对了5.1 三个必做的sanity check写完PPO训练循环之后不要急着跑完整训练先做三个检查检查一ratio初始值是否为1。在第一次update之前new_log_probs和old_log_probs应该完全相等ratio应该全是1.0。如果不是说明log prob的计算路径有问题。检查二KL初始值是否为0。第一次update之前policy和reference是同一个模型KL应该是0。如果不是说明KL的计算或者模型加载有问题。检查三单步过拟合。拿一条固定数据反复做update看reward能不能过拟合到很高。如果跑了100步reward都不涨说明梯度信号有问题可能是advantage算错了或者loss写反了。5.2 训练过程中的监控指标正式训练时我一般会监控这几个指标指标正常范围异常含义policy loss-0.1 ~ 0.1持续为正说明advantage符号可能反了value loss逐渐下降持续不降说明value model没学好KL0.01 ~ 0.5超过1说明policy偏离太远ratio mean0.9 ~ 1.1偏离太多说明clip_ratio设太小reward mean逐渐上升不涨说明学习率太小或reward有问题clip fraction0.1 ~ 0.3太高说明更新步长太大这些数值不是绝对的不同任务会有差异但量级和趋势是通用的。如果某个指标明显偏离先检查对应的代码逻辑再调超参。5.3 一个最小可复现的调试流程如果你发现训练不稳定可以按这个流程逐步排查把batch size降到1去掉gradient accumulation看单条数据能不能正常更新。把clip_ratio设成很大的值比如10关掉clipping看训练是否稳定。如果稳定了说明clipping逻辑有问题。把KL penalty设成0看reward能不能涨。如果能涨但KL爆炸说明KL系数需要调大。把value model冻住只用policy的梯度更新看reward能不能涨。如果能涨说明value model是瓶颈。这个流程能帮你快速定位问题出在哪个环节比盲目调参高效得多。6. 从loop结构看RL Infra的扩展方向把PPO的训练循环拆清楚之后你会发现RL Infra的很多扩展方向都是围绕这个loop做优化。比如异步rollout就是把rollout循环和update循环解耦用单独的worker做生成训练进程只管更新比如replay buffer就是把历史rollout数据存下来提高数据利用率比如分布式PPO就是把batch维度切到多卡上每张卡负责一部分rollout和update。这些扩展的底层逻辑都是一样的识别loop中的瓶颈环节然后针对性地优化。rollout慢就优化生成update慢就优化梯度计算显存不够就优化模型驻留策略。你把基础loop搞明白了这些扩展都是水到渠成的事。我自己在实现异步rollout的时候踩过一个坑rollout worker和trainer之间的数据同步如果做得不好会出现policy版本不一致的问题——worker用旧policy生成的数据trainer用新policy做updateratio会偏离1很多。解决方案是给每个rollout数据打上policy版本号trainer只接受版本号匹配的数据或者用importance sampling做修正。代码部分我尽量保持了框架无关的伪代码风格你换成DeepSpeed、Megatron或者FSDP都能套进去。核心是把loop的每一层逻辑理清楚知道每一步在干什么、为什么这么干、不这么干会出什么问题。这比记住某个框架的API重要得多。