news 2026/9/13 2:27:56

verl 中的 Optimal Token Baseline(OTB):基于 Token 级路径方差的最优基线优势估计实践指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
verl 中的 Optimal Token Baseline(OTB):基于 Token 级路径方差的最优基线优势估计实践指南

verl 中的 Optimal Token Baseline(OTB):基于 Token 级路径方差的最优基线优势估计实践指南

【免费下载链接】verlverl/HybridFlow: A Flexible and Efficient RL Post-Training Framework项目地址: https://gitcode.com/GitHub_Trending/ve/verl

verl(HybridFlow)在examples/otb_trainer中提供了 Optimal Token Baseline(OTB)训练器的完整参考实现:与 GRPO 那种“每个 prompt 组一个均值基线”的做法不同,OTB 在每个时间步为同一 prompt 组内的所有采样轨迹计算一个由累计路径方差加权的最优基线,从而在 token 粒度上降低优势估计的方差。读完本文,你将掌握 OTB 的数学原理与源码实现、optimal_token_baselinetir_optimal_token_baseline两种变体的区别、calculate_sum_pi_squared=True的开关作用,以及如何基于run_qwen3_8b_fsdp.sh在 NVIDIA GPU 上用 vLLM 采样 + FSDP 训练跑通 OTB。

一、什么是 Optimal Token Baseline

1.1 动机:从“组均值基线”到“时间步最优基线”

在以 GRPO 为代表的 group-based RL 训练中,每个 prompt 会采样 N 条轨迹,基线通常取该组奖励的均值,即对整条轨迹(或整个组)使用单一基线。这种做法的局限在于:一条轨迹内部不同位置的不确定性是不同的——有的时间步模型“很有把握”(概率集中),有的时间步模型“非常随机”(概率接近均匀分布)。把所有位置同等对待,会带来不必要的方差。

OTB 的核心思想是:为每个时间步 t 单独计算一个最优基线,并且用“累计路径方差”作为该时间步在组内加权聚合时的权重。verl 的 core_algos.py 中对该理论的描述为:

  • 每个时间步的最优基线:B_t* = E[G_t × W_t] / E[W_t]
  • 其中W_t = Σ_{j=1}^t ||s_j||²是到时间步 t 为止的累计路径方差代理
  • 而每个时间步的方差代理为||s_j||² = 1 - 2π_j + Σπ²

W_t刻画了轨迹“到目前为止已经实现的能量/不确定性”,累计值越大,说明该轨迹在高方差路径上走得越远,其收益在基线估计中的权重就越高——直观上,给“走钢丝走得越久”的轨迹更大的话语权,能更准确地估计该时刻应当扣减的基线。

1.2 与 TIR 变体的关系

OTB 存在两个注册在 core_algos.py 中的变体:

枚举值字符串名适用场景
OPTIMAL_TOKEN_BASELINEoptimal_token_baseline单轮 / 通用场景
TIR_OPTIMAL_TOKEN_BASELINEtir_optimal_token_baselineTIR(Tool-Integrated Reasoning,多轮工具调用)场景

两者的数学形式完全一致(B_t* = E[G_t × W_t] / E[W_t]),差别在于处理变长轨迹与掩码的方式

  • compute_optimal_token_baseline_advantage直接在[bs, response_length]的稠密矩阵上按response_mask分组计算;
  • compute_multi_turn_optimal_token_baseline_advantage先把每个轨迹的w_cumulativetoken_returns按真实长度“压实”(unpad)到[bs*n, max_response_length]再计算基线,最后再写回原始[bs, turn * response_length]布局,因此天然适配 TIR 多轮(multi-turn)数据的拼接结构。

如果你在跑 TIR 多轮训练,用环境变量ADV_ESTIMATOR=tir_optimal_token_baseline即可切换到该变体(见 run_qwen3_8b_fsdp.sh 第 11 行)。

二、OTB 在 verl 中的注册与调用链

2.1 优势估计器注册机制

verl 用一个可扩展的注册表管理所有优势估计器(GAE、GRPO、RLOO、REMAX、OPO、GPG 等),OTB 只是其中之一。定义位于 core_algos.py:

class AdvantageEstimator(str, Enum): ... OPTIMAL_TOKEN_BASELINE = "optimal_token_baseline" TIR_OPTIMAL_TOKEN_BASELINE = "tir_optimal_token_baseline"

注册装饰器@register_adv_est(...)(core_algos.py)把两个实现函数分别绑定到这两个字符串名上:

  • compute_optimal_token_baseline_advantage(core_algos.py)
  • compute_multi_turn_optimal_token_baseline_advantage(core_algos.py)

因此配置algorithm.adv_estimator=optimal_token_baseline就能在训练时被解析并调用对应实现。

2.2 训练器中的调用点

在 PPO 训练器 ray_trainer.py 中,当检测到adv_estimator属于 OTB 两个变体之一时,会断言 batch 中必须存在sum_pi_squared

assert "sum_pi_squared" in data.batch, ( "Step-dependent optimal baseline requires sum_pi_squared from actor. " "Please set actor.calculate_sum_pi_squared=True in config." ) adv_kwargs["sum_pi_squared"] = data.batch["sum_pi_squared"] adv_kwargs["old_log_probs"] = data.batch["old_log_probs"] adv_kwargs["rollout_is_weights"] = data.batch.get("rollout_is_weights", None)

这段代码同时说明了三个关键事实:

  1. calculate_sum_pi_squared=True是硬性前置条件——不开这个开关,训练会在优势计算处直接断言失败;
  2. OTB 使用old_log_probs(采样时的策略概率)来构造路径方差代理w_t = 1 - 2·exp(old_log_probs) + sum_pi_squared
  3. 当启用 Rollout Correction(rollout_is_weights存在)时,OTB 还会把重要性采样权重平方ρ̄²(t)乘到W_t上,用于在截断 IS 下最小化 MSE。

2.3sum_pi_squared从哪里来

sum_pi_squared(记为 Σπ²,即模型在词表上的概率平方和)由 actor 前向时逐 token 计算,相关开关定义在 actor.py:

calculate_sum_pi_squared: bool = False

在 FSDP 引擎 transformer_impl.py 中,当calculate_sum_pi_squared=True时,前向会调用verl_F.calculate_sum_pi_squared_from_logits(logits)并把结果写入model_output["sum_pi_squared"]。其数值实现位于 torch_functional.py:

def calculate_sum_pi_squared_from_logits(logits: torch.Tensor): """ Formula: Σπ² = exp(logsumexp(2*logits) - 2*logsumexp(logits)) """ return torch.exp(torch.logsumexp(2.0 * logits, dim=-1) - 2.0 * torch.logsumexp(logits, dim=-1))

该式利用 log-sum-exp 的数值技巧直接给出词表概率平方和的精确值,全程不显式计算 softmax 概率,数值稳定性更好。需要特别注意的是:FSDP 实现中会检查calculate_sum_pi_squared=Trueuse_fused_kernels=True不能同时开启(transformer_impl.py),因为融合核路径不产出该统计量。

三、OTB 优势计算的核心步骤(源码级拆解)

compute_optimal_token_baseline_advantage(core_algos.py)的完整流程可以拆成五步:

Step 1:计算每个时间步的收益(reward-to-go)

returns = (token_level_rewards * response_mask).flip(dims=[-1]).cumsum(dim=-1).flip(dims=[-1])

对 token 级奖励从序列末尾向前做 cumsum,得到每个位置的累计回报G_t

Step 2:逐时间步方差代理

pi_t = torch.exp(old_log_probs) w_per_timestep = 1 - 2 * pi_t + sum_pi_squared

1 - 2π + Σπ²得到每个位置的“不确定性度量”。当策略在该位置完全确定(π 为 one-hot,Σπ²=1)时该项为 0;当策略接近均匀分布时该项接近 1,从而在数值上捕捉每个位置的熵/方差信息。

Step 3:累计路径方差

w_cumulative = (w_per_timestep * response_mask).cumsum(dim=-1)

W_t = Σ_{j=1}^t w_j,即从轨迹起点累计到当前时间步的方差代理之和。

Step 4:按 prompt 分组计算每时间步最优基线

numerator = (returns_group * w_cumulative_group * mask_group).sum(dim=0) denominator = (w_cumulative_group * mask_group).sum(dim=0) + epsilon baseline_per_step = numerator / denominator

对组内所有轨迹,按时间步做加权平均B_t* = Σ[G_t × W_t] / Σ[W_t]epsilon=1e-8防除零)。注意两个边界行为:

  • 组内只有 1 条轨迹时,不计算基线,优势直接等于收益;
  • handle_zero_tail=True(默认开启)时,会把组内最长轨迹超出第二长轨迹长度的那段“尾巴”的基线置零——因为那段只有一条轨迹参与,没有可比的组内参照。

Step 5:计算优势

advantages = (returns - baselines) * response_mask

每个 token 的优势A_t = G_t - B_t*,最终乘response_mask屏蔽 padding 位置。

TIR 变体(core_algos.py)的数学流程一致,但多了一步“按真实长度压实 → 在[bs*n, max_response_length]紧凑张量上计算 → 写回原始多轮布局”的转换(Step 4),这正是它能正确处理 TIR 多轮数据的关键。

四、官方标准脚本:run_qwen3_8b_fsdp.sh 逐段解读

verl 为 OTB 提供了官方 canonical 脚本 run_qwen3_8b_fsdp.sh,配置矩阵为文本任务 + vLLM 采样 + FSDP 训练 + NVIDIA GPU。脚本顶部给出了可调环境变量及其默认值:

环境变量默认值作用
MODEL_PATHQwen/Qwen3-8B基座模型(HuggingFace Hub 路径或本地路径)
NNODES/NGPUS_PER_NODE1/8节点数与每节点 GPU 数
ADV_ESTIMATORoptimal_token_baseline优势估计器;TIR 场景设为tir_optimal_token_baseline
TRAIN_BATCH_SIZE128训练 batch 大小
PPO_MINI_BATCH_SIZE128PPO mini-batch 大小
MAX_PROMPT_LENGTH1024最大 prompt 长度
MAX_RESPONSE_LENGTH2048最大响应长度
PPO_MAX_TOKEN_LEN_PER_GPU24576动态 bsz 下每 GPU 最大 token 数
ACTOR_LR1e-6actor 学习率
ENTROPY_COEFF0熵正则系数
ROLLOUT_TP2vLLM 张量并行度
ROLLOUT_GPU_MEM_UTIL0.75vLLM GPU 显存利用率
ROLLOUT_N8每个 prompt 采样的轨迹数(组大小)
TOTAL_EPOCHS/SAVE_FREQ/TEST_FREQ15/20/5训练总轮数、保存频率、评测频率
PROJECT_NAME/EXPERIMENT_NAMEverl_otb_gsm8k_math/qwen3_8b_vllm_fsdpwandb/日志项目与实验名

4.1 数据部分(DATA)

DATA=( algorithm.adv_estimator=${adv_estimator} algorithm.use_kl_in_reward=False data.train_files="$train_files" data.val_files="$val_files" data.train_batch_size=${train_batch_size} data.max_prompt_length=${max_prompt_length} data.max_response_length=${max_response_length} data.filter_overlong_prompts=True data.truncation='error' )

脚本默认混用 GSM8K 与 MATH 两个数据集(train_files=['$HOME/data/gsm8k/train.parquet', '$HOME/data/math/train.parquet'],测试集同理),使用前请按 prepare_data.rst 准备好 parquet 格式的数据。值得注意的是algorithm.use_kl_in_reward=False——OTB 示例走的是“KL 从损失侧约束(use_kl_loss=False时则完全不做 KL 约束)+ 组内优势估计”的路线,与经典的 in-reward KL 惩罚不同。

4.2 actor 部分:OTB 的两个关键开关

ACTOR=( actor_rollout_ref.actor.optim.lr=${actor_lr} actor_rollout_ref.actor.ppo_mini_batch_size=${ppo_mini_batch_size} actor_rollout_ref.actor.use_dynamic_bsz=True actor_rollout_ref.actor.ppo_max_token_len_per_gpu=${ppo_max_token_len_per_gpu} actor_rollout_ref.actor.use_kl_loss=False actor_rollout_ref.actor.entropy_coeff=${entropy_coeff} actor_rollout_ref.actor.calculate_sum_pi_squared=True actor_rollout_ref.actor.fsdp_config.param_offload=False actor_rollout_ref.actor.fsdp_config.optimizer_offload=False )
  • actor_rollout_ref.actor.calculate_sum_pi_squared=TrueOTB 的必需开关。如前文所述,缺失它会在 ray_trainer.py 处断言失败。同时注意不要与use_fused_kernels=True同时启用。
  • use_dynamic_bsz=True+ppo_max_token_len_per_gpu=24576:开启动态 batch 大小,以 token 数为单位控制显存。
  • entropy_coeff=0:示例默认不施加熵正则。

4.3 采样与参考策略部分(ROLLOUT / REF)

ROLLOUT=( actor_rollout_ref.rollout.name=vllm actor_rollout_ref.rollout.tensor_model_parallel_size=${rollout_tp} actor_rollout_ref.rollout.gpu_memory_utilization=${rollout_gpu_mem_util} actor_rollout_ref.rollout.n=${rollout_n} actor_rollout_ref.rollout.log_prob_use_dynamic_bsz=True actor_rollout_ref.rollout.log_prob_max_token_len_per_gpu=${ppo_max_token_len_per_gpu} ) REF=( actor_rollout_ref.ref.log_prob_use_dynamic_bsz=True actor_rollout_ref.ref.log_prob_max_token_len_per_gpu=${ppo_max_token_len_per_gpu} actor_rollout_ref.ref.fsdp_config.param_offload=True )
  • 采样由 vLLM 承担,tensor_model_parallel_size=2表示 rollout 阶段做 2 路张量并行,n=8表示每个 prompt 采 8 条轨迹(OTB 的组大小)。
  • 参考策略(ref)只用于计算 KL(尽管本示例use_kl_in_reward=False),其 FSDP 参数卸载(param_offload=True)以节省显存。

4.4 训练器与启动方式(TRAINER / LAUNCH)

TRAINER=( trainer.balance_batch=True trainer.critic_warmup=0 trainer.logger='["console","wandb"]' trainer.project_name=${project_name} trainer.experiment_name=${experiment_name} trainer.n_gpus_per_node=${NGPUS_PER_NODE} trainer.nnodes=${NNODES} trainer.save_freq=${save_freq} trainer.test_freq=${test_freq} trainer.total_epochs=${total_epochs} )

critic_warmup=0说明 OTB 不需要 value network(与 GRPO 同族,属于无 critic 的优势估计),因此也无需 critic 预热。启动部分(脚本 104-123 行)直接调用python3 -m verl.trainer.main_ppo传入上述全部参数数组;若环境变量VERL_USE_UV(默认 1)且设备为 GPU,则改用uv run --frozen --all-packages --extra vllm --extra fsdp启动 driver 与 Ray worker,保证 vLLM × FSDP 依赖组合与仓库锁定的uv.lock一致。运行时需在 verl 仓库根目录下执行:

bash examples/otb_trainer/run_qwen3_8b_fsdp.sh

五、关键配置速查

OTB 训练的最小配置面:

配置项推荐值说明
algorithm.adv_estimatoroptimal_token_baselineTIR 多轮场景用tir_optimal_token_baseline
actor_rollout_ref.actor.calculate_sum_pi_squaredTrue必需,否则训练断言失败;勿与use_fused_kernels=True同开
actor_rollout_ref.rollout.n>= 2(示例 8)组大小;组内只有 1 条轨迹时 OTB 退化为无基线
algorithm.use_kl_in_rewardFalse(示例)示例采用组内优势路线
actor_rollout_ref.actor.use_dynamic_bsz+ppo_max_token_len_per_gpuTrue/24576动态 bsz 控制显存

六、实践要点与注意事项

  1. 两个开关必须同时就位algorithm.adv_estimator=optimal_token_baselineactor_rollout_ref.actor.calculate_sum_pi_squared=True。前者决定调用 OTB 实现,后者保证 batch 中携带sum_pi_squared张量;缺任一都会中断训练。
  2. 理解handle_zero_tail的边界处理:默认开启时,组内最长轨迹的“无人可比”尾巴基线被置 0,该段优势等于原始收益,避免用单一轨迹的统计量污染基线。
  3. rollout IS 兼容:OTB 通过rollout_is_weightsρ̄²(t)缩放W_t,可平滑接入 algorithm.py 中的 Rollout Correction 配置,缓解 rollout 与训练策略不一致带来的偏差。
  4. 数据准备与运行环境:脚本默认数据位于$HOME/data/gsm8k$HOME/data/math,训练入口是verl.trainer.main_ppo,运行前请确认依赖(vLLM、FSDP extras)已按 install.rst 安装。

OTB 在 verl 中的实现路径清晰:配置开关 → actor 前向产出sum_pi_squared(torch_functional.py)→ 训练器组装参数(ray_trainer.py)→ 核心算法按组逐时间步加权求基线(core_algos.py)。无论你是想复现 OTB 训练,还是希望把它作为自定义优势估计器的起点(verl 的register_adv_est注册表支持以字符串名注册新实现),本文给出的脚本、源码链路与配置矩阵都足以支撑你直接上手。

【免费下载链接】verlverl/HybridFlow: A Flexible and Efficient RL Post-Training Framework项目地址: https://gitcode.com/GitHub_Trending/ve/verl

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/13 2:25:26

数据结构与算法面试速通指南:高频考点与备考策略

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/13 2:22:46

kohya_ss LoRA 训练完整指南:从零安装到练出第一个角色模型

kohya_ss LoRA 训练完整指南:从零安装到练出第一个角色模型 【免费下载链接】kohya_ss 项目地址: https://gitcode.com/GitHub_Trending/ko/kohya_ss kohya_ss 解决一件事:不碰命令行,也能给 Stable Diffusion 训练自己的 LoRA 和自定…

作者头像 李华
网站建设 2026/9/13 2:16:22

深入理解HOOK机制:从消息钩子到API Hook的底层原理与实战

HOOK这个词,搞Windows开发的人天天挂嘴边,做Web开发的人也经常听到,可你真要让人一句话说清楚它是什么、能干什么,能一口气讲明白的人真不多。我直接上两段能跑的代码,带你把HOOK的底层逻辑、实现方式和坑点一次性捋顺…

作者头像 李华
网站建设 2026/9/13 2:15:42

职场短剧AI配音选型指南:小云雀与OiiOii核心差异解析

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华