GRPO 逐行代码 walkthrough · 干货攻略

  • 链接: https://x.com/rasbt/status/2106374448060248318
  • 分类: x-tips
  • 来源: X @rasbt
  • 作者: Jay
  • 更新: 2026-10-05
  • 仓库: rasbt/reasoning-from-scratch

这是什么

这是一条关于 GRPO(Group Relative Policy Optimization)算法逐行代码实现的实战分享。

2026 年 10 月 3 日,Daniel Priscu(@daniel_priscu)在 X 上发了一条感慨:

"grpo walked through with actual code is rare, most explainers stop at the diagram"

Sebastian Raschka(@rasbt)随后回复确认:

"@daniel_priscu And it's used to train a pretty solid model too"

这里的"it"指的是 Raschka 在《Build a Reasoning Model (From Scratch)》第 6 章中实现的 GRPO 代码——完整可运行的 Python 脚本,逐行实现了 advantage 计算、reward 评估、logprob 差分和 GRPO loss,已经用这个代码在 Qwen3-0.6B Base 上把 MATH-500 准确率从 15.2% 提升到 47.4%(与同尺寸官方 Qwen3 reasoning 模型 48.2% 基本持平)。

配套脚本还包含多 GPU(FSDP)版本和 batch 训练版本,可直接跑 MATH 数据集。


为什么值得关注

谁在推——@rasbt

Sebastian Raschka 是前威斯康星大学统计学教授、《Build a Large Language Model (From Scratch)》和《Build a Reasoning Model (From Scratch)》作者,在 ML/AI 领域有超过 50 万关注者。他的分享以代码级实操 + 实验数据著称,每条结论都有实验支撑,不发空泛结论。

解决什么问题

GRPO 相关的代码 walkthrough 极度稀缺。 目前大多数 GRPO 资料只有算法框图,或者直接调用 HuggingFace TRL 库的黑盒脚本。Raschka 的第 6 章做了三件事:

  1. 逐行实现 GRPO 四大核心步骤:advantage 计算 → reward 获取 → sequence logprob → GRPO loss
  2. 直接在 Qwen3-0.6B Base 上训练,用 12k MATH 训练集跑出可量化的 benchmark 提升
  3. 配套多 GPU 脚本,不只是 notebook 示意图,而是生产级别的训练脚本

核心数值锚点(官方来源)

模型 MATH-500 准确率 平均输出 tokens
Qwen3-0.6B Base(训练前) 15.2% 78.85
Qwen3-0.6B GRPO 训练后(50 步 / 512 tokens / 8 rollouts) 47.4% 586.11
官方 Qwen3 reasoning 模型(同尺寸) 48.2% 1369.79

结论:GRPO 训练后的 0.6B 模型与官方 reasoning 模型准确率基本持平(47.4% vs 48.2%),同时输出 token 缩短了 57%(586 vs 1369),推理效率大幅提升。


核验过程

官方来源

来源 读取内容 核验结论
rasbt/reasoning-from-scratch GitHub 主仓库 README 全书章节结构、ch06 代码位置、配套 bonus 材料 ✅ 确认 ch06 GRPO 代码存在
rasbt/reasoning-from-scratch ch06/02_rlvr_grpo_scripts_intro README GRPO 基准对比表(10 种方法)、运行命令、显存需求表 ✅ 读取完整
rasbt/reasoning-from-scratch ch06/02_rlvr_grpo_scripts_intro/rlvr_grpo_original_no_kl.py(raw 文件) 逐行 GRPO 实现:sample_response / sequence_logprob / compute_grpo_loss / 训练循环 ✅ 读取完整(495 行)
Raschka LinkedIn 帖子(2026-01) "从 15% 到 47% 准确率,与官方 Qwen3 reasoning 模型同尺寸基本持平" ✅ 交叉确认
Raschka Substack(@rasbt) 同 LinkedIn 表述 + "12k MATH 训练集" + "多 GPU 脚本" ✅ 交叉确认
YouTube 视频:Build A Reasoning Model From Scratch 6(2026-10-03) 视频时间戳 1:22:33 确认 "15.2% → 47%"、"almost as good as this one" ✅ 视频字幕交叉确认

交叉验证关键说法

说法 来源A(官方) 来源B(验证) 结论
Qwen3-0.6B Base MATH-500 = 15.2% GitHub README 基准表 Row 1 YouTube 1:22:33 视频原话 ✅ 确认
GRPO 训练后 47.4%(50步/no-KL) GitHub README 基准表 Row 5 LinkedIn/Substack "15% → 47%" ✅ 确认(精确值 47.4%,约等于 47%)
与官方 Qwen3 reasoning 模型 48.2% 基本持平 GitHub README 基准表 Row 2 YouTube 1:22:49 "almost as good" ✅ 确认
输出 token 从 ~79 降至 ~586 GitHub README 基准表 YouTube 1:23:03 ✅ 确认
KL divergence term 移除提升性能 ch06 README 原文 DAPO/Dr. GRPO/Olmo3 paper 引用 ✅ 确认(文献引用存在)
训练超过 50 步性能反而下降 GitHub README 注释 YouTube 1:23:49 视频原话 ✅ 确认

无法核验项

  • 具体训练时长("两小时"):YouTube 视频 1:22:33 提及,但这是针对特定硬件(未披露 GPU 型号)的经验值,无法独立复现
  • 12k MATH 训练集的具体来源:GitHub README 提及 https://github.com/rasbt/math_full_minus_math500,但该仓库内容未逐一核验
  • "pretty solid model" 的定性描述:Raschka 主观评价,无法量化

上手步骤

1. 环境准备

git clone --depth 1 https://github.com/rasbt/reasoning-from-scratch.git
cd reasoning-from-scratch

# 推荐用 uv(速度更快)
uv venv
source .venv/bin/activate

# 安装核心依赖
pip install torch torchvision torchaudio
pip install transformers datasets accelerate math_verify matplotlib pandas
python -m ipykernel install --user --name=raschka-reasoning

2. 下载 Qwen3-0.6B Base 模型

from transformers import AutoModelForCausalLM, AutoTokenizer

model_id = "Qwen/Qwen3-0.6B-Base"
tok = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(model_id, device_map="auto")

注意:必须用 Base 模型(非 Instruct),因为 GRPO 是在 pre-trained base 上激发 reasoning 能力的训练阶段。

3. 评估基座模型基准(15.2%)

uv run ../../ch03/02_math500-verifier-scripts/evaluate_math500.py \
  --dataset_size 500 \
  --which_model base

4. 核心:逐行理解 GRPO 四大步骤

GRPO 的核心实现在 rlvr_grpo_original_no_kl.py(495 行),四大步骤如下:

Step 1: 采样多个 response(rollout)并计算 reward

def reward_rlvr(answer_text, ground_truth):
    """验证 answer 是否与 ground truth 匹配"""
    extracted = extract_final_candidate(answer_text, fallback=None)  # 要求 \boxed{}
    if not extracted:
        return 0.0
    correct = grade_answer(extracted, ground_truth)
    return float(correct)

# 对每个 prompt 采样 num_rollouts 个 response
for _ in range(num_rollouts):
    token_ids, prompt_len, text = sample_response(model, tokenizer, prompt, ...)
    reward = reward_rlvr(text, example["answer"])
    roll_rewards.append(reward)

Step 2: 计算 advantage(组内相对优势)

rewards = torch.tensor(roll_rewards, device=device)
advantages = (rewards - rewards.mean()) / (rewards.std() + 1e-4)
# GRPO 的核心:不用 value network,而是用组内相对排名

为什么不用 value network?Raschka 在视频中解释:GRPO 来自 DeepSeekMath,用同 group 内采样结果的均值作为 baseline,比单独训练一个 value network 更省显存,也更稳定。

Step 3: 计算 sequence log probability

def sequence_logprob(model, token_ids, prompt_len):
    logits = model(token_ids.unsqueeze(0)).squeeze(0).float()
    logprobs = torch.log_softmax(logits, dim=-1)
    targets = token_ids[1:]
    selected = logprobs[:-1].gather(1, targets.unsqueeze(-1)).squeeze(-1)
    return selected[prompt_len - 1:].sum()

Step 4: 计算 GRPO loss 并反向传播

logps = torch.stack([sequence_logprob(model, token_ids, prompt_len)
                     for token_ids, prompt_len in rollout_data])

pg_loss = -(advantages.detach() * logps).mean()  # 注意: advantages 不反传
loss = pg_loss  # 没有 KL term(no-KL 版本的关键区别)

5. 运行 GRPO 训练(单 GPU)

uv run rlvr_grpo_original_no_kl.py \
  --num_rollouts 8 \
  --steps 100 \
  --max_new_tokens 512

推荐先用 --steps 50 跑一轮(Raschka 的实验显示 50 步效果最好,100 步反而下降)。

6. 多 GPU(FSDP)版本

uv run rlvr_grpo_original_no_kl_batched_fsdp.py \
  --num_rollouts 8 \
  --steps 50 \
  --max_new_tokens 512 \
  --num_gpus 4

7. 评估训练结果

# 替换 checkpoint_path 为实际生成的路径
uv run ../../ch03/02_math500-verifier-scripts/evaluate_math500.py \
  --dataset_size 500 \
  --which_model base \
  --checkpoint_path checkpoints/rlvr_grpo_original_no_kl/qwen3-0.6B-rlvr-grpo-step00050.pth

8. 训练可视化

uv run plot_metrics.py \
  --csv logs/rlvr_grpo_original_no_kl_metrics.csv \
  --moving_average 20

坑与适用边界

⚠️ KL term 是关键抉择,不同任务结论不同

Raschka 的实验明确显示:移除 KL divergence term 后 GRPO 性能大幅提升(47.4% vs 33.4%),但这是针对 math reasoning 的结论。GitHub README 明确注释引用了 DAPO、Dr. GRPO、Olmo 3 等 paper 建议移除 KL term。对于非 math 任务(如代码生成),KL term 可能仍是必要的——不要无脑照搬。

⚠️ 50 步 vs 100 步:不是越多越好

Raschka 视频原话(1:23:49):"when I trained it even longer, then it got even lower, um, 30 and so forth, it collapsed。"

实验数据(官方 README Row 5 & 6): - 50 步:47.4% - 100 步:44.0% - 200 步:更低(崩溃)

Vanilla GRPO 在 50 步后不稳定,Chapter 7 会介绍改进版本(Olmo3 mod、DeepSeek V3.2 mod)来延长稳定训练窗口。

⚠️ 显存需求不低,单卡需要 ~20GB

官方显存表: | num_rollouts | max_new_tokens | 所需 RAM | |-------------|--------------|--------| | 8 | 512 | 20.31 GB | | 8 | 1024 | 30.50 GB | | 4 | 512 | 12.80 GB |

消费级 24GB 显存的 RTX 4090 可以跑 --max_new_tokens 512 --num_rollouts 8,但 30GB+ 需求意味着 A6000 或 H100 更稳妥。

⚠️ reward 设计决定上限

当前实现用的是 \boxed{} 答案匹配(math_verify 库),对于没有唯一数值答案的开放式问题无法使用。如果你的任务不是 math,需要自行设计 reward function。

⚠️ Base vs Instruct 模型选择

必须用 Base 模型开始 GRPO,而不是 Instruct 模型。Raschka 强调 Instruct 模型已经过 SFT,可能与 GRPO 的 RL 目标冲突。

✅ 最适合的场景

  • 想从代码层面理解 GRPO 而不是调用黑盒库
  • 在消费级 GPU(24GB)上快速验证 GRPO 效果
  • 构建 math reasoning 专用小模型(0.6B~1.7B)
  • 准备深入 Chapter 7(GRPO 改进变体)前的预备学习

一句话结论

Raschka 第 6 章的 GRPO walkthrough 是目前少有的「逐行代码 + 真实 benchmark + 可直接跑」的学习资源;Qwen3-0.6B Base 经 50 步 GRPO 训练后 MATH-500 从 15.2% 提升到 47.4%(与官方同尺寸 reasoning 模型持平),但 vanilla GRPO 超过 50 步容易崩溃,建议先跑通基础版再尝试 Chapter 7 的改进变体。


核验来源(按优先级):

  1. rasbt/reasoning-from-scratch GitHub ch06/02_rlvr_grpo_scripts_intro README — GRPO 基准对比表 ✅
  2. rasbt/reasoning-from-scratch ch06/02_rlvr_grpo_scripts_intro/rlvr_grpo_original_no_kl.py(raw)— 逐行代码实现 ✅
  3. Raschka LinkedIn 帖子 — "15% → 47%" 声明 + 官方 reasoning 模型对比 ✅
  4. Raschka Substack(@rasbt)— 同 LinkedIn 表述 + 训练集来源 ✅
  5. YouTube: Build A Reasoning Model From Scratch 6(2026-10-03)— 视频 1:22:33 确认基准数据 ✅
  6. rasbt/reasoning-from-scratch GitHub 主仓库 README — 全书结构确认 ✅

不确定处: - "两小时"训练时间的具体硬件配置未披露,无法独立复现 - 12k MATH 训练集(rasbt/math_full_minus_math500)的具体处理细节未逐一核验 - "pretty solid model" 为 Raschka 主观描述,无法量化