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_baseline与tir_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_BASELINE | optimal_token_baseline | 单轮 / 通用场景 |
TIR_OPTIMAL_TOKEN_BASELINE | tir_optimal_token_baseline | TIR(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_cumulative与token_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)这段代码同时说明了三个关键事实:
calculate_sum_pi_squared=True是硬性前置条件——不开这个开关,训练会在优势计算处直接断言失败;- OTB 使用
old_log_probs(采样时的策略概率)来构造路径方差代理w_t = 1 - 2·exp(old_log_probs) + sum_pi_squared; - 当启用 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=True与use_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_PATH | Qwen/Qwen3-8B | 基座模型(HuggingFace Hub 路径或本地路径) |
NNODES/NGPUS_PER_NODE | 1/8 | 节点数与每节点 GPU 数 |
ADV_ESTIMATOR | optimal_token_baseline | 优势估计器;TIR 场景设为tir_optimal_token_baseline |
TRAIN_BATCH_SIZE | 128 | 训练 batch 大小 |
PPO_MINI_BATCH_SIZE | 128 | PPO mini-batch 大小 |
MAX_PROMPT_LENGTH | 1024 | 最大 prompt 长度 |
MAX_RESPONSE_LENGTH | 2048 | 最大响应长度 |
PPO_MAX_TOKEN_LEN_PER_GPU | 24576 | 动态 bsz 下每 GPU 最大 token 数 |
ACTOR_LR | 1e-6 | actor 学习率 |
ENTROPY_COEFF | 0 | 熵正则系数 |
ROLLOUT_TP | 2 | vLLM 张量并行度 |
ROLLOUT_GPU_MEM_UTIL | 0.75 | vLLM GPU 显存利用率 |
ROLLOUT_N | 8 | 每个 prompt 采样的轨迹数(组大小) |
TOTAL_EPOCHS/SAVE_FREQ/TEST_FREQ | 15/20/5 | 训练总轮数、保存频率、评测频率 |
PROJECT_NAME/EXPERIMENT_NAME | verl_otb_gsm8k_math/qwen3_8b_vllm_fsdp | wandb/日志项目与实验名 |
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=True:OTB 的必需开关。如前文所述,缺失它会在 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_estimator | optimal_token_baseline | TIR 多轮场景用tir_optimal_token_baseline |
actor_rollout_ref.actor.calculate_sum_pi_squared | True | 必需,否则训练断言失败;勿与use_fused_kernels=True同开 |
actor_rollout_ref.rollout.n | >= 2(示例 8) | 组大小;组内只有 1 条轨迹时 OTB 退化为无基线 |
algorithm.use_kl_in_reward | False(示例) | 示例采用组内优势路线 |
actor_rollout_ref.actor.use_dynamic_bsz+ppo_max_token_len_per_gpu | True/24576 | 动态 bsz 控制显存 |
六、实践要点与注意事项
- 两个开关必须同时就位:
algorithm.adv_estimator=optimal_token_baseline与actor_rollout_ref.actor.calculate_sum_pi_squared=True。前者决定调用 OTB 实现,后者保证 batch 中携带sum_pi_squared张量;缺任一都会中断训练。 - 理解
handle_zero_tail的边界处理:默认开启时,组内最长轨迹的“无人可比”尾巴基线被置 0,该段优势等于原始收益,避免用单一轨迹的统计量污染基线。 - rollout IS 兼容:OTB 通过
rollout_is_weights以ρ̄²(t)缩放W_t,可平滑接入 algorithm.py 中的 Rollout Correction 配置,缓解 rollout 与训练策略不一致带来的偏差。 - 数据准备与运行环境:脚本默认数据位于
$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),仅供参考