从零实现RLHF-PPO:强化学习训练Loop与工程避坑指南

发布时间:2026/10/6 20:19:01
从零实现RLHF-PPO:强化学习训练Loop与工程避坑指南
1. RL训练基础从Loop到RLHF-PPO的完整拆解搞强化学习训练框架这件事我踩过的坑比大多数人写过的代码都多。很多人一上来就盯着PPO的clip系数调参结果连一个最基础的training loop都没跑通loss曲线跟心电图似的乱跳。这篇东西就是写给那些准备自己动手搭RL训练流程的人——不管你是刚入门想搞清楚RLHF到底怎么跑起来还是已经跑过几轮但总觉得哪里不对劲想回头补基础都能从这里找到能直接抄的代码结构和踩坑经验。先把话说清楚RL训练的核心不是什么高深算法而是一个能稳定跑起来的loop。这个loop要处理环境交互、数据收集、策略更新、模型同步这几件事每一件单独拎出来都不难但串在一起就容易出问题。而RLHF-PPO本质上就是在这个loop上加了一层人类偏好奖励模型把环境给的reward换成了模型打分。理解了基础loopPPO和RLHF就是在这个骨架上加东西。我下面会从最朴素的policy gradient loop开始一步步推到RLHF-PPO的完整实现代码用PyTorch写尽量保持可运行、可复现。你不需要有分布式训练的经验单卡就能跟着跑。2. 最朴素的RL Loop长什么样2.1 一个loop的四个核心动作任何RL训练loop不管多复杂拆开来看就是四件事在循环采样Rollout用当前策略和环境交互收集一批轨迹数据评估Evaluate算每条轨迹的回报或者算advantage更新Update用这批数据更新策略参数同步Sync把新策略同步到采样端如果是分离架构这四步听起来简单但每一步都有坑。采样阶段最容易出问题的是数据分布偏移——你采的时候用的是旧策略更新的时候策略已经变了这个gap就是off-policy问题的根源。PPO的clip机制本质上就是在控制这个gap不要太大。我见过太多人写loop的时候把采样和更新混在一起一边采一边更新结果策略更新太快采样数据完全失效。正确的做法是先采一批固定大小的数据然后在这批数据上做多轮更新这就是PPO里n_epochs参数的由来。2.2 用代码把loop写出来先看一个最简版本的loop骨架不涉及PPO就是纯policy gradientimport torch import torch.nn as nn import torch.optim as optim class PolicyNet(nn.Module): def __init__(self, state_dim, action_dim): super().__init__() self.fc nn.Sequential( nn.Linear(state_dim, 64), nn.Tanh(), nn.Linear(64, action_dim) ) def forward(self, x): return torch.softmax(self.fc(x), dim-1) def collect_trajectories(env, policy, max_steps200): state env.reset() log_probs [] rewards [] states [] actions [] for _ in range(max_steps): state_tensor torch.FloatTensor(state).unsqueeze(0) probs policy(state_tensor) dist torch.distributions.Categorical(probs) action dist.sample() next_state, reward, done, _ env.step(action.item()) log_probs.append(dist.log_prob(action)) rewards.append(reward) states.append(state) actions.append(action.item()) state next_state if done: break return states, actions, log_probs, rewards def compute_returns(rewards, gamma0.99): returns [] R 0 for r in reversed(rewards): R r gamma * R returns.insert(0, R) returns torch.FloatTensor(returns) # 标准化这一步很关键 returns (returns - returns.mean()) / (returns.std() 1e-8) return returns def train_loop(env, policy, optimizer, num_iterations1000): for it in range(num_iterations): # 1. 采样 states, actions, log_probs, rewards collect_trajectories(env, policy) # 2. 评估 returns compute_returns(rewards) log_probs torch.stack(log_probs) # 3. 更新 loss -(log_probs * returns).mean() optimizer.zero_grad() loss.backward() optimizer.step() if it % 50 0: print(fIter {it}, Loss: {loss.item():.4f}, fAvg Return: {sum(rewards):.2f})这段代码能跑但有几个问题必须点出来。第一compute_returns里的标准化是必须的不标准化的话梯度方差会大到没法训练。第二这个loop是纯on-policy的采一批更新一次数据用完就扔样本效率极低。第三没有advantage的概念直接用return做权重方差还是太大。2.3 从loop到PPO为什么要加clip上面那个loop最大的问题是步长不可控。策略梯度更新的时候如果某一步更新太大策略直接崩掉后面采的数据全是垃圾训练就废了。PPO的clip就是来解决这个问题的。PPO的核心思想是新策略和旧策略的概率比不要偏离太远。具体做法是在loss里加一个clip操作def ppo_loss(new_log_probs, old_log_probs, advantages, clip_epsilon0.2): ratio torch.exp(new_log_probs - old_log_probs) surr1 ratio * advantages surr2 torch.clamp(ratio, 1 - clip_epsilon, 1 clip_epsilon) * advantages return -torch.min(surr1, surr2).mean()这个clip的含义是当ratio超过1epsilon或者低于1-epsilon的时候梯度就被截断了策略不会因为一个batch的数据更新太猛。clip_epsilon一般取0.1到0.2我实测下来0.2在大多数任务上比较稳如果训练不稳定可以降到0.1。注意clip只对advantage为正的样本起“限制增大”的作用对advantage为负的样本起“限制减小”的作用。这个细节很多人没搞清楚导致调参的时候方向搞反。3. RLHF-PPO把人类偏好塞进loop里3.1 RLHF的三段式流程RLHF不是一步到位的它分三个阶段阶段目标产出SFT用人类示范数据微调基座模型会听话的初始策略Reward Model用人类偏好对比数据训练打分模型能打分的奖励函数PPO用RM的分数作为reward训练策略对齐后的最终模型很多人直接跳到第三步结果策略模型连基本指令都跟不好PPO训练出来的东西全是胡言乱语。SFT阶段不能省它是给PPO提供一个合理的起点否则PPO要从随机初始化开始探索根本训不动。Reward Model的训练数据是成对的给定同一个prompt有两个回答A和B人类标注哪个更好。RM的loss是def rm_loss(chosen_rewards, rejected_rewards): # chosen_rewards和rejected_rewards是RM对两个回答的打分 return -torch.log(torch.sigmoid(chosen_rewards - rejected_rewards)).mean()这个loss的含义是让chosen的分数比rejected高差值越大loss越小。训练好的RM就是PPO阶段的reward函数。3.2 PPO在RLHF里的四个模型RLHF-PPO最让人头晕的是它同时涉及四个模型Actor当前正在训练的策略模型Critic估计状态价值的模型用来算advantageReward Model打分的参数冻结Reference ModelSFT阶段的模型副本参数冻结用来算KL惩罚为什么要Reference Model因为PPO训练的时候策略会拼命讨好RMRM本身是有漏洞的策略会找到RM打高分但实际很烂的回答。KL惩罚就是限制策略不要偏离SFT模型太远def compute_kl_penalty(log_probs, ref_log_probs): # 近似KL散度 kl log_probs - ref_log_probs return kl最终的reward是total_reward rm_score - kl_coef * kl_penaltykl_coef一般取0.01到0.1太小了策略会跑偏太大了策略学不动。我一般从0.02开始试。3.3 完整RLHF-PPO的loop结构把上面的东西串起来RLHF-PPO的loop是这样的def rlhf_ppo_train(prompts, actor, critic, reward_model, ref_model, optimizer_actor, optimizer_critic, config): for iteration in range(config.num_iterations): # 1. 采样用actor生成回答 with torch.no_grad(): responses, old_log_probs generate_responses(actor, prompts) # 2. 打分RM给rewardref算KL with torch.no_grad(): rm_scores reward_model(prompts, responses) ref_log_probs ref_model.compute_log_probs(prompts, responses) # 3. 算reward和advantage kl_penalty old_log_probs - ref_log_probs rewards rm_scores - config.kl_coef * kl_penalty with torch.no_grad(): values critic(prompts, responses) advantages compute_gae(rewards, values, config.gamma, config.lam) # 4. 多轮更新 for epoch in range(config.ppo_epochs): new_log_probs actor.compute_log_probs(prompts, responses) new_values critic(prompts, responses) # Actor loss actor_loss ppo_loss(new_log_probs, old_log_probs, advantages, config.clip_epsilon) # Critic loss critic_loss nn.MSELoss()(new_values, rewards) # 更新 optimizer_actor.zero_grad() actor_loss.backward() optimizer_actor.step() optimizer_critic.zero_grad() critic_loss.backward() optimizer_critic.step()这个结构看起来清晰但实际跑起来问题一堆。下面我逐个说。4. 实操中真正会卡住你的地方4.1 采样阶段的显存爆炸生成回答的时候如果batch size开太大KV cache会直接把显存吃满。我一开始用batch_size64跑7B模型A100 80G直接OOM。解决办法有两个一是用gradient checkpointing二是把生成和训练分开做生成的时候用更小的batch。生成阶段还有一个坑采样温度。温度太高生成的东西乱七八糟RM打分很低advantage全是负的训练信号很弱。温度太低生成的东西千篇一律策略学不到多样性。我一般用0.7到1.0之间的温度具体看任务。4.2 Advantage估计的坑GAEGeneralized Advantage Estimation是PPO里算advantage的标准方法但它在RLHF里有个特殊问题序列长度不一致。不同回答的长度不一样value的bootstrap怎么处理我的做法是在每个序列的最后一个token处把advantage截断不跨序列传播。具体实现的时候用一个mask把padding位置和序列结束后的位置mask掉def compute_gae(rewards, values, gamma0.99, lam0.95, maskNone): advantages torch.zeros_like(rewards) last_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] last_gae delta gamma * lam * last_gae advantages[t] last_gae if mask is not None and mask[t] 0: last_gae 0 returns advantages values return advantages, returns注意mask的处理非常关键如果不maskpadding位置的value会污染advantage的计算训练会莫名其妙地不稳定。4.3 KL惩罚的方向问题KL惩罚的符号很容易搞反。我们要的是策略不要偏离ref太远所以惩罚应该加在reward上让偏离大的回答reward变低。但KL本身是log_probs - ref_log_probs这个值可正可负直接减的话方向不一定对。正确的做法是用KL的估计量保证非负def compute_kl_penalty(log_probs, ref_log_probs): # 用k3估计量保证非负 log_ratio ref_log_probs - log_probs kl torch.exp(log_ratio) - log_ratio - 1 return kl这个估计量在log_ratio接近0的时候近似等于0.5 * log_ratio^2永远非负。我一开始用简单的差值结果策略在某些样本上反而被鼓励偏离ref训练直接崩了。4.4 常见问题速查表现象可能原因排查方向loss突然变成nan学习率太大或KL系数太小降lr到1e-6kl_coef加到0.05reward一直不涨RM打分有问题或advantage全负检查RM输出分布打印advantage统计生成结果重复温度太低或KL惩罚太强温度调到0.8kl_coef降到0.01显存OOMbatch太大或KV cache没释放减小batch生成后手动清cache训练速度极慢四个模型串行跑用FSDP或DeepSpeed做模型并行5. 代码实现里的关键细节5.1 模型初始化和参数冻结RLHF-PPO里四个模型的初始化顺序很重要。Actor从SFT模型加载Critic从SFT模型加载或者从RM加载Reward Model和Reference Model直接加载冻结def init_models(config): # Actor从SFT加载 actor AutoModelForCausalLM.from_pretrained(config.sft_path) # Critic从SFT加载但输出维度改成1 critic AutoModelForSequenceClassification.from_pretrained( config.sft_path, num_labels1 ) # RM和Ref冻结 reward_model AutoModelForSequenceClassification.from_pretrained( config.rm_path, num_labels1 ) ref_model AutoModelForCausalLM.from_pretrained(config.sft_path) for p in reward_model.parameters(): p.requires_grad False for p in ref_model.parameters(): p.requires_grad False return actor, critic, reward_model, ref_modelCritic的输出维度必须是1因为它估计的是标量value。很多人直接用CausalLM做critic输出维度是vocab_size然后取第一个token的logit这样也能跑但效率低。5.2 生成阶段的log_prob计算PPO需要old_log_probs这个是在生成的时候算的。但生成的时候用的是采样log_prob需要重新算一遍def generate_with_log_probs(model, prompts, max_new_tokens256, temperature0.8): with torch.no_grad(): outputs model.generate( prompts, max_new_tokensmax_new_tokens, temperaturetemperature, do_sampleTrue, return_dict_in_generateTrue, output_scoresTrue ) sequences outputs.sequences # 重新算log_probs log_probs compute_log_probs(model, sequences, prompts.shape[1]) return sequences, log_probs这里有个效率问题generate的时候已经算了一遍forward重新算log_probs又算一遍浪费了一倍计算。优化方法是直接用output_scores里的logits算log_prob但要注意scores里存的是采样后的logits需要做log_softmax。5.3 训练时的梯度累积RLHF-PPO的batch通常很大因为要采足够多的数据但显存有限必须用梯度累积def train_step(actor, critic, batch, config, optimizer_actor, optimizer_critic): micro_batch_size config.micro_batch_size num_micro_batches len(batch) // micro_batch_size for i in range(num_micro_batches): micro_batch batch[i * micro_batch_size:(i 1) * micro_batch_size] actor_loss compute_actor_loss(actor, micro_batch, config) critic_loss compute_critic_loss(critic, micro_batch, config) # 累积梯度 (actor_loss / num_micro_batches).backward() (critic_loss / num_micro_batches).backward() # 统一更新 torch.nn.utils.clip_grad_norm_(actor.parameters(), config.max_grad_norm) torch.nn.utils.clip_grad_norm_(critic.parameters(), config.max_grad_norm) optimizer_actor.step() optimizer_critic.step() optimizer_actor.zero_grad() optimizer_critic.zero_grad()梯度裁剪不能省RLHF的梯度方差很大不裁剪的话偶尔会出现梯度爆炸。6. 一些实战经验和避坑建议6.1 从小模型开始验证流程我强烈建议先用小模型比如GPT-2或者Qwen-0.5B把整个流程跑通确认loss能降、reward能涨再换大模型。大模型跑一次要几个小时流程有问题的话调试成本太高。小模型上跑通了换大模型只需要改配置。6.2 监控指标不能只看lossRLHF训练要看的不只是loss更重要的是平均reward应该缓慢上升如果一直平或者下降说明有问题KL散度应该稳定在一个合理范围突然飙升说明策略跑偏了回答长度如果长度突然变化很大说明策略在钻RM的空子RM打分分布如果所有回答的分数都差不多说明RM区分度不够我一般每50步打印一次这些指标画成曲线看趋势。6.3 关于超参数的几点体会PPO的超参数很多但真正重要的就几个参数推荐值说明learning_rate1e-6 ~ 5e-6比SFT小一个量级clip_epsilon0.1 ~ 0.2训练不稳定就降到0.1kl_coef0.01 ~ 0.05从0.02开始试ppo_epochs2 ~ 4太多会过拟合当前batchbatch_size64 ~ 256越大越稳但越慢gamma0.99一般不用改lam0.95GAE的标准值学习率是最关键的RLHF的lr一定要比SFT小因为策略已经比较好了大lr会直接破坏掉。我见过有人用1e-4跑PPO结果策略直接崩成随机输出。6.4 一个容易被忽略的细节padding的处理RLHF里prompt和response的长度都不一样padding是必须的。但padding会引入两个问题一是attention mask要正确设置二是loss计算要mask掉padding位置。def compute_actor_loss(actor, batch, config): input_ids batch[input_ids] attention_mask batch[attention_mask] labels batch[labels] # 只对response部分算loss outputs actor(input_ids, attention_maskattention_mask) logits outputs.logits # 只对response部分算log_probs log_probs compute_log_probs_from_logits(logits, labels, attention_mask) # PPO loss ratio torch.exp(log_probs - batch[old_log_probs]) surr1 ratio * batch[advantages] surr2 torch.clamp(ratio, 1 - config.clip_epsilon, 1 config.clip_epsilon) * batch[advantages] loss -torch.min(surr1, surr2) # mask掉padding loss (loss * batch[loss_mask]).sum() / batch[loss_mask].sum() return lossloss_mask要精确到token级别prompt部分和padding部分都mask掉只对response的有效token算loss。这个细节不做的话训练信号会被稀释收敛速度慢很多。6.5 关于分布式训练单卡跑7B模型的RLHF基本不现实四个模型加起来显存需求太大。实际生产环境一般用FSDP或者DeepSpeed ZeRO-3做参数分片。但分布式会引入新的问题模型同步、梯度all-reduce、生成时的负载均衡。我的建议是先用单卡小模型把逻辑跑通然后上FSDP的时候重点检查三件事一是actor和critic的optimizer state有没有正确分片二是reward model和ref model的forward有没有用no_grad三是生成阶段的KV cache有没有正确释放。分布式这块坑太深一篇文章讲不完后面可以单独开一篇讲FSDP下的RLHF实现。6.6 一个实用的调试技巧如果训练不收敛先别急着调超参数按这个顺序排查用固定的一批数据手动算一遍loss看数值对不对把kl_coef设成0看纯RM reward能不能涨把clip_epsilon设成很大的值比如10相当于去掉clip看是不是clip的问题把actor的学习率降到1e-7看是不是lr太大检查RM的打分是否合理拿几个样本人工看一下这个顺序能覆盖90%的常见问题。我调试的时候基本就是按这个流程走一般半小时内能定位到问题。7. 从loop到RLHF一些延伸思考把基础loop和RLHF-PPO拆开看之后会发现RLHF本质上就是在loop的reward计算环节做文章。基础RL里reward是环境给的RLHF里reward是RM给的再加上KL约束。理解了这一点很多变体就很好理解了。比如DPODirect Preference Optimization本质上就是跳过RM和PPO直接用偏好数据优化策略把RL问题转化成了监督学习问题。GRPOGroup Relative Policy Optimization则是去掉了critic用组内相对分数代替advantage估计。这些方法都是在loop的不同环节做简化或替换。我自己在实际项目里的体会是先把基础loop跑稳再考虑上RLHF。很多人一上来就搞RLHF结果连最基本的policy gradient都没跑通遇到问题根本不知道是loop的问题还是RLHF的问题。基础打牢了后面加什么都是在这个骨架上挂东西。代码这块我建议自己从头写一遍不要直接抄现成的框架。抄的时候觉得都懂自己写的时候才会发现每个细节都有坑。我第一遍写RLHF-PPO的时候光是KL惩罚的符号就搞反了两次advantage的mask处理错了三次这些都是自己写才能暴露出来的问题。