Skip to content

GSM8K PPO:用一道数学题串完整流程

GSM8K 是学习 PPO/RLVR 的好入口:prompt 是小学数学题,答案可以规则验证,reward 逻辑比开放聊天更透明。读这一页时,请把它当作一条样本从 parquet 走到 PPO loss 的路线图。

这页原来哪里不够适合新手

原来的版本已经讲了“数据、rollout、reward、advantage、update”,但还缺三样东西:

  • 缺字段级源码:promptdata_sourcereward_model.ground_truth 到底在哪写入、在哪读取。
  • 缺训练主循环锚点:reward、old logprob、ref logprob、value、advantage、actor update 的顺序需要对应到 RayPPOTrainer.fit()
  • 缺公式和指标:PPO/GAE 的最小公式、critic/scorecritic/rewards 的区别、rollout.n 对 batch 的放大都要明确。

数据预处理:parquet 里必须有什么

入口脚本是 examples/data_preprocess/gsm8k.py。它做了 4 件事:

  1. openai/gsm8kquestionanswer
  2. 在问题后加一句 instruction:要求最后答案放在 #### 后。
  3. 用正则从标准答案中抽取 #### 42 这种 final answer。
  4. 写出 verl 训练需要的字段。

关键字段长这样:

python
{
    "data_source": "openai/gsm8k",
    "prompt": [{"role": "user", "content": question}],
    "ability": "math",
    "reward_model": {"style": "rule", "ground_truth": solution},
    "extra_info": {"split": split, "index": idx, "answer": answer_raw, "question": question_raw},
}

新手最容易漏的是 reward_model["ground_truth"]。reward function 不是直接知道标准答案,而是 reward manager 从 non_tensor_batch["reward_model"]["ground_truth"] 里取。

MATH 数据在 examples/data_preprocess/math_dataset.py,字段结构相同,但 data_sourceDigitalLearningGmbH/MATH-lighteval,答案格式是 \boxed{}。PPO Qwen3 示例默认同时训练 GSM8K + MATH,所以 reward scorer 会按 data_source 分流。

启动脚本怎么读

先准备数据:

bash
python3 examples/data_preprocess/gsm8k.py --local_save_dir ~/data/gsm8k
python3 examples/data_preprocess/math_dataset.py --local_save_dir ~/data/math

再看 PPO 示例:

bash
cd examples/ppo_trainer
MODEL_PATH=Qwen/Qwen3-8B \
TRAIN_BATCH_SIZE=1024 PPO_MINI_BATCH_SIZE=256 ROLLOUT_N=1 \
bash run_qwen3_8b_fsdp.sh

不要急着跑。先把脚本当作配置索引:

配置在脚本里对应含义
algorithm.adv_estimator=gaeDATA 数组用 GAE,需要 critic
data.train_filesGSM8K + MATH train parquetprompt 来源
data.train_batch_sizeTRAIN_BATCH_SIZE每步取多少个 prompt
actor_rollout_ref.rollout.nROLLOUT_N 默认 1每个 prompt 采样几条 response
actor_rollout_ref.actor.ppo_mini_batch_sizePPO_MINI_BATCH_SIZEactor 更新 mini-batch,全局数
critic.model.pathCRITIC_MODEL_PATHvalue model,PPO/GAE 需要
actor_rollout_ref.rollout.nameINFER_BACKENDvLLM/SGLang/TensorRT-LLM 等 rollout 后端

Megatron 版本 examples/ppo_trainer/run_qwen3_8b_megatron.sh 算法字段基本一样,只是 actor/critic/ref 多了 megatron.tensor_model_parallel_sizepipeline_model_parallel_size,并设置 model_engine=megatron。新手先读 FSDP 版更轻。

一条样本在 RayPPOTrainer.fit() 里怎么走

源码主线在 verl/trainer/ppo/ray_trainer.py::fit()。可以按下面 10 步读:

  1. dataloader 取一个 prompt batch,并给每个原始 prompt 写 uid
  2. gen_batch.repeat(repeat_times=rollout.n, interleave=True) 把 prompt 展开成多条待采样轨迹。
  3. async_rollout_manager.generate_sequences() 调用 rollout backend 生成 response。
  4. 原 batch 也按 rollout.n repeat,然后和生成结果 union
  5. compute_response_mask()attention_mask 切出 response token mask。
  6. reward manager 计算 rm_scoresextract_reward() 得到 reward_tensor
  7. _compute_old_log_prob() 重算 actor 旧策略 logprob,并顺便统计 actor/entropy
  8. 如果启用 reference policy,_compute_ref_log_prob() 得到 ref_log_prob
  9. PPO/GAE 需要 critic,_compute_values() 得到 values
  10. compute_advantage() 调用 compute_gae_advantage_return(),之后 _update_critic()_update_actor()

这就是你看日志里 timing_s/gentiming_s/rewardtiming_s/old_log_probtiming_s/reftiming_s/valuestiming_s/advtiming_s/update_actor 的来源。

reward:GSM8K 如何判分

默认 scorer 分派在 verl/utils/reward_score/__init__.py::default_compute_score()

text
data_source == "openai/gsm8k" -> verl/utils/reward_score/gsm8k.py

gsm8k.py 的关键点:

  • 只在 response 最后 300 个字符里找答案,避免超长正则太慢。
  • strict 模式要求匹配最后一个 #### number
  • 答案等于 ground truth 返回 score,默认 1.0。
  • 没有 #### 返回 0。
  • 格式正确但答案错误返回 format_score,当前源码默认 0.0。

注意本地 docs/preparation/reward_function.rst 仍说“格式正确错答给 0.1”,但当前源码默认是 0.0。学习和调试时以 verl/utils/reward_score/gsm8k.py 为准。

reward manager 在 verl/workers/reward_manager/naive.py 中把最终 reward 写到 response 最后一个有效 token:

text
reward_tensor[i, valid_response_length - 1] = reward

这就是为什么后面的 token_level_scores.sum(-1) 可以得到每条 response 的 outcome reward。

PPO/GAE 公式怎样落到源码

PPO 的核心是:

text
ratio_t = exp(logprob_new_t - logprob_old_t)
loss_t = -min(ratio_t * A_t, clip(ratio_t, 1-eps, 1+eps) * A_t)

对应 verl/trainer/ppo/core_algos.py::compute_policy_loss_vanilla()

GAE 的核心是:

text
delta_t = reward_t + gamma * V_{t+1} - V_t
A_t = delta_t + gamma * lam * A_{t+1}
return_t = A_t + V_t

对应 compute_gae_advantage_return()。因为 GSM8K 是 outcome reward,大多数 token 的 reward 是 0,最后有效 token 有最终分数;GAE 会借助 value 把这个最终结果传播回前面的 response token。

指标怎么对应这条链路

指标来源先看什么
critic/score/*metric_utils.compute_data_metrics()token_level_scores.sum(-1) 统计rule reward 原始分
critic/rewards/*token_level_rewards.sum(-1)如果 use_kl_in_reward=True,这里已经扣 KL
critic/advantages/*compute_gae_advantage_return() 输出是否有可学习信号
critic/vf_explained_varreturns vs valuescritic 是否学到 baseline
actor/pg_clipfraccompute_policy_loss_vanilla()PPO 更新是否经常撞 clip
actor/ppo_klcompute_policy_loss_vanilla()新旧策略差异
actor/entropy_compute_old_log_prob() 后聚合 entropy生成分布是否过早变窄
response_length/clip_ratiometric_utils.compute_data_metrics()是否大量回答打满 max_response_length

最适合新手的读源码顺序

  1. examples/data_preprocess/gsm8k.py:字段从哪来。
  2. examples/ppo_trainer/run_qwen3_8b_fsdp.sh:配置怎样覆盖。
  3. verl/utils/reward_score/gsm8k.py:答案怎么判。
  4. verl/workers/reward_manager/naive.py:reward 怎么写到最后一个 token。
  5. verl/trainer/ppo/ray_trainer.py::fit():训练流水线。
  6. verl/trainer/ppo/core_algos.py::compute_gae_advantage_return()compute_policy_loss_vanilla():公式落地。
  7. verl/trainer/ppo/metric_utils.py:日志指标怎么来的。

本节参考与延伸阅读

  • examples/ppo_trainer/README.md
  • examples/ppo_trainer/run_qwen3_8b_fsdp.sh
  • examples/ppo_trainer/run_qwen3_8b_megatron.sh
  • examples/data_preprocess/gsm8k.py
  • examples/data_preprocess/math_dataset.py
  • verl/utils/reward_score/gsm8k.py
  • verl/workers/reward_manager/naive.py
  • verl/trainer/ppo/ray_trainer.py
  • verl/trainer/ppo/core_algos.py
  • verl/trainer/ppo/metric_utils.py
  • docs/examples/gsm8k_example.rst
  • docs/start/quickstart.rst
  • docs/preparation/prepare_data.rst
  • PPO paper
  • GAE paper

面向源码阅读的 verl 学习文档。