
CS336:从零手搓LLM-后训练部分
从零实现 SFT、Expert Iteration 和 GRPO,训练一个 1.5B 参数的语言模型在 MATH 数据集上学会数学推理。
一、实验动机
语言模型有一个令人兴奋的能力——思维链推理(Chain-of-Thought Reasoning):在给出最终答案前,先生成一步步的推理过程,显著提升复杂问题的解决准确率。
本实验的目标是:使用 Qwen 2.5 Math 1.5B 这个预训练的数学基础模型,通过三种递进的训练方法,让它学会在 MATH(高中竞赛数学题)上做推理:
方法 | 核心思想 | 监督信号 |
|---|---|---|
SFT(监督微调) | 用 DeepSeek R1 的推理链数据做模仿学习 | 人类/强模型标注的推理链 |
Expert Iteration | 模型自己采样 → 筛选正确回答 → 再做 SFT,循环迭代 | Reward 函数(答案对错) |
GRPO(组相对策略优化) | 用强化学习的策略梯度,根据 reward 信号直接优化模型 | Reward 函数(答案对错) |
二、基础组件
2.1 提示词格式
所有实验统一使用 R1-Zero 提示词,模型被训练输出特定格式:
<think>
一步步的推理过程...
</think>
<answer>
\boxed{最终答案}
</answer>这种格式便于后续自动化解析答案和判分。
2.2 答案判分
判分由 r1_zero_reward_fn 完成,返回三个分数:
{
"format_reward": 0.0 | 1.0, # 格式是否正确(含 <think>/<answer> 标签)
"answer_reward": 0.0 | 1.0, # 答案是否数学等价于标准答案
"reward": 0.0 | 1.0 # 总奖励 = 格式分 × 答案分
}答案等价判断经过三层 fallback:
字符串规范化匹配:去掉空格、统一
\frac写法、补齐花括号等SymPy 表达式等价:将
\frac{1}{2}和0.5转为数学表达式比较LaTeX 解析:用
math_verify做更深层的语义等价
2.3 梯度累积
1.5B 模型单卡显存有限,采用 梯度累积(Gradient Accumulation)技术:将 batch 拆成多个 microbatch,每次只做 loss.backward(),等累积够 gradient_accumulation_steps 步后再统一更新参数。
三、SFT(监督微调)
原理
SFT 是最简单直接的方法——把 DeepSeek R1 生成的"问题 → 推理链 → 答案"数据喂给 Qwen 模型,用交叉熵损失训练:
只计算输出部分的 loss(用 response_mask 屏蔽 prompt 和 padding)。
关键实现
在 sft_utils.py 中实现了以下核心函数:
函数 | 作用 |
|---|---|
| 分词 prompt 和 output,构造 |
| 通过模型 forward pass 获取每 token 的 log-probability |
| 计算逐 token 熵,监控模型置信度变化 |
| 只在 response 位置求和归一化 |
| 单步 SFT 训练(含梯度累积和 backward) |
实验结论
用全量数据 SFT 可达到 15%+ 验证准确率
仅用正确回答(过滤错误样本)训练的模型性能反而更好——这直接启发了 Expert Iteration
四、Expert Iteration(专家迭代)
原理
Expert Iteration 是一个自举循环:
flowchart TD
A["当前模型 π_θ"] --> B["对每道题采样 G 条回答"]
B --> C["用 reward_fn 判对错"]
C --> D["筛掉错误回答\n只保留 reward=1 的 (q, o) 对"]
D --> E["用保留的数据做 SFT\n更新 π_θ"]
E --> A核心洞察:不需要人类标注推理链,模型自己生成、自己筛选、自己学习。
实验结果
5 轮迭代后,准确率持续提升
关键超参数:每道题的采样数 G、每轮的 SFT epoch 数、训练 batch 大小
五、GRPO(组相对策略优化)
这是本作业最核心的部分——用强化学习的方法训练语言模型。
5.1 语言模型 = 策略(Policy)
在 RL 视角下:
RL 概念 | LLM 场景 |
|---|---|
状态 (s_t) | 已生成的前缀 token 序列 |
动作 (a_t) | 生成下一个 token |
策略 (\pi_\theta(a_t \mid s_t)) | 模型的 softmax 输出概率分布 |
轨迹 (\tau) | 一条完整回答(从 |
奖励 (R(\tau)) | 答案是否正确(0 或 1) |
5.2 策略梯度(Policy Gradient)
朴素策略梯度的核心公式:
直观理解:好回答(reward=1)的所有 token 被"鼓励"(概率增大),坏回答的 token 被"压制"(概率减小)。
代码实现仅一行:
def compute_naive_policy_gradient_loss(raw_rewards, policy_log_probs):
return -raw_rewards * policy_log_probs # 负号 = 最小化 → 等价于梯度上升5.3 Baseline 降低方差
直接用原始 reward(只有 0/1)作为梯度权重,方差很大。引入 baseline 后:
GRPO 的巧思:不用额外训练 Critic 网络,而是用同一道题的 G 条回答的组内均值当作 baseline。
# 组归一化(Group Normalization)
avg = mean(rewards_in_group) # baseline = 组内均值
advantages = [(r - avg) for r in rewards] # 相对优势
# 可选:再除以组内标准差
advantages = [a / (std + eps) for a in advantages]5.4 GRPO-Clip 防止策略突变
在 off-policy 训练中(同一批采样数据上做多次梯度更新),新旧策略差异过大可能导致训练崩溃。Clip 机制限制策略更新幅度:
def compute_grpo_clip_loss(advantages, policy_log_probs, old_log_probs, cliprange):
ratio = torch.exp(policy_log_probs - old_log_probs) # π_new / π_old
loss = -torch.min(
ratio * advantages, # 无限制项
torch.clamp(ratio, 1-cliprange, 1+cliprange) * advantages # Clip 项
)
return loss5.5 完整 GRPO 训练流程
flowchart TD
A["1. vLLM 采样:对每道题生成 G 条回答"] --> B["2. 保存 old_log_probs(旧策略的 log-prob)"]
B --> C["3. reward_fn 打分 → run_compute_group_normalized_rewards()"]
C --> D["4. 组归一化:advantages = (r - mean) / std"]
D --> E["5. tokenize + get_response_log_probs 获取 policy_log_probs"]
E --> F["6. compute_grpo_clip_loss() 计算 per-token loss"]
F --> G["7. masked_mean() 只在 response 位置求均值"]
G --> H["8. loss /= gradient_accumulation_steps → loss.backward()"]
H --> I["9. optimizer.step()"]
I -->|下一轮| A5.6 实现的 6 个 GRPO 核心函数
全部位于 grpo_utils.py,通过 14 个测试用例验证:
函数 | 分值 | 作用 |
|---|---|---|
| 2 | 组归一化:raw_rewards → advantages |
| 1 | 朴素策略梯度: |
| 2 | GRPO-Clip: |
| 1 | 统一入口,分发到 |
| 1 | 只在 response 位置求均值 |
| 3 | 完整的前向+反向微批次训练 |
六、实验消融总结
GRPO 实验中进行了多组消融对比:
消融主题 | 对比项 | 核心发现 |
|---|---|---|
学习率 |
| 太大会发散,太小收敛慢, |
Baseline 效果 |
| 有 baseline 方差更低,收敛更稳定 |
长度归一化 |
| 除以固定常数比除实际长度更稳定 |
标准差归一化 |
| 去掉标准差归一化(Dr.GRPO 方案)避免过难/过易题目主导梯度 |
On-policy vs Off-policy | 1 epoch vs 多 epoch | Off-policy 配合 Clip 可有效复用采样数据 |
Clip 消融 | GRPO-Clip vs GRPO-No-Clip | 多步 off-policy 时 Clip 必不可少,防止崩溃 |
Prompt 消融 | R1-Zero prompt vs question-only | 简单 prompt 可能因与预训练数据更匹配而表现更好 |
七、核心代码片段
Advantage 的计算
def run_compute_group_normalized_rewards(reward_fn, rollout_responses,
repeated_ground_truths, group_size,
advantage_eps, normalize_by_std):
raw_rewards, group_normalized_rewards, curr_group = [], [], []
for i, response in enumerate(rollout_responses):
raw_reward = reward_fn(response, repeated_ground_truths[i])["reward"]
raw_rewards.append(raw_reward)
curr_group.append(raw_reward)
if (i + 1) % group_size == 0:
avg = sum(curr_group) / group_size
curr_norm = [n - avg for n in curr_group]
if normalize_by_std:
curr_norm = [n / (advantage_eps + statistics.stdev(curr_group))
for n in curr_norm]
group_normalized_rewards.extend(curr_norm)
curr_group = []
return (torch.tensor(group_normalized_rewards),
torch.tensor(raw_rewards), {})完整的微批次训练步
def grpo_microbatch_train_step(policy_log_probs, response_mask,
gradient_accumulation_steps, loss_type,
raw_rewards, advantages, old_log_probs, cliprange):
loss, meta = compute_policy_gradient_loss(
policy_log_probs, loss_type, raw_rewards, advantages, old_log_probs, cliprange
)
loss = masked_mean(loss, response_mask) # 只在 response token 求均值
loss /= gradient_accumulation_steps # 梯度累积调整
loss.backward() # 反向传播
return loss, {}八、关键 Insight
SFT 是基础,RL 是放大器:SFT 教会模型基本格式和推理框架,GRPO 在此基础上通过 reward 信号进一步优化。缺少 SFT 冷启动,纯 RL 很难收敛。
组归一化(Group Normalization)是 GRPO 的核心创新:同一道题的 G 条回答互相对比,自动生成 baseline,避免了传统 Actor-Critic 方法需要额外训练 Critic 网络的负担。
Clip 机制至关重要:在 off-policy 训练中,没有 Clip 约束的策略更新可能因为单步梯度过大导致模型崩溃。Clip 是 PPO/GRPO 稳定性的基石。
答案解析器(Reward Function)决定了训练质量:如果 reward 函数不能准确判断答案对错,整个 RL 训练的方向就会偏。三层 fallback(字符串 → SymPy → LaTeX)的设计体现了工程中对鲁棒性的追求。
系统工程不可忽视:梯度累积、vLLM 高效推理、多 GPU 调度(一个 GPU 跑训练、一个 GPU 跑推理)都是让 1.5B 模型训练可行的关键工程手段。
九、总结
这次作业完整覆盖了从监督微调到强化学习的整个对齐训练管线:
SFT(学格式) → Expert Iteration(自举提升) → GRPO(RL 优化)通过亲手实现 6 个 GRPO 核心函数和完整的 SFT 工具链,深入理解了策略梯度、Advantage 估计、Clip 机制、组归一化等 RL 核心概念——这些正是 DeepSeek R1、OpenAI o1 等前沿推理模型背后的关键技术。
本文基于 CS336 Spring 2025 Assignment 5 实验编写,代码仓库:github.com/stanford-cs336/assignment5-alignment

Comments
Discuss this project
Emoji supported. Comments appear immediately.