ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

RLHF-PPO训练框架实战:从Loop设计到四模型协同的完整拆解

RLHF-PPO训练框架实战:从Loop设计到四模型协同的完整拆解 1. RL训练基础从Loop到RLHF-PPO的完整拆解搞强化学习训练框架这件事我踩过的坑比大多数人跑过的实验都多。最开始我以为RL Infra就是把模型丢进去、把loss跑起来就完事了结果第一次跑PPO的时候reward曲线直接起飞然后崩掉排查了整整两天才发现是advantage normalization的位置放错了。从那以后我就明白一个道理RL训练的基础设施核心不在于算法有多花哨而在于你对loop的理解有多深、对PPO每个环节的数据流有多清楚。这篇内容主要面向两类人一是刚接触RLHF训练、想搞清楚PPO到底在训练循环里干了什么的工程师二是已经跑过一些RL实验但对训练框架的loop设计、数据流转、关键参数选择还比较模糊的从业者。我会从最基础的训练循环讲起把RLHF-PPO的完整链路拆开配上可运行的代码骨架把每个环节为什么这么设计、参数怎么算、坑在哪里都说清楚。关键词RL、RLHF、PPO、loop、code会贯穿全文但不会为了堆词而堆词每个概念都会落到实际代码和操作上。先说结论性的认知RLHF中的PPO训练本质上是一个四模型协同的loop工程。Actor负责生成、Critic负责评估、Reward Model负责打分、Reference Model负责约束。这四个模型在每一轮训练循环里各司其职任何一个环节的数据流断了或者参数配错了整个训练就会崩。很多人觉得PPO难难的不是算法公式难的是把这个loop工程搭稳。2. 训练Loop的整体设计与思路拆解2.1 为什么RL训练需要一个显式的Loop结构监督学习训练是一个很直接的流程取batch、前向、算loss、反向、更新。整个循环里数据是静态的模型只跟固定的标签打交道。但RL训练不一样它的数据是模型自己生成的每一轮循环产生的数据都会影响下一轮模型的行为。这就意味着你不能像监督学习那样把数据加载和模型更新解耦必须设计一个显式的loop结构来管理“生成-评估-更新”这个闭环。我见过不少人一开始写RL训练代码习惯性地套监督学习的trainer模板结果发现数据加载器根本没法用因为样本不是预先存在的而是当前策略实时产出的。这就是为什么RL Infra的第一个核心问题就是loop怎么设计。一个典型的RLHF-PPO训练loop包含以下阶段采样阶段Actor生成response、评估阶段Reward Model打分、Reference Model算KL、优势估计阶段Critic算value、计算GAE、更新阶段PPO clip loss反向传播。这四个阶段在每一轮循环里顺序执行但它们的计算图关系和数据依赖关系需要仔细设计。注意采样阶段必须用torch.no_grad()包裹否则你会在显存里保留整个生成过程的计算图几轮下来直接OOM。这个坑我在第一次跑PPO的时候踩过当时以为是batch size太大调小了一半还是爆后来才发现是生成阶段没关梯度。2.2 四模型协同的架构选型与显存权衡RLHF-PPO最让人头疼的就是四个模型同时存在Actor、Critic、Reward、Reference。如果每个模型都是7B规模即使用bf16精度光模型权重就要占掉大约56GB显存7B × 2 bytes × 4。这还没算优化器状态、梯度、激活值。所以实际工程中必须做架构选型上的权衡。常见的方案有三种。第一种是Actor和Critic共享底层参数只在最后一层分叉出policy head和value head这样能省掉一个完整模型的显存。第二种是Reference Model和Reward Model用更小的模型比如Actor用7BReference用同一个初始权重的7B但冻结Reward用1.5B或者3B。第三种是Actor和Reference共享同一个模型实例通过开关adapter来切换Reference阶段直接禁用adapter即可。我个人的经验是如果你的显存预算在80GB单卡级别建议采用Actor-Critic共享底座加独立Reference的方案。Reward Model可以单独部署在一个更小的卡上或者用API调用。如果显存更紧张那就把Reference和Actor做成同一个模型加LoRA adapter切换训练时只更新adapter部分Reference阶段把adapter关掉就是原始模型。这个方案实测下来显存占用能压到原来的60%左右。方案显存占用实现复杂度适用场景四模型独立极高低多卡A100/H100集群Actor-Critic共享底座高中单卡80GBActor-Reference共享LoRA中高单卡40-80GB全部共享多Head低极高实验性场景2.3 Loop的粒度选择Step-level还是Episode-level训练loop的粒度也是一个需要想清楚的问题。所谓step-level就是每生成一个token就做一次评估和更新episode-level则是等整个response生成完毕后再统一评估和更新。绝大多数RLHF-PPO的实现都是episode-level的因为Reward Model通常只对完整response打分而且PPO的advantage估计需要完整的轨迹。但这里有个细节虽然更新是episode-level的但KL penalty的计算可以是token-level的。也就是说在生成每个token的时候同时计算当前policy和reference policy的log prob差值累积起来作为KL项。这样做的好处是KL约束更精细不会因为response长度不同导致KL项被稀释。我在实际项目中的做法是生成阶段同时记录每个token的log probActor和Reference各一份生成结束后统一计算KL和Reward然后做GAE和PPO更新。3. 核心细节解析与实操要点3.1 PPO中Advantage估计的完整计算链路PPO的核心在于advantage的估计而advantage的估计依赖Critic给出的value。整个计算链路是这样的首先Critic对每个token位置输出一个value估计然后我们用Reward Model给出的最终reward和KL penalty组合成每个token的即时reward接着用GAEGeneralized Advantage Estimation把即时reward和value估计结合起来算出advantage。具体来说即时reward的构造是最后一个token位置加上Reward Model的打分减去KL penalty中间token位置只有KL penalty的负值。这里有个容易搞错的地方KL penalty是加在reward里的不是单独作为一个loss项。很多人会把KL直接加到loss里这样做的效果和加在reward里是不一样的。加在reward里意味着KL会影响advantage进而影响policy和value的更新方向加在loss里则只是约束policy不要偏离太远不参与advantage计算。实践中加在reward里的效果更稳定。GAE的计算涉及两个参数gamma和lambda。gamma是折扣因子通常设0.99或者1.0lambda是GAE的平滑参数通常设0.95。计算方式是反向遍历整个序列每一步的advantage等于即时reward加上gamma乘以下一步的value再减去当前value再加上gamma乘以lambda乘以下一步的advantage。这个计算过程用代码表示就是几行循环但顺序不能错必须从最后一个token往前算。def compute_gae(rewards, values, gamma0.99, lam0.95): advantages [] gae 0 for t in reversed(range(len(rewards))): if t len(rewards) - 1: next_value 0 else: next_value values[t 1] delta rewards[t] gamma * next_value - values[t] gae delta gamma * lam * gae advantages.insert(0, gae) returns [adv val for adv, val in zip(advantages, values)] return advantages, returns提示advantage算完之后一定要做normalization减均值除标准差。这一步对训练稳定性影响巨大不做normalization的话policy gradient的方差会非常大训练很容易崩。但注意normalization的维度是每个batch内做不要跨batch累积统计量。3.2 Reward Model打分与KL penalty的配合方式Reward Model的打分通常只在最后一个token位置给出一个标量。但KL penalty是每个token都有的。这两者怎么配合我的做法是构造一个和序列等长的reward数组初始化为0然后把KL penalty的负值填到每个位置最后把Reward Model的打分加到最后一个位置。这样Critic在学习的时候中间位置的value会逐渐学会预测累积的KL penalty最后一个位置的value会学会预测KL penalty加最终reward。KL penalty的系数设置也很讲究。系数太小policy会跑偏生成的内容会变得奇怪系数太大policy几乎不更新训练没有效果。常用的初始值是0.01到0.1之间。我一般从0.04开始试如果发现KL散度增长太快就调大如果KL一直很低但reward也不涨就调小。还有一个技巧是使用adaptive KL control设定一个目标KL值然后根据实际KL动态调整系数。这个在HuggingFace的PPO实现里有现成的但自己写的话也不复杂就是根据KL和目标值的比值来缩放系数。3.3 PPO Clip Loss的实现细节与常见错误PPO的clip loss看起来简单但实现的时候有几个细节特别容易出错。首先是ratio的计算ratio等于当前policy的log prob减去旧policy的log prob然后取exp。这里旧policy的log prob是在采样阶段记录的必须在更新之前就固定下来不能在更新过程中重新计算。我见过有人每次更新step都重新算一遍旧log prob结果ratio永远是1clip完全失效。其次是clip的范围通常设0.2。但要注意clip是对ratio做clip不是对loss做clip。loss的最终形式是取unclipped loss和clipped loss的最大值当advantage为正时或最小值当advantage为负时。这个min/max的选择取决于advantage的符号写反了会导致训练方向完全错误。def ppo_clip_loss(new_log_probs, old_log_probs, advantages, clip_ratio0.2): ratio torch.exp(new_log_probs - old_log_probs) surr1 ratio * advantages surr2 torch.clamp(ratio, 1 - clip_ratio, 1 clip_ratio) * advantages policy_loss -torch.min(surr1, surr2).mean() return policy_loss还有一个容易忽略的点是value loss的裁剪。有些实现会对value loss也做clip防止value更新过猛。这个不是PPO原论文的要求但在实践中确实能提升稳定性。做法是计算value的clip版本然后取原value loss和clip value loss的最大值。4. 实操过程与核心环节实现4.1 从零搭建一个最小可运行的PPO训练循环我现在把整个训练循环的骨架代码写出来你可以直接拿去改。这个骨架包含了采样、评估、GAE计算、PPO更新四个核心环节去掉了分布式和混合精度的部分方便理解主干逻辑。import torch import torch.nn as nn from torch.optim import Adam class PPOTrainer: def __init__(self, actor, critic, reward_model, ref_model, lr1e-5, clip_ratio0.2, kl_coef0.04, gamma0.99, lam0.95, ppo_epochs4): self.actor actor self.critic critic self.reward_model reward_model self.ref_model ref_model self.optimizer Adam( list(actor.parameters()) list(critic.parameters()), lrlr ) self.clip_ratio clip_ratio self.kl_coef kl_coef self.gamma gamma self.lam lam self.ppo_epochs ppo_epochs def sample(self, prompts, max_new_tokens256): with torch.no_grad(): sequences, old_log_probs self.actor.generate( prompts, max_new_tokensmax_new_tokens ) ref_log_probs self.ref_model.log_prob(sequences) values self.critic(sequences) rewards self.reward_model(sequences) return sequences, old_log_probs, ref_log_probs, values, rewards def compute_rewards(self, rewards, old_log_probs, ref_log_probs): kl old_log_probs - ref_log_probs token_rewards -self.kl_coef * kl token_rewards[:, -1] rewards return token_rewards def train_step(self, prompts): sequences, old_log_probs, ref_log_probs, values, rewards \ self.sample(prompts) token_rewards self.compute_rewards( rewards, old_log_probs, ref_log_probs ) advantages, returns compute_gae( token_rewards, values, self.gamma, self.lam ) advantages (advantages - advantages.mean()) / \ (advantages.std() 1e-8) for _ in range(self.ppo_epochs): new_log_probs self.actor.log_prob(sequences) new_values self.critic(sequences) policy_loss ppo_clip_loss( new_log_probs, old_log_probs, advantages, self.clip_ratio ) value_loss nn.MSELoss()(new_values, returns) loss policy_loss 0.5 * value_loss self.optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_( self.actor.parameters(), 1.0 ) self.optimizer.step() return policy_loss.item(), value_loss.item()这个骨架里每个函数的作用都很明确。sample负责生成和收集旧数据compute_rewards负责构造token级别的rewardtrain_step把整个loop串起来。实际使用时你需要根据具体模型替换generate和log_prob的实现。4.2 关键参数的选取与计算过程PPO训练里有几个参数直接决定成败我逐个说下怎么选。学习率Actor的学习率通常比监督学习小一个数量级1e-6到5e-6是比较安全的范围。Critic的学习率可以稍大1e-5左右。如果发现policy loss震荡厉害先降学习率。Clip ratio默认0.2但在RLHF场景下我建议从0.1开始试。因为RLHF的reward信号通常比较稀疏clip太宽会导致更新步长过大。KL系数前面说了0.04是个不错的起点。但更好的做法是用adaptive KL目标KL设为0.01到0.02。具体计算是如果当前KL大于目标KL的1.5倍系数乘以1.5如果小于目标KL的0.5倍系数除以1.5。Batch size和mini-batch采样batch size决定了每次loop收集多少数据通常设64到256条prompt。PPO的mini-batch是在更新阶段把采样数据分成更小的块通常设8到32。mini-batch太小会导致梯度噪声大太大则失去PPO多epoch更新的意义。PPO epochs通常设4。太多会导致policy偏离旧policy太远clip频繁触发更新效率反而下降。注意KL系数的调整不要每步都做建议每10到20个训练step调整一次。频繁调整会让训练不稳定因为KL本身有波动。4.3 训练过程中的监控指标与日志记录跑RL训练不监控指标等于盲开。我必看的指标有这几个policy loss、value loss、KL散度、reward均值、response长度、clip fraction。其中clip fraction是ratio被clip的比例如果这个值长期高于0.3说明clip ratio设小了或者学习率太大了。KL散度如果持续增长不收敛说明KL系数太小。Reward均值应该整体上升但会有波动如果一直不涨说明reward信号有问题或者advantage计算有bug。我习惯每10个step打一次日志每100个step存一次checkpoint。日志里除了上述指标还会记录当前KL系数、学习率、grad norm。grad norm突然变大通常是训练要崩的前兆这时候可以手动降低学习率或者提前停止这个epoch的更新。5. 常见问题与排查技巧实录5.1 Reward不涨反降的排查思路Reward不涨是RLHF训练最常见的问题。排查顺序我一般是这样的先看KL散度如果KL很小接近0说明policy几乎没更新检查学习率是不是太小、KL系数是不是太大、advantage是不是全被normalize成接近0了。如果KL正常但reward不涨检查Reward Model的打分是否合理可以拿几条生成结果人工看一下有时候是Reward Model本身有问题。如果KL在涨但reward不涨说明policy在偏离但没往正确的方向偏检查advantage的符号是不是反了或者reward的构造是不是把KL penalty加错了位置。还有一个隐蔽的坑是prompt的分布问题。如果训练prompt太单一policy很容易过拟合到某一种回答模式reward一开始涨很快然后卡住。解决办法是增加prompt的多样性或者在reward里加一个多样性惩罚项。5.2 显存溢出与计算图泄漏的定位方法显存溢出在RL训练里太常见了。定位方法很简单在训练loop的每个阶段后面打印torch.cuda.memory_allocated()看哪个阶段显存涨得最多。通常采样阶段是显存大户因为要生成完整序列。如果采样阶段显存就爆了减小max_new_tokens或者减小采样batch size。如果是更新阶段爆检查是不是在采样时没关梯度导致计算图被保留到了更新阶段。计算图泄漏的典型症状是显存随着训练step逐渐增长最后OOM。原因是某些tensor被意外保留在了计算图里。检查方法是看有没有在torch.no_grad()外面调用了需要梯度的模型或者有没有把带梯度的tensor存到了list里跨step累积。问题现象可能原因排查方法解决方案显存逐步增长计算图泄漏打印各阶段显存检查no_grad包裹采样阶段OOM序列太长打印生成长度分布减小max_new_tokens更新阶段OOMbatch太大打印batch shape减小mini-batchKL爆炸系数太小监控KL曲线增大KL系数Reward不涨学习率太小检查grad norm增大学习率Clip fraction过高学习率太大监控clip比例降低学习率5.3 训练不稳定时的应急处理清单训练不稳定的时候不要慌按这个清单逐项检查第一advantage有没有做normalization第二旧log prob有没有在更新前固定第三KL penalty有没有加对位置第四value loss有没有和policy loss一起更新第五梯度裁剪有没有开第六学习率有没有设太大。这六项覆盖了90%的不稳定原因。我个人的经验是RLHF-PPO训练前100个step是最容易崩的熬过前100步之后通常会稳定下来。所以前期可以设一个warmup学习率从很小的值线性增加到设定值同时KL系数可以设大一点等训练稳定后再降下来。这个策略我用了很多次实测能显著降低早期崩溃的概率。6. 代码工程化的一些个人体会把上面的骨架代码变成能长期跑的训练框架还有几件事要做。第一是数据加载要异步采样和更新可以流水线化采样下一批数据的同时更新当前批这样GPU利用率能提升30%以上。第二是checkpoint要存全量状态包括actor、critic、optimizer、当前KL系数、训练step数不然断点续训的时候KL系数对不上会导致训练行为突变。第三是日志要结构化用wandb或者tensorboard记录方便对比不同超参的实验结果。我自己在实际操作中的体会是RL Infra的复杂度不在于单个模块有多难而在于模块之间的耦合关系。PPO的loop里每个环节的输出都是下一个环节的输入任何一个环节的数据格式或者数值范围出了问题都会在后续环节被放大。所以写代码的时候一定要在每个环节之间加assert检查shape、检查数值范围、检查有没有NaN。这些assert在调试阶段能帮你省下大量时间。最后再分享一个小技巧如果你在调试PPO的实现可以先用一个极简的玩具环境验证比如让Reward Model直接返回response长度的负值然后看policy能不能学会生成更短的response。这个玩具实验能在几分钟内跑完如果这个都跑不通那肯定是PPO实现有bug不用去大模型上浪费时间。等玩具环境跑通了再上真实模型效率会高很多。
返回列表