ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

torchtune.rlhf 模块深度解析:PPO 与 DPO 核心 API 实现及训练实战

2026/9/17 19:43:43 拓冰建站 浏览量
torchtune.rlhf 模块深度解析:PPO 与 DPO 核心 API 实现及训练实战 torchtune.rlhf 模块深度解析PPO 与 DPO 核心 API 实现及训练实战【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune本文基于 torchtune 官方 API 参考文档docs/source/api_ref_rlhf.rst展开系统讲解torchtune.rlhf模块中用于 PPO 与 DPO 两类 RLHFReinforcement Learning from Human Feedback算法的组件与损失函数奖励计算get_rewards_ppo、GAE 优势估计estimate_advantages、序列截断truncate_sequence_at_first_stop_token以及三大损失PPOLoss、DPOLoss、RSOLoss。读完后你将理解每个 API 的输入输出张量形状与语义并能结合仓库中现成的 recipe 与配置文件如 lora_dpo_distributed.py、1B_full_ppo_low_memory_single_device.yaml把这套 API 真正用起来。一、模块总览torchtune.rlhf 导出什么API 参考文档api_ref_rlhf.rst声明该模块是Components and losses for RLHF algorithms like PPO and DPO用于 PPO、DPO 等 RLHF 算法的组件与损失函数核心条目包括导出符号类别作用estimate_advantagesPPO 工具函数基于 GAE 估计优势值与回报get_rewards_ppoPPO 工具函数计算含 KL 惩罚的 PPO 逐 token 奖励truncate_sequence_at_first_stop_token序列处理在首个 stop token 处截断并填充loss.PPOLoss损失模块PPO 裁剪策略损失 值函数损失loss.DPOLoss损失模块Direct Preference Optimization 损失loss.RSOLoss损失模块Rejection Sampling Optimizationhinge损失从源码结构看见 rlhf/init.py除上述六个条目外模块还导出了一批支撑函数与数据结构它们在实际 recipe 中同样关键序列处理sequence_processing.pylogits_to_logprobs、batched_logits_to_logprobs、get_batch_log_probs、truncate_sequence_for_logprobs掩码统计工具rewards.pymasked_mean、masked_sum、masked_var、whiten、get_reward_penalty_mask类型定义_types.pyTrajectoryPPO 轨迹张量集合、PPOStatsPPO 统计指标、ChosenRejectedOutputsDPO 的 chosen/rejected 前向输出。这套模块被 recipe 直接调用例如 DPO 分布式训练脚本中from torchtune.rlhf import ChosenRejectedOutputs并用rlhf.get_batch_log_probs(...)计算序列对数概率见 lora_dpo_distributed.py 第 36、615 行附近。二、序列处理truncate_sequence_at_first_stop_token在 PPO 生成轨迹后必须先在第一个 stop token 处截断奖励惩罚、logprob 计算都依赖这一步。该函数的 docstring 自带一个完整 doctest直观说明行为import torch from torchtune import rlhf stop_token_ids torch.tensor([2, 869]) sequences torch.tensor([ [869, 30, 869], [2, 30, 869], [869, 30, 2], [50, 30, 869], [13, 30, 2], [13, 30, 5], [13, 2, 20], [13, 2, 2], [2, 2, 2], ]) padding_mask, truncated rlhf.truncate_sequence_at_first_stop_token( sequences, stop_token_ids, fill_value0 ) # 首个 stop token 之后的位置被填 0padding_mask 标记这些位置以上示例逐行摘自 sequence_processing.py 第 29-73 行的文档示例。参数说明参数类型说明sequencestorch.Tensor形状[b, seq_len]或[seq_len]stop_tokenstorch.Tensorstop token 集合如 EOS、padfill_valueint默认0截断后填充值通常用pad_id返回值(padding_mask, sequences)其中padding_mask为 bool 张量True表示该位置已被截断填充。实现细节函数先用torch.isin(sequences, stop_tokens)找到 stop 位置再用cumsum判断是否已出现首个 stop token 之后的位置满足条件的位置被就地写为fill_valuesequence_processing.py 第 75-79 行。注意这会原地修改输入张量且序列需保证一个 stop token 之后不再出现有效 token的假设。与之配合的还有truncate_sequence_for_logprobs(query_response_logits, context_length)对(query, response)拼接后的 logits形状[b, context_length response_length, vocab_size]截取 response 部分返回[:, context_length - 1 : -1]因为预测第 t 个 response token 需要用第 t-1 个位置的 logits。三、logit → logprob 工具链RLHF 损失的核心输入是生成序列每个 token 的对数概率。模块提供三层工具3.1 logits_to_logprobs 与 batched_logits_to_logprobslogprobs rlhf.logits_to_logprobs(logits, sequences, temperature1.0)logits_to_logprobs对logits / temperature做log_softmax后按sequences逐 token gather返回形状[b, response_length]sequence_processing.py 第 82-99 行。batched_logits_to_logprobs同样的计算但按chunk_size默认 4分块处理用 float32 输出张量承接结果目的是降低大词表 logits 的显存峰值。3.2 get_batch_log_probsDPO 的关键函数logps rlhf.get_batch_log_probs( logits, # (b, s, v) 模型直接输出 labels, # (b, s) 真值 token label_pad_token_id-100, # 该 id 的 label 被忽略 return_average_logprobsFalse, )实现逻辑先将 labels 右移一位labels[:, 1:]与logits[:, :-1, :]对齐把label_pad_token_id默认取torchtune.data.CROSS_ENTROPY_IGNORE_IDX即 -100位置替换为 0 以安全 gather然后逐 token 取 logprob。两种返回模式return_average_logprobsFalse默认返回(per_token_log_probs * loss_mask).sum(-1)形状(b,)即序列总 logprobreturn_average_logprobsTrue用masked_mean返回按有效 token 数平均的 logprob可避免长序列在 DPO 中天然占据更大权重的问题DPO 原论文训练器即提供该选项。DPO recipe 正是用它完成一次前向拿到 chosen/rejected 的 logpsall_log_probs rlhf.get_batch_log_probs(all_logits, concatenated_labels)再按批次前一半/后一半拆分lora_dpo_distributed.py 第 592-625 行。四、PPO 组件奖励、优势与损失4.1 get_rewards_ppoKL 惩罚 末尾奖励total_reward, kl, kl_reward rlhf.get_rewards_ppo( scores, # (b,) 奖励模型打分 logprobs, # (b, response_len) 策略模型 logprobs ref_logprobs, # (b, response_len) 参考模型 logprobs kl_coeff, # floatKL 惩罚系数 valid_score_idxsNone, # (b,) 每条序列最后一个有效 token 的位置 )实现上rewards.py 第 49-94 行逐 token 计算kl logprobs - ref_logprobs得到kl_reward -kl_coeff * kl奖励模型打分只加在每条序列的最后一个有效位置若给定valid_score_idxs则用scatter_add_精确投放否则直接加在[:, -1]。这样 PPO 的总奖励就是逐 token 的 KL 负奖励 序列末端的奖励模型分返回三个形状均为(b, response_len)的张量。与之配套的是get_reward_penalty_mask(padding_masks, seq_lens, penalise_no_eosTrue, min_response_lengthNone)由于序列已按 stop token 截断完全没有 padding就意味着未出现 stop token可据此惩罚不生成 EOS的轨迹同时惩罚长度小于min_response_length的过短回答。这些参数在 PPO 配置中可见min_response_length: 18、penalise_no_eos: True、reward_penalty: -31B_full_ppo_low_memory_single_device.yaml 第 144-149 行。4.2 estimate_advantagesGAE 优势估计advantages, returns rlhf.estimate_advantages( values, # (b, response_len) 值函数预测 rewards, # (b, response_len) 上面得到的 total_reward gamma, # 折扣因子 lmbda, # GAE-Lambda masksNone, # (b, response_len) 有效 token 掩码 )实现是标准的逆序 GAE 循环rewards.py 第 182-238 行delta_t r_t gamma * V(s_{t1}) - V(s_t) A_t delta_t gamma * lambda * A_{t1} returns advantages values最后对优势做whiten跨 batch 白化以降低方差被 mask 掉的 padding 位置优势置零。whiten、masked_mean、masked_var等实现源自 HuggingFace TRL 的trl/core.py源码注释中标明了出处其中whiten(x, mask)在有 mask 时用masked_mean/masked_var计算再(x - mean) * rsqrt(var 1e-8)可选shift_mean加回均值。对应的 PPO 超参在配置中的示例gamma: 1、lmbda: 0.95、kl_coeff: 0.011B_full_ppo_low_memory_single_device.yaml 第 158-168 行。4.3 PPOLoss裁剪策略损失 值函数损失loss_fn torchtune.rlhf.loss.PPOLoss( epsilon0.1, # 策略比裁剪范围 value_clip_range0.2, # 值函数裁剪范围 value_coeff0.1, # 值损失系数 ) loss, policy_loss, value_loss, ratios, clipfrac loss_fn( pi_old_logprobs, # 旧策略 logprobs pi_logprobs, # 当前策略 logprobs advantages, phi_old_values, # 旧值函数预测 phi_values, # 当前值函数预测 returns, padding_masksNone, # 策略损失的有效 token 掩码 value_padding_masksNone, # 值损失的有效 token 掩码 )实现要点loss/ppo.py比率与裁剪ratios exp(pi_logprobs - pi_old_logprobs)夹在[1-epsilon, 1epsilon]策略损失取max(-A * clipped_ratio, -A * ratio)即 pessimistic悲观目标clipfrac记录被裁剪样本的占比是监控 PPO 更新幅度的关键指标值损失对phi_values相对phi_old_values做value_clip_range裁剪取max((V-returns)^2, (V_clip-returns)^2)再乘 0.5提供padding_masks时用masked_mean只统计有效 token避免 padding 稀释损失返回的policy_loss、value_loss、ratios.mean()、clipfrac均已detach()供 recipe 记录指标。PPOStats_types.py则以 NamedTuple 形式收纳loss / policy_loss / value_loss / ratios / clipfrac / approx_policy_kls六项指标。注意源码 docstring 中说明的返回元组包含 5 个元素loss、policy_loss、value_loss、ratios、clipfrac而PPOStats额外含approx_policy_kls由 recipe 层面补充估算。PPO 的轨迹数据统一封装在Trajectory中其字段与形状在 _types.py 中有完整注释query_responses (b, contextmax_gen)、logprobs/ref_logprobs/values (b, max_gen)、scores (b,)、masks、position_ids、response_padding_masks、value_padding_masks、value_seq_idxs、seq_lens。4.4 PPO 实战配置仓库提供了单卡 PPO 全参微调 recipe ppo_full_finetune_single_device.py 与配套配置TinyLlama 1B1B_full_ppo_low_memory_single_device.yaml。配置结构上同时维护四个 checkpointerpolicy、ref_policy、value、reward——参考策略与奖励模型在训练中权重冻结价值模型通常用已训练好的奖励模型权重初始化配置注释明确说明了这一设计。启动命令tune run ppo_full_finetune_single_device --config llama2/1B_full_ppo_low_memory_single_device其中与torchtune.rlhf直接相关的超参数段# GAE 超参 gamma: 1 lmbda: 0.95 # PPO 超参 loss: _component_: torchtune.rlhf.loss.PPOLoss epsilon: 0.2 value_coeff: 0.1 value_clip_range: 0.2 kl_coeff: 0.01 # 惩罚相关 min_response_length: 18 penalise_no_eos: True reward_penalty: -3 stop_token_ids: [2, 29889]五、偏好优化损失DPOLoss 与 RSOLoss5.1 公共基类与数据结构DPOLoss和RSOLoss都继承自PreferenceLossloss/dpo.py其约定is_reference_free属性基类返回False表示损失需要参考模型 logprobs若某损失可无参考模型运行recipe 会据此跳过参考前向输入类型为ChosenRejectedOutputs_types.pychosen_logps (b,)、rejected_logps (b,)、chosen_logits、rejected_logits输出统一为三元组(losses, chosen_rewards, rejected_rewards)后两者用于计算 reward accuracy 等监控指标。5.2 DPOLossdpo torchtune.rlhf.loss.DPOLoss(beta0.1, label_smoothing0.0) losses, chosen_rewards, rejected_rewards dpo(policy_inputs, reference_inputs)核心计算loss/dpo.py 第 79-120 行logits (π_chosen - π_rejected) - (ref_chosen - ref_rejected) loss -logsigmoid(β * logits) * (1 - label_smoothing) - logsigmoid(-β * logits) * label_smoothingbeta温度参数文档建议0.1~0.5默认 0.1beta - 0时参考模型被忽略label_smoothing编码标注不确定性计算保守的 DPO 损失默认 0返回的chosen_rewards / rejected_rewards定义为β * (π_logps - ref_logps)detachrecipe 用它统计rewards/accuracies (chosen rejected).float().mean()。实战接入DPO recipe 通过一次chosenrejected 拼接前向拿到两组输出concatenated_forward再用disable_adapter上下文禁用 LoRA 适配器对同一 batch 做一次 no-grad 前向得到参考输出lora_dpo_distributed.py 第 695-704 行——即参考模型并非独立副本而是关掉 LoRA 后的底座模型这是 LoRA DPO 省显存的关键。配置文件示例8B_lora_dpo.yamlloss: _component_: torchtune.rlhf.loss.DPOLoss beta: 0.1 label_smoothing: 0启动方式tune run --nnodes 1 --nproc_per_node 2 lora_dpo_distributed --config llama3_1/8B_lora_dpo该 recipe 支持 FSDP、激活检查点/卸载、bf16、梯度累积总 batch batch_size * GPU 数 * gradient_accumulation_steps指标日志包含rewards/chosen、rewards/rejected、rewards/accuracies、log_probs/*、logits/*等lora_dpo_distributed.py 第 776-800 行。5.3 RSOLossrso torchtune.rlhf.loss.RSOLoss(gamma0.1)RSORejection Sampling Optimizationhinge 损失与 DPO 共享 logratio 差计算但损失改为losses relu(1 - gamma * logits)直觉上DPO 相当于对偏好数据做 logistic 回归而 RSO 是其 SVM 式hinge 损失的对应物源码 docstring 引自论文描述。注意该类被deprecated装饰器标记即将在未来版本中弃用新项目建议直接使用DPOLoss。六、测试与延伸阅读工具函数测试test_rewards.py、test_sequence_processing.py损失函数测试tests/torchtune/rlhf/loss/目录覆盖 DPO/PPO 损失数值行为Recipe 级测试test_lora_dpo_distributed.py、test_ppo_full_finetune_single_device.py验证了上文配置与 API 的组合可用性。小结torchtune.rlhf用一组形状明确的张量工具函数奖励、优势、截断、logprob 三个可config.instantiate的损失模块把 PPO 与 DPO 两条 RLHF 路线的数学公式落成了可直接嵌入 recipe 的 PyTorch 组件。PPO 路线关注get_rewards_ppo → estimate_advantages → PPOLoss的调用链与gamma/lmbda/kl_coeff/epsilon等超参DPO 路线只需get_batch_log_probs → DPOLoss(beta...)两步并可用禁用 LoRA 适配器充当参考模型的技巧大幅降低显存。所有 API 的精确签名与张量形状以 torchtune/rlhf 源码注释为准。【免费下载链接】torchtunePyTorch native post-training library项目地址: https://gitcode.com/GitHub_Trending/to/torchtune创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考