三藏签名
< Back to projectsCS336:从零手搓LLM-后训练部分

CS336:从零手搓LLM-后训练部分

LLMRLSFT

从零实现 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:

  1. 字符串规范化匹配:去掉空格、统一 \frac 写法、补齐花括号等

  2. SymPy 表达式等价:将 \frac{1}{2}0.5 转为数学表达式比较

  3. LaTeX 解析:用 math_verify 做更深层的语义等价

2.3 梯度累积

1.5B 模型单卡显存有限,采用 梯度累积(Gradient Accumulation)技术:将 batch 拆成多个 microbatch,每次只做 loss.backward(),等累积够 gradient_accumulation_steps 步后再统一更新参数。


三、SFT(监督微调)

原理

SFT 是最简单直接的方法——把 DeepSeek R1 生成的"问题 → 推理链 → 答案"数据喂给 Qwen 模型,用交叉熵损失训练:

LSFT=tlogpθ(otq,o<t)\mathcal{L}_{\text{SFT}} = -\sum_{t} \log p_\theta(o_t \mid q, o_{<t})

只计算输出部分的 loss(用 response_mask 屏蔽 prompt 和 padding)。

关键实现

sft_utils.py 中实现了以下核心函数:

函数

作用

tokenize_prompt_and_output

分词 prompt 和 output,构造 response_mask

get_response_log_probs

通过模型 forward pass 获取每 token 的 log-probability

compute_entropy

计算逐 token 熵,监控模型置信度变化

masked_normalize

只在 response 位置求和归一化

sft_microbatch_train_step

单步 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)

一条完整回答(从 <think></answer>

奖励 (R(\tau))

答案是否正确(0 或 1)

5.2 策略梯度(Policy Gradient)

朴素策略梯度的核心公式:

θJ(θ)=Eτ[tθlogπθ(atst)R(τ)]\nabla_\theta J(\theta) = \mathbb{E}_{\tau}\left[ \sum_t \nabla_\theta \log \pi_\theta(a_t|s_t) \cdot R(\tau) \right]

直观理解:好回答(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 后:

θJ=E[tθlogπθ(atst)(R(τ)b)]\nabla_\theta J = \mathbb{E}\left[ \sum_t \nabla_\theta \log \pi_\theta(a_t|s_t) \cdot (R(\tau) - b) \right]

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 机制限制策略更新幅度:

L=min(πθπθoldA,clip ⁣(πθπθold,1ϵ,1+ϵ)A)\mathcal{L} = -\min\left( \frac{\pi_\theta}{\pi_{\theta_{\text{old}}}} \cdot A,\quad \text{clip}\!\left(\frac{\pi_\theta}{\pi_{\theta_{\text{old}}}}, 1-\epsilon, 1+\epsilon\right) \cdot A \right)

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 loss

5.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 -->|下一轮| A

5.6 实现的 6 个 GRPO 核心函数

全部位于 grpo_utils.py,通过 14 个测试用例验证:

函数

分值

作用

run_compute_group_normalized_rewards

2

组归一化:raw_rewards → advantages

compute_naive_policy_gradient_loss

1

朴素策略梯度:-A × log π

compute_grpo_clip_loss

2

GRPO-Clip:-min(ratio×A, clip(ratio)×A)

compute_policy_gradient_loss

1

统一入口,分发到 no_baseline / reinforce_with_baseline / grpo_clip

masked_mean

1

只在 response 位置求均值

grpo_microbatch_train_step

3

完整的前向+反向微批次训练


六、实验消融总结

GRPO 实验中进行了多组消融对比:

消融主题

对比项

核心发现

学习率

1e-6 ~ 1e-4

太大会发散,太小收敛慢,1e-5 附近最优

Baseline 效果

no_baseline vs reinforce_with_baseline

有 baseline 方差更低,收敛更稳定

长度归一化

masked_mean vs masked_normalize

除以固定常数比除实际长度更稳定

标准差归一化

normalize_by_std=True/False

去掉标准差归一化(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

  1. SFT 是基础,RL 是放大器:SFT 教会模型基本格式和推理框架,GRPO 在此基础上通过 reward 信号进一步优化。缺少 SFT 冷启动,纯 RL 很难收敛。

  2. 组归一化(Group Normalization)是 GRPO 的核心创新:同一道题的 G 条回答互相对比,自动生成 baseline,避免了传统 Actor-Critic 方法需要额外训练 Critic 网络的负担。

  3. Clip 机制至关重要:在 off-policy 训练中,没有 Clip 约束的策略更新可能因为单步梯度过大导致模型崩溃。Clip 是 PPO/GRPO 稳定性的基石。

  4. 答案解析器(Reward Function)决定了训练质量:如果 reward 函数不能准确判断答案对错,整个 RL 训练的方向就会偏。三层 fallback(字符串 → SymPy → LaTeX)的设计体现了工程中对鲁棒性的追求。

  5. 系统工程不可忽视:梯度累积、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.

No comments yet.