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 章做了三件事:
- 逐行实现 GRPO 四大核心步骤:advantage 计算 → reward 获取 → sequence logprob → GRPO loss
- 直接在 Qwen3-0.6B Base 上训练,用 12k MATH 训练集跑出可量化的 benchmark 提升
- 配套多 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 的改进变体。
核验来源(按优先级):
- rasbt/reasoning-from-scratch GitHub ch06/02_rlvr_grpo_scripts_intro README — GRPO 基准对比表 ✅
- rasbt/reasoning-from-scratch ch06/02_rlvr_grpo_scripts_intro/rlvr_grpo_original_no_kl.py(raw)— 逐行代码实现 ✅
- Raschka LinkedIn 帖子 — "15% → 47%" 声明 + 官方 reasoning 模型对比 ✅
- Raschka Substack(@rasbt)— 同 LinkedIn 表述 + 训练集来源 ✅
- YouTube: Build A Reasoning Model From Scratch 6(2026-10-03)— 视频 1:22:33 确认基准数据 ✅
- rasbt/reasoning-from-scratch GitHub 主仓库 README — 全书结构确认 ✅
不确定处: - "两小时"训练时间的具体硬件配置未披露,无法独立复现 - 12k MATH 训练集(rasbt/math_full_minus_math500)的具体处理细节未逐一核验 - "pretty solid model" 为 Raschka 主观描述,无法量化