跳转到正文

15.1 GRPO 训练机制

上一章我们深入了 DPO 的理论与实践,看到它可以直接从固定的偏好数据里学习:同一个 prompt 下,chosen 应该比 rejected 更可能出现。现在我们回到在线训练:模型不再只读别人已经标好的偏好对,而是在训练过程中自己生成回答、自己得到反馈、再用反馈更新自己。

GRPO 的入口是同题多答。给定同一道题,模型一次生成多个回答;奖励函数分别给这些回答打分;然后只在这一组回答内部比较谁更好。它表面上像"让模型多试几次",真正解决的问题是:

没有 Critic 的时候,模型怎么判断某个回答是比预期好,还是比预期差?

一个直观答案是:拿它和同一道题的其他回答比。GRPO 就是沿着这个思路,把同题多答变成可以训练的策略优化方法。

本节沿着一次完整的 GRPO 训练轨迹来讲:先看同题多答怎样产生组内比较,再解释为什么"和同题其他回答比"可以替代 Critic,接着写出优势、概率比值和裁剪目标,最后回到手写代码和 GSM8K 训练实验。

Mermaid diagram

这张图先表达一个最基本的训练信号:同一道题多答几次,每个回答都有分数;高于同组平均分的回答以后更容易出现,低于同组平均分的回答以后更少出现

GRPO 的入口

一个带数字的微缩例子

用一个具体例子走一遍。假设题目是:

小明有 3 个苹果,又买了 2 个,现在一共有几个?

模型对同一道题一次写出 4 个回答,规则打分如下:

回答模型写了什么分数
1"3 + 2 = 5,所以答案是 5。"1.5
2"答案是 5。"1.0
3"应该是 6。"0.0
4"不确定,可能是 4。"0.0

这 4 个分数的平均分是:

于是模型会这样理解这组回答:

回答和平均分比较之后怎么学
1明显比平均好,以后多生成它
2也比平均好,稍微多生成它
3比平均差,以后少生成它
4比平均差,以后少生成它

这里的"比平均分高多少、低多少",后面会被正式叫做优势。在这个例子里,优势就是"这份回答在同题四个回答里表现得比平均好还是差"。

把语言模型放进强化学习框架

为了用 RL 语言讲清楚 GRPO,先把对应关系列出来:

强化学习概念在数学推理模型里是什么
状态 题目 prompt 加上已经写出的推理步骤,也就是
动作 下一步生成的 token,也就是
轨迹 一整段推理过程和最终答案
奖励 答案是否正确、格式是否符合要求
策略 当前正在训练的语言模型

对一道题 来说,模型生成完整回答 就相当于走完一条轨迹。被训练的对象仍然是语言模型策略

需要澄清一个常见误会:GRPO 不是一个新的模型,也不只是"组内归一化"这个公式。GRPO 是一种在线训练策略模型的方法。 训练方式是:对同一个 prompt 一次生成多个回答,把这些回答放在同一组里打分,并比较:这个回答在同组里是否高于平均水平? 最后更新策略时仍然使用 PPO-style 的 ratio + clip,避免新策略离旧策略太远。

用一句话概括:

GRPO = 在线组采样 + 规则/奖励打分 + 组内相对优势 + PPO-style 裁剪更新。

把开头的苹果题翻译成这句话:同一道题一次生成 4 个回答,这是在线组采样;用答案正确性和格式给分,这是规则/奖励打分;用 这样的差值判断好坏,这是组内相对优势;最后让好回答概率上升、差回答概率下降,但每次只小步调整,这就是 PPO-style 裁剪更新

PPO Critic 的痛点

要理解 GRPO 为什么这样设计,先看它要替代的 Critic 有什么问题。

Critic 是什么

在 PPO 这类 Actor-Critic 方法里,Actor 是负责生成回答的策略模型,Critic 则像一个"价值评估器":它不直接生成回答,而是估计"当前已经写到这里,后面大概能拿到多少总奖励"。用公式写就是价值函数:

其中 是当前状态——对语言模型来说可以粗略理解为"prompt 加上已经生成的前几个 token"; 是 Critic 自己的参数。Critic 的作用是给策略更新提供一个基线:如果某个回答的真实奖励比 Critic 预估的更高,就说明这个回答比预期好,应该提高概率;如果比预期低,就应该降低概率。

如果照 PPO 的路线走,优势大致写成:

这句话的意思是:不要只看奖励高不高,要看它有没有比 Critic 的预期更好。这在传统强化学习里很自然,但在 LLM 数学推理里就很重。

Critic 在 LLM 训练中的三大问题

1. 吃显存:Critic 与 Actor 同等规模,PPO 需要同时装下 Actor + Critic + Reference + RM 四个模型。

2. 训练不稳定:价值函数 需要从"部分生成的文本"预测"最终得分",但 LLM 序列很长(500+ tokens),监督信号只在末尾才有,方差极大。

3. 工程复杂:四个模型各有一套优化器、学习率、梯度裁剪配置,调参难度指数级增长。

回顾第 6 章基线分析第 7 章优势函数,Critic 的核心作用是提供基线来降低方差。如果不需要单独训练网络就能得到基线,Critic 就可以退休了——这就是 GRPO 的出发点。

GRPO 的核心 与 组内归一化替代 Critic

GRPO 的想法出奇地简单:不再单独训练 Critic,而是用同一个 prompt 下多个回答的平均分临时充当基线。DeepSeekMath 论文提出 GRPO 时,明确说它 "foregoes the critic model",并用组内分数来估计基线。

GRPO 从 PPO 中替换 Critic 基线

因此,GRPO 与 PPO 的关系可以概括为:

  • PPO 问:这个回答比 Critic 预估的平均水平好吗?
  • GRPO 问:这个回答比同一道题的其他回答好吗?
  • PPO 和 GRPO 都还会用概率比值和裁剪,避免一次更新过大。

论文脉络:GRPO 来自 DeepSeekMath 论文 DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models。它不是完全抛弃 PPO,而是在 PPO 框架里去掉 Critic,用组内相对奖励构造优势。

GRPO 把 PPO 里的 Critic 基线换成"同一道题的一组回答的平均分":原来问"这个回答比 Critic 预期好吗",现在问"这个回答比同题其他回答好吗"。这就是"组内相对优势"的直觉:同一道题里,谁比平均好,就多学谁;谁比平均差,就少生成谁

代码地图

下面是一份最小手写 GRPO 代码地图。它不是 trl 的工程源码,而是把 GRPO 的数学结构摊开给你看:每个公式后面都能回到这份代码里的某几行。

 13# [A] 组采样:每个 prompt 生成 group_size 个回答,并保留原始 token 边界
 14def sample_groups(model, tokenizer, prompts, group_size=8, max_new_tokens=256):
 15    expanded_prompts = [prompt for prompt in prompts for _ in range(group_size)]
 16    prompt_batch = tokenizer(expanded_prompts, padding=True, return_tensors="pt")
 17    prompt_batch = {
 18        key: value.to(model.device)
 19        for key, value in prompt_batch.items()
 20    }
 21 
 22    pad_token_id = tokenizer.pad_token_id
 23    if pad_token_id is None:
 24        pad_token_id = tokenizer.eos_token_id
 25 
 26    with torch.no_grad():
 27        output_ids = model.generate(
 28            **prompt_batch,
 29            do_sample=True,
 30            temperature=1.0,
 31            max_new_tokens=max_new_tokens,
 32            pad_token_id=pad_token_id,
 33            eos_token_id=tokenizer.eos_token_id,
 34        )
 35 
 36    prompt_width = prompt_batch["input_ids"].size(1)
 37    completion_ids = output_ids[:, prompt_width:]
 38    completion_mask = completion_mask_until_eos(
 39        completion_ids,
 40        tokenizer.eos_token_id,
 41    )
 42    attention_mask = torch.cat(
 43        [prompt_batch["attention_mask"], completion_mask.to(torch.long)],
 44        dim=1,
 45    )
 46 
 47    responses = tokenizer.batch_decode(completion_ids, skip_special_tokens=True)
 48    group_ids = torch.arange(len(prompts), device=model.device).repeat_interleave(
 49        group_size
 50    )
 51    batch = {
 52        "input_ids": output_ids,
 53        "attention_mask": attention_mask,
 54        "completion_mask": completion_mask,
 55    }
 56    return responses, group_ids, batch
 80# [C] 组内优势:用同题目的回答均值替代 Critic 基线
 81def group_advantages(rewards, group_size=8, eps=1e-8):
 82    grouped_rewards = rewards.view(-1, group_size)
 83    group_mean = grouped_rewards.mean(dim=1, keepdim=True)
 84    group_std = grouped_rewards.std(dim=1, keepdim=True, correction=1)
 85 
 86    advantages = (grouped_rewards - group_mean) / (group_std + eps)
 87    advantages = torch.where(
 88        group_std < eps,
 89        torch.zeros_like(advantages),
 90        advantages,
 91    )
 92    return advantages.reshape(-1)
 93 
 94 
 95# [D] 只返回回答部分的逐 token log probability,形状为 [B, T]
 96def per_token_logprobs(model, input_ids, attention_mask, completion_length):
 97    outputs = model(input_ids=input_ids, attention_mask=attention_mask)
 98    logits = outputs.logits[:, :-1, :]
 99    target_ids = input_ids[:, 1:]
100 
101    all_token_logprobs = logits.log_softmax(dim=-1)
102    picked_logprobs = all_token_logprobs.gather(
103        dim=-1,
104        index=target_ids.unsqueeze(-1),
105    ).squeeze(-1)
106    return picked_logprobs[:, -completion_length:]
107 
108 
109def masked_sequence_mean(values, mask):
110    """每段回答先按有效 token 求平均,防止长回答获得更大权重。"""
111    mask = mask.to(values.dtype)
112    token_count = mask.sum(dim=-1).clamp_min(1.0)
113    return (values * mask).sum(dim=-1) / token_count
116# [E-F] 原始 GRPO:逐 token ratio、clip、KL,再按回答长度归一化
117def grpo_objective_from_logprobs(
118    new_logprobs,
119    old_logprobs,
120    ref_logprobs,
121    completion_mask,
122    advantages,
123    clip_eps=0.2,
124    kl_coef=0.04,
125):
126    token_ratio = torch.exp(new_logprobs - old_logprobs)
127    token_advantages = advantages.unsqueeze(-1)
128    unclipped = token_ratio * token_advantages
129    clipped_ratio = torch.clamp(token_ratio, 1.0 - clip_eps, 1.0 + clip_eps)
130    clipped = clipped_ratio * token_advantages
131 
132    # DeepSeekMath 式 (4):D_KL(policy || ref) 的逐 token 无偏正值估计
133    log_ratio_ref = ref_logprobs - new_logprobs
134    per_token_kl = torch.exp(log_ratio_ref) - log_ratio_ref - 1.0
135 
136    per_token_objective = torch.minimum(unclipped, clipped) - kl_coef * per_token_kl
137    per_response_objective = masked_sequence_mean(
138        per_token_objective,
139        completion_mask,
140    )
141    loss = -per_response_objective.mean()
142 
143    policy_loss = -masked_sequence_mean(
144        torch.minimum(unclipped, clipped),
145        completion_mask,
146    ).mean()
147    approx_kl = masked_sequence_mean(per_token_kl, completion_mask).mean()
148    metrics = {
149        "loss": loss.detach(),
150        "policy_loss": policy_loss.detach(),
151        "approx_kl": approx_kl.detach(),
152        "mean_ratio": masked_sequence_mean(token_ratio, completion_mask).mean().detach(),
153    }
154    return loss, metrics
155 
156 
157def grpo_loss(
158    policy_model,
159    ref_model,
160    batch,
161    advantages,
162    old_logprobs=None,
163    clip_eps=0.2,
164    kl_coef=0.04,
165):
166    completion_length = batch["completion_mask"].size(1)
167    new_logprobs = per_token_logprobs(
168        policy_model,
169        batch["input_ids"],
170        batch["attention_mask"],
171        completion_length,
172    )
173 
174    # 单次更新时 old policy 就是采样 policy;detach 保留 ratio 的梯度。
175    if old_logprobs is None:
176        old_logprobs = new_logprobs.detach()
177 
178    with torch.no_grad():
179        ref_logprobs = per_token_logprobs(
180            ref_model,
181            batch["input_ids"],
182            batch["attention_mask"],
183            completion_length,
184        )
185 
186    return grpo_objective_from_logprobs(
187        new_logprobs,
188        old_logprobs,
189        ref_logprobs,
190        batch["completion_mask"],
191        advantages,
192        clip_eps,
193        kl_coef,
194    )
197# [G] 训练步骤:采样、打分、组内归一化、再反向传播
198def train_step(
199    policy_model,
200    ref_model,
201    optimizer,
202    tokenizer,
203    prompts,
204    ground_truths,
205    group_size=8,
206):
207    responses, _, batch = sample_groups(
208        policy_model,
209        tokenizer,
210        prompts,
211        group_size,
212    )
213    rewards = score_responses(
214        responses,
215        ground_truths,
216        group_size,
217        policy_model.device,
218    )
219    advantages = group_advantages(rewards, group_size)
220 
221    loss, metrics = grpo_loss(policy_model, ref_model, batch, advantages)
222    optimizer.zero_grad()
223    loss.backward()
224    optimizer.step()
225    return metrics
226 
227 
228# [H] GRPO 训练循环:每轮都在线生成新回答
229def train_grpo(policy_model, ref_model, optimizer, tokenizer, dataloader):
230    ref_model.eval()
231    for prompts, ground_truths in dataloader:
232        metrics = train_step(
233            policy_model,
234            ref_model,
235            optimizer,
236            tokenizer,
237            prompts,
238            ground_truths,
239        )
240        print(
241            "loss=",
242            float(metrics["loss"]),
243            "kl=",
244            float(metrics["approx_kl"]),
245        )

这份代码可以分成八块:

标记代码部分后文会解释什么
[A]sample_groups为什么每个 prompt 要生成多个回答
[B]rule_reward / score_responses奖励从哪里来,为什么数学题不需要 RM
[C]group_advantages组内均值如何替代 Critic 基线
[D]per_token_logprobs如何保留回答中每个 token 的
[E]grpo_objective_from_logprobs逐 token 的 ratioclip 和策略更新
[F]per_token_kl为什么 KL 也必须在逐 token 层面计算
[G]train_step采样、打分、优势、loss、反向传播如何接起来
[H]train_grpo为什么 GRPO 是在线训练,每轮都生成新回答

从 PPO 改到 GRPO:到底替换了哪几行

如果不改成 GRPO,而是继续按 PPO / RLHF 的方式训练,代码直觉通常是这样:

python
# PPO / RLHF:在线生成,然后让 Critic 估计逐 token 优势
responses, completion_mask = policy_old.generate(prompts)
old_per_token_logps = per_token_logprobs(policy_old, prompts, responses).detach()

rewards = reward_model(prompts, responses)
advantages = critic_based_advantages(prompts, responses, rewards)

new_per_token_logps = per_token_logprobs(policy, prompts, responses)
token_ratio = torch.exp(new_per_token_logps - old_per_token_logps)
per_token_objective = torch.min(
    token_ratio * advantages,
    torch.clamp(token_ratio, 1 - clip_eps, 1 + clip_eps) * advantages,
)
ppo_loss = -masked_sequence_mean(per_token_objective, completion_mask).mean()

这里的 critic 就是前面说的价值模型。它的工作不是生成答案,而是估计一个基线:这个 prompt 和当前回答前缀,大概应该拿多少分。然后 PPO 用 rewards - values 得到优势,判断某个回答是"比预期好"还是"比预期差"。

GRPO 的改法很集中:保留在线生成、概率比值和裁剪,但不再训练 Critic;优势改成从同一个 prompt 的一组回答里算出来

python
# 同一个 prompt 生成 G 个回答,然后做组内比较
responses, completion_mask = generate_many(policy_old, prompts, num_generations=G)
old_per_token_logps = per_token_logprobs(policy_old, prompts, responses).detach()

rewards = reward_fn(prompts, responses)
rewards_by_group = rewards.view(batch_size, G)

group_mean = rewards_by_group.mean(dim=1, keepdim=True)
group_std = rewards_by_group.std(dim=1, keepdim=True)
response_advantages = (
    (rewards_by_group - group_mean) / (group_std + 1e-4)
).view(-1)

new_per_token_logps = per_token_logprobs(policy, prompts, responses)
token_ratio = torch.exp(new_per_token_logps - old_per_token_logps)
per_token_objective = torch.min(
    token_ratio * response_advantages[:, None],
    torch.clamp(token_ratio, 1 - clip_eps, 1 + clip_eps)
    * response_advantages[:, None],
)
grpo_loss = -masked_sequence_mean(per_token_objective, completion_mask).mean()

这里有一个容易漏掉的层次:结果监督只给每段回答一个奖励,所以同一回答里的 token 共享一个 ;新旧策略的概率比值、裁剪和 KL 仍然分别作用于每个 token。最后先对每段回答的有效 token 求平均,再对组内回答求平均。

把真正变化的几行单独拎出来,就是:

diff
  responses = policy_old.generate(prompts)
  rewards = reward_model_or_rule(prompts, responses)
- values = critic(prompts, responses)
- advantages = rewards - values

+ rewards_by_group = rewards.view(batch_size, G)
+ group_mean = rewards_by_group.mean(dim=1, keepdim=True)
+ group_std = rewards_by_group.std(dim=1, keepdim=True)
+ advantages = ((rewards_by_group - group_mean) / (group_std + 1e-4)).view(-1)

  loss = ppo_style_clipped_loss(logps_new, logps_old, advantages)

所以 GRPO 保留了 PPO 的逐 token 概率比值和裁剪,改变的是优势的来源:它用同题回答的组内相对奖励代替单独训练的 Critic。

先以原论文为验收标准

DeepSeekMath 原文的式(3)写成三层平均:先对一段回答的 个 token 求平均,再对同一问题的 个回答求平均,最后对问题分布取期望。式中的比值也带有 token 下标

因此,下面两种写法含义不同:

第二种写法等于 。它会把回答长度带进比值,而且整段回答只裁剪一次,不是 DeepSeekMath 式(3)。原论文的式(4)同样以 token 为单位计算 KL 估计。本文后面的公式和代码都以这两条原始定义为准。DeepSeekMath 原文公式

再对照 TRL 的工程实现

当前 Hugging Face TRL 的真实源码也保留了这个 token 维度。2026-08-28 查看 GRPOTrainer 时,可以看到这些对应关系:

  1. GRPOTrainer 的初始化参数里有 reward_funcs,它可以是奖励模型,也可以是普通 Python 函数。也就是说,数学题这类任务可以直接用规则函数打分,不一定要先训练 RM。
  2. self.num_generations = args.num_generations 对应公式里的 ,也就是每个 prompt 生成几个回答
  3. 源码会把 rewards reshape 成 (-1, num_generations),计算 mean_grouped_rewards 和组内 std_rewards,再得到 advantages = rewards - mean_grouped_rewards,必要时除以标准差。
  4. _get_per_token_logps_and_entropies 返回每个回答 token 的 log probability;损失部分计算逐 token 的 coef_1 = exp(log_ratio),再用 torch.clamp 得到 coef_2,最后逐 token 取 min
  5. loss_type="grpo" 先用回答掩码对每段回答求 token 平均,再对 batch 求平均。这对应原论文的 loss_type="bnpo" 则把整个 batch 的有效 token 一起平均,回答长度权重不同。当前 TRL 的损失实现

TRL 现在还支持 sequence-level importance sampling、DAPO/BNPO/DR-GRPO 等扩展,而且 GRPOConfig 的默认值会随版本演进。复现原始 DeepSeekMath 公式时,需要显式选择 importance_sampling_level="token"loss_type="grpo"num_iterations=1beta=0.04,并关闭后来加入的 KL 偏差修正。后文的配置会把这些选项全部写出来。对照 TRL 0.24 的 GRPOConfig当前 main 分支的 GRPOConfig,可以看到 loss_type 等默认值的变化,以及新版新增且默认开启的 use_bias_correction_kl

对照 PPOTrainer,差别就更清楚:PPOTrainer 需要 reward_modelvalue_model,并用 value_model 产生优势估计;GRPOTrainer 不需要单独的 value_model,它把同题多答的组内相对分数直接变成优势。

GRPO 的完整公式

前面用直觉和代码 diff 看过 GRPO 怎么工作。这一节严格按 DeepSeekMath 的式(3)和式(4)展开。需要同时保留两个层次:一段回答只有一个结果奖励,因此它的 token 共享同一个优势;策略比值、裁剪和 KL 则都在 token 层面计算。

下面四小节依次给出样本结构、组内优势、逐 token 的 PPO Clip 和逐 token 的 KL 惩罚,最后用一张数据流图收尾。

样本结构与组采样

GRPO 的训练样本不是"一个 prompt 配一个回答",而是一个 prompt 配一组回答。假设一个 batch 里有多个题目,用 表示第几个题目,用 表示这个题目下第几个回答:

每个字母的意思是:

  • :第 个 prompt,也就是一道题或一个问题。
  • :group size,每个 prompt 生成几个回答。代码里的 num_generations=8 就是
  • :第 个 prompt 下生成的第 个回答。
  • :生成这批回答时使用的旧策略。它负责采样数据。
  • :正在被更新的新策略。它负责学习,让好回答更可能出现。

采样过程可以写成:

符号 表示"从某个分布中采样"。每个回答生成后,都要得到一个奖励:

这里 是奖励函数, 是一个标量。数学题里, 可以很简单:答案对就加分,格式规范也加分。GRPO 的关键不是"奖励函数一定很复杂",而是:同一道题下的多个回答会放在一起比较

代码里对应的是 [A] 组采样

 13# [A] 组采样:每个 prompt 生成 group_size 个回答,并保留原始 token 边界
 14def sample_groups(model, tokenizer, prompts, group_size=8, max_new_tokens=256):
 15    expanded_prompts = [prompt for prompt in prompts for _ in range(group_size)]
 16    prompt_batch = tokenizer(expanded_prompts, padding=True, return_tensors="pt")
 17    prompt_batch = {
 18        key: value.to(model.device)
 19        for key, value in prompt_batch.items()
 20    }
 21 
 22    pad_token_id = tokenizer.pad_token_id
 23    if pad_token_id is None:
 24        pad_token_id = tokenizer.eos_token_id
 25 
 26    with torch.no_grad():
 27        output_ids = model.generate(
 28            **prompt_batch,
 29            do_sample=True,
 30            temperature=1.0,
 31            max_new_tokens=max_new_tokens,
 32            pad_token_id=pad_token_id,
 33            eos_token_id=tokenizer.eos_token_id,
 34        )
 35 
 36    prompt_width = prompt_batch["input_ids"].size(1)
 37    completion_ids = output_ids[:, prompt_width:]
 38    completion_mask = completion_mask_until_eos(
 39        completion_ids,
 40        tokenizer.eos_token_id,
 41    )
 42    attention_mask = torch.cat(
 43        [prompt_batch["attention_mask"], completion_mask.to(torch.long)],
 44        dim=1,
 45    )
 46 
 47    responses = tokenizer.batch_decode(completion_ids, skip_special_tokens=True)
 48    group_ids = torch.arange(len(prompts), device=model.device).repeat_interleave(
 49        group_size
 50    )
 51    batch = {
 52        "input_ids": output_ids,
 53        "attention_mask": attention_mask,
 54        "completion_mask": completion_mask,
 55    }
 56    return responses, group_ids, batch

替代 Critic 的基线

GRPO 的核心思路在这里兑现:对同一个问题 ,先采样 个回答得到 个奖励 ,再做两步处理——减均值替代 Critic,除标准差归一化尺度。两步合起来得到组内优势:

其中 是组内均值, 是组内标准差, 是一个很小的数,防止标准差为 0 时除以 0。DeepSeekMath 原文只写 std,没有规定分母使用 还是 。示例跟随 PyTorch torch.std 的默认定义使用 Bessel 校正;这只会改变优势的整体缩放,不改变同组回答的正负顺序。

第一步:减均值替代 Critic。回顾 PPO 优势 ,本质是"奖励减基线"。Critic 学的 就是对"在这个 prompt 下平均能拿多少分"的估计。而组内均值 是这个估计的直接样本版本——同一道题的 个回答就是 次蒙特卡洛采样,平均起来就是无偏估计。代换:

这一步保证了: 的正负和"是否好于组平均"完全对齐;组内期望 ,与 Critic 基线的性质一致。

第二步:除标准差归一化尺度。不同题目的奖励尺度差异巨大——简单题组内奖励可能在 之间波动(, ),难题组内可能在 之间波动(, )。如果只减均值不除标准差,简单题和难题的梯度尺度相同——但简单题已经掌握了,不应该再主导梯度。除以 把所有题目的优势尺度拉到接近 1:

在统计学里这个变换叫 z-score 标准化,几何含义是把每组奖励平移到原点、缩放到单位方差,让不同分布可以在同一坐标轴上比较。

两步合起来读: 表示这个回答比同组平均好,应该提高概率; 表示比平均差,应该降低概率; 表示和平均差不多,不需要太强更新。这种"控制变量"式的组内比较也比跨样本的绝对评分更稳定——同一组内的回答共享相同的 prompt,唯一差异是模型生成的随机性。它也和人类偏好的本质对齐:判断本来就是"A 比 B 好"这种比较式的,不是"A 得 87 分"这种绝对的。

边界情形与代码对应

如果同一组回答奖励全都一样, 会接近 0,代码会把优势设成 0。这表示这道题暂时没有可学习的差异:大家都对,或者大家都错,模型不知道该更偏向哪一个回答。 的作用是避免 的数值问题。

代码里对应的是 [C] 组内优势

 80# [C] 组内优势:用同题目的回答均值替代 Critic 基线
 81def group_advantages(rewards, group_size=8, eps=1e-8):
 82    grouped_rewards = rewards.view(-1, group_size)
 83    group_mean = grouped_rewards.mean(dim=1, keepdim=True)
 84    group_std = grouped_rewards.std(dim=1, keepdim=True, correction=1)
 85 
 86    advantages = (grouped_rewards - group_mean) / (group_std + eps)
 87    advantages = torch.where(
 88        group_std < eps,
 89        torch.zeros_like(advantages),
 90        advantages,
 91    )
 92    return advantages.reshape(-1)

代码对应关系:

  • grouped_rewards = rewards.view(-1, group_size):把一维奖励列表重新排成"每行一个 prompt、每行 个回答"的形状。
  • group_mean = grouped_rewards.mean(dim=1, keepdim=True):计算每个 prompt 的
  • group_std = grouped_rewards.std(dim=1, keepdim=True):计算每个 prompt 的
  • advantages = (grouped_rewards - group_mean) / (group_std + eps):实现
  • torch.where(group_std < eps, 0, advantages):如果一组回答没有差异,就不给这组样本训练信号。

一句话总结:GRPO = PPO 的裁剪机制 + 用组内排名替代 Critic。下面两小节就把"PPO 的裁剪机制"完整展开。

策略比值与 PPO Clip:先保留 token,再计算 ratio

语言模型生成第 个 token 时,会根据问题和已经生成的前缀给它一个条件概率:

这些概率通常很小。一段回答的联合概率要把它们全部相乘,几十个小数连乘后很容易小到计算机无法稳定表示。代码因此先计算对数概率:

对数把乘法变成加法,所以整段回答的对数概率确实等于 。这条性质适合计算整段回答的概率,却不表示所有算法都应该立刻把 token 维度求和。原始 GRPO 要在每个 token 上分别计算比值和裁剪,因此代码必须保留形状为 的逐 token 对数概率。

个 token 的新旧策略比值是:

如果 ,新策略把这个 token 的条件概率提高了 20%;如果 ,新策略把它降低了 20%。回答级优势 会广播给这段回答中的每个有效 token。原论文的裁剪目标是:

这里的顺序很重要:先对每个 token 算 ,再分别裁剪,最后对一段回答的有效 token 求平均。假设一段三 token 回答的比值是 。原始 GRPO 会得到三个比值并分别裁剪成 。如果先把对数概率求和,得到的回答级比值会变成

此时整段回答只剩一个比值,再把它裁剪成 。三个 token 原本不同的变化被压成一个数,回答越长,连乘带来的长度效应也越强。这正是旧示例出错的地方。

为什么要裁剪?这批回答由 生成。新策略训练得越久,这批数据越不能代表它当前会生成什么。逐 token 裁剪限制每个生成决策能利用旧数据改变多少。详细推导见策略更新的约束机制

116# [E-F] 原始 GRPO:逐 token ratio、clip、KL,再按回答长度归一化
117def grpo_objective_from_logprobs(
118    new_logprobs,
119    old_logprobs,
120    ref_logprobs,
121    completion_mask,
122    advantages,
123    clip_eps=0.2,
124    kl_coef=0.04,
125):
126    token_ratio = torch.exp(new_logprobs - old_logprobs)
127    token_advantages = advantages.unsqueeze(-1)
128    unclipped = token_ratio * token_advantages
129    clipped_ratio = torch.clamp(token_ratio, 1.0 - clip_eps, 1.0 + clip_eps)
130    clipped = clipped_ratio * token_advantages
131 
132    # DeepSeekMath 式 (4):D_KL(policy || ref) 的逐 token 无偏正值估计
133    log_ratio_ref = ref_logprobs - new_logprobs
134    per_token_kl = torch.exp(log_ratio_ref) - log_ratio_ref - 1.0
135 
136    per_token_objective = torch.minimum(unclipped, clipped) - kl_coef * per_token_kl
137    per_response_objective = masked_sequence_mean(
138        per_token_objective,
139        completion_mask,
140    )
141    loss = -per_response_objective.mean()
142 
143    policy_loss = -masked_sequence_mean(
144        torch.minimum(unclipped, clipped),
145        completion_mask,
146    ).mean()
147    approx_kl = masked_sequence_mean(per_token_kl, completion_mask).mean()
148    metrics = {
149        "loss": loss.detach(),
150        "policy_loss": policy_loss.detach(),
151        "approx_kl": approx_kl.detach(),
152        "mean_ratio": masked_sequence_mean(token_ratio, completion_mask).mean().detach(),
153    }
154    return loss, metrics

代码中的对应关系是:

  • new_logprobsold_logprobs 的形状都是 ,每个位置保存一个回答 token 的对数概率。
  • token_ratio = exp(new_logprobs - old_logprobs) 逐位置实现
  • advantages.unsqueeze(-1) 把每段回答的一个优势广播到它的所有 token。
  • minimum(unclipped, clipped) 在每个 token 上选择更保守的目标。
  • masked_sequence_mean(..., completion_mask) 先对每段回答的有效 token 求平均;外层 .mean() 再让每段回答拥有相同权重。

直接写 (loss * mask).sum() / mask.sum() 会把整个 batch 的 token 一起平均,长回答拥有更多权重。TRL 把这种归约命名为 bnpo;原始 GRPO 对应的是每段回答先除以自己的长度。优化器执行最小化,所以代码最后对要最大化的目标加负号。

KL 惩罚:每个 token 都要和 Reference 比较

DeepSeekMath 还在每个回答 token 上加入 KL 惩罚,让 Policy 不要离 Reference 太远。对第 个 token,原文式(4)使用:

这个形式不是凭空选的,它满足三个关键性质。

性质一:每个 token 的估计值非负。令 ,则 。求导 ,所以 是凸函数,在 (即 )处取最小值 。任何 都给出正值,这避免了朴素估计 在单次采样上可能为负的问题。

性质二:是 的无偏估计。注意 ,所以 ;而 。代回去:

性质三:在小偏差处退化为二次型。把 处 Taylor 展开:,所以

几何含义 作为 的函数是一条 形曲线,最低点在 (Policy = Reference),开口由 主导。这正是"越偏离惩罚越大"在数学上的写照——而二次型主导意味着梯度在偏离小时温和、偏离大时变陡,避免一次性把策略推得太远。

把逐 token 裁剪和逐 token KL 合在一起,DeepSeekMath 式(3)的单个问题目标是:

训练损失是 。这里 是 KL 惩罚权重,对应代码里的 kl_coef。DeepSeekMath 实验使用 DeepSeekMath 式(3)、式(4)与实验设置

126    token_ratio = torch.exp(new_logprobs - old_logprobs)
127    token_advantages = advantages.unsqueeze(-1)
128    unclipped = token_ratio * token_advantages
129    clipped_ratio = torch.clamp(token_ratio, 1.0 - clip_eps, 1.0 + clip_eps)
130    clipped = clipped_ratio * token_advantages
131 
132    # DeepSeekMath 式 (4):D_KL(policy || ref) 的逐 token 无偏正值估计
133    log_ratio_ref = ref_logprobs - new_logprobs
134    per_token_kl = torch.exp(log_ratio_ref) - log_ratio_ref - 1.0
135 
136    per_token_objective = torch.minimum(unclipped, clipped) - kl_coef * per_token_kl
137    per_response_objective = masked_sequence_mean(
138        per_token_objective,
139        completion_mask,
140    )
141    loss = -per_response_objective.mean()
142 
143    policy_loss = -masked_sequence_mean(
144        torch.minimum(unclipped, clipped),
145        completion_mask,
146    ).mean()
147    approx_kl = masked_sequence_mean(per_token_kl, completion_mask).mean()
148    metrics = {
149        "loss": loss.detach(),
150        "policy_loss": policy_loss.detach(),
151        "approx_kl": approx_kl.detach(),
152        "mean_ratio": masked_sequence_mean(token_ratio, completion_mask).mean().detach(),
153    }
154    return loss, metrics

代码对应关系:

  • log_ratio_ref = ref_logprobs - new_logprobs:逐 token 实现
  • per_token_kl = exp(log_ratio_ref) - log_ratio_ref - 1:逐 token 实现
  • per_token_objective = minimum(...) - kl_coef * per_token_kl:先在同一个 token 上合并裁剪目标和 KL。
  • masked_sequence_mean:用回答掩码排除 prompt 和 EOS 后的 padding,再执行原论文的

一次完整训练的七步

把所有步骤连起来,GRPO 的一次训练就是:

  1. 对每个 prompt 采样 个回答。
  2. 用规则或奖励函数给每个回答打分。
  3. 在同一个 prompt 的组内计算
  4. 保留回答 token 的 log probability,逐 token 算
  5. 对每个 token 分别执行 PPO-style clip,并计算 Reference KL。
  6. 每段回答先按有效 token 求平均,再对组内回答求平均。
  7. 反向传播,只更新 Policy。

完整的 GRPO 数据流如下图:

Mermaid diagram

GRPO 训练实验 与 GSM8K + 规则奖励

公式讲完后,看一次真实的 GRPO 训练。本节用一个最小可跑的实验:在 GSM8K 上用规则奖励训练 Qwen2.5-1.5B。

为什么不需要 RM

GSM8K 是一个包含 8500 道小学数学应用题的数据集,每道题都有明确的数值答案。这恰好是一个有"客观正确答案"的场景——不需要 RM,直接用规则判断答案是否正确:

  • 答案正确:
  • 格式规范(有清晰的推理步骤):
  • 答案错误:
python
# 1. 规则奖励函数(不需要 RM!)
import re

def rule_based_reward(prompt: str, response: str, ground_truth: str) -> float:
    reward = 0.0
    # 格式分:检查 \boxed{...}
    if re.search(r'\\boxed\{[^}]+\}', response):
        reward += 0.5
    # 答案分:提取最终答案并比较
    answer_match = re.search(r'\\boxed\{([^}]+)\}', response)
    if answer_match:
        model_answer = answer_match.group(1).strip()
        try:
            if abs(float(model_answer) - float(ground_truth)) < 0.01:
                reward += 1.0
        except ValueError:
            if model_answer == ground_truth:
                reward += 1.0
    return reward

# 测试
prompt = "Janet 的鸡蛋盒子每天能装 16 个鸡蛋。她每天早上吃 3 个,下午用 4 个烤松饼。她每周能卖多少个鸡蛋?"
good = "首先计算每天剩余的鸡蛋数:16 - 3 - 4 = 9 个\n每周有 7 天,所以每周能卖:9 × 7 = 63 个\n\\boxed{63}"
bad = "我觉得大概能卖 50 个左右吧。\\boxed{50}"
print(rule_based_reward(prompt, good, '63'))  # 1.5
print(rule_based_reward(prompt, bad, '63'))   # 0.5

注意这里的关键区别:不需要训练任何 RM,规则就是裁判。数学题有标准答案,直接比较就行。这种"可验证奖励"正是 RLVR 的核心思想。

在手写代码地图中,奖励函数对应的是 [B]。它只接收回答和标准答案,返回一个标量奖励:

 59# [B] 规则奖励:数学答案正确、格式规范就给分
 60def rule_reward(response, ground_truth):
 61    reward = 0.0
 62    boxed = re.search(r"\\boxed\{([^}]+)\}", response)
 63 
 64    if boxed:
 65        reward += 0.5
 66        if boxed.group(1).strip() == str(ground_truth).strip():
 67            reward += 1.0
 68 
 69    return reward
 70 
 71 
 72def score_responses(responses, ground_truths, group_size=8, device="cpu"):
 73    rewards = []
 74    for i, response in enumerate(responses):
 75        prompt_id = i // group_size
 76        rewards.append(rule_reward(response, ground_truths[prompt_id]))
 77    return torch.tensor(rewards, dtype=torch.float32, device=device)

运行 GRPO 训练

我们使用 trl 库提供的 GRPO 实现。和 PPO 相比,GRPO 不需要 Critic 模型:

python
# 2. GRPO 训练代码(简化示意)
from trl import GRPOTrainer, GRPOConfig
from transformers import AutoModelForCausalLM, AutoTokenizer
from datasets import load_dataset

model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-1.5B-Instruct")
tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2.5-1.5B-Instruct")
if tokenizer.pad_token is None:
    tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "left"

paper_grpo_options = dict(
    beta=0.04,                         # 原论文的 KL 系数
    epsilon=0.2,
    num_iterations=1,                  # 每批采样后只更新一次
    scale_rewards="group",             # 同一问题的组内标准化
    importance_sampling_level="token", # 逐 token ratio
    loss_type="grpo",                  # 每段回答先除以自身长度
)
# 仓库固定的 TRL 0.24 尚无此选项;新版 TRL 才需要显式关闭它。
if "use_bias_correction_kl" in GRPOConfig.__dataclass_fields__:
    paper_grpo_options["use_bias_correction_kl"] = False

config = GRPOConfig(
    output_dir="./grpo_gsm8k",
    num_generations=8,                 # 教学规模;原论文为 64
    per_device_train_batch_size=8,     # 全局 batch 必须能被组大小整除
    max_completion_length=1024,
    learning_rate=1e-6,                # 原论文实验设置
    num_train_epochs=1,
    **paper_grpo_options,
)

gsm8k = load_dataset("openai/gsm8k", "main")
trainer = GRPOTrainer(
    model=model,
    args=config,
    train_dataset=gsm8k["train"],
    reward_funcs=[rule_based_reward],  # 直接传入规则奖励函数
    processing_class=tokenizer,
)

trainer.train()  # 开始训练——不需要 Critic,不需要 RM
trainer.save_model("./grpo_gsm8k/final_model")

这里把算法口径和实验规模分开处理。num_generations=8 是为了降低教学实验的显存需求;DeepSeekMath 的正式实验设置为每个问题采样 64 个回答、batch size 1024、学习率 、KL 系数 0.04,并在每次探索后只更新一次。其余显式参数用于防止 TRL 的新版默认值把训练切换到 DAPO、BNPO、sequence-level importance sampling 或带偏差修正的 KL。仓库固定的 TRL 0.24 配置源码还没有 use_bias_correction_kl当前 main 分支配置新增了该字段且默认开启,所以示例只在检测到字段时将它关闭。TRL 的配置说明也提示原始 loss_type="grpo" 可能带来长度偏差,并推荐后续变体;这不改变本节复现原论文式(3)的选择。

如果把 GRPOTrainer 内部最关键的训练步骤摊开,就是"先组采样,再打分,再算优势,再更新策略":

197# [G] 训练步骤:采样、打分、组内归一化、再反向传播
198def train_step(
199    policy_model,
200    ref_model,
201    optimizer,
202    tokenizer,
203    prompts,
204    ground_truths,
205    group_size=8,
206):
207    responses, _, batch = sample_groups(
208        policy_model,
209        tokenizer,
210        prompts,
211        group_size,
212    )
213    rewards = score_responses(
214        responses,
215        ground_truths,
216        group_size,
217        policy_model.device,
218    )
219    advantages = group_advantages(rewards, group_size)
220 
221    loss, metrics = grpo_loss(policy_model, ref_model, batch, advantages)
222    optimizer.zero_grad()
223    loss.backward()
224    optimizer.step()
225    return metrics
226 
227 
228# [H] GRPO 训练循环:每轮都在线生成新回答
229def train_grpo(policy_model, ref_model, optimizer, tokenizer, dataloader):
230    ref_model.eval()
231    for prompts, ground_truths in dataloader:
232        metrics = train_step(
233            policy_model,
234            ref_model,
235            optimizer,
236            tokenizer,
237            prompts,
238            ground_truths,
239        )
240        print(
241            "loss=",
242            float(metrics["loss"]),
243            "kl=",
244            float(metrics["approx_kl"]),
245        )

推理步骤的变化

GRPO 训练最令人兴奋的观察是模型推理方式的变化:

训练前(直接猜答案):

题目:小明有 15 个苹果,给了小红 3 个,又给了小刚 5 个,还剩多少个?
回答:我觉得还剩 7 个。\boxed{7}

训练后(展示推理过程):

题目:小明有 15 个苹果,给了小红 3 个,又给了小刚 5 个,还剩多少个?
回答:
让我一步一步算:
- 小明一开始有 15 个苹果
- 给了小红 3 个:15 - 3 = 12
- 又给了小刚 5 个:12 - 5 = 7
- 所以还剩 7 个
\boxed{7}

模型从"直接猜答案"变成了"先列算式再计算"——这不是我们教它的,而是模型在 GRPO 训练过程中自己"领悟"出来的。因为展示推理步骤能提高答案正确率(拿到更高的规则奖励),所以 GRPO 的优化压力自然地选择了这条路径。

Mermaid diagram

实验对比与参数调优

显存占用对比

模型大小PPO 显存(4 模型)GRPO 显存(2 模型)节省比例
1.5B~24 GB~14 GB~42%
7B~80 GB~48 GB~40%
14B~160 GB~96 GB~40%
70B~640 GB~384 GB~40%

GRPO 省掉了 Critic(和 Actor 同等规模)和 RM 两个模型,通常能减少 30-40% 的显存占用。在实际工程中,这意味着原本需要 8 张 A100 的训练任务,现在 5 张就够了。

组内方差的演化

GRPO 的核心创新是用组内归一化替代 Critic。在训练初期,同一个问题的 8 个回答质量差异很大(方差高)。随着训练推进,组内回答质量趋于一致(方差降低),大部分回答都能答对。

训练初期(Episode 10):
  问题 "15 - 3 - 5 = ?" 的 8 个回答:[3, 7, 12, 7, 15, 7, 8, 10]
  组内方差:高(答案五花八门)
  归一化优势:[−1.2, +0.1, +0.8, +0.1, +1.5, +0.1, −0.3, +0.6]

训练中期(Episode 100):
  同一问题的 8 个回答:[7, 7, 7, 8, 7, 7, 7, 7]
  组内方差:低(大部分答对了)
  归一化优势:[0, 0, 0, −0.5, 0, 0, 0, 0]

训练后期(Episode 300):
  同一问题的 8 个回答:[7, 7, 7, 7, 7, 7, 7, 7]
  组内方差:接近零(全部答对)
  归一化优势:全部接近零 → 无梯度信号

当组内方差降为零时,优势全部为零,没有梯度信号了——模型在这个问题上"毕业"了。这正是我们想要的行为:训练信号自然地转移到还没掌握的题目上。

k 值的选择

k(组大小)是 GRPO 最关键的超参数,它直接影响组内归一化的质量:

k 值采样成本归一化质量适用场景
2低(每个问题只采 2 次)差(均值和标准差不稳定)快速验证
4中等一般资源有限时
8较高良好默认推荐
16很好(统计量更稳定)追求上限
64很高极好大规模训练
python
# GRPO 组内归一化的简单实现
import numpy as np

def grpo_group_normalize(rewards: list[float]) -> list[float]:
    rewards = np.array(rewards, dtype=float)
    mean, std = rewards.mean(), rewards.std()
    if std < 1e-8:
        return np.zeros_like(rewards)
    return (rewards - mean) / std

# 8 个回答的奖励
rewards = [1.5, 0.0, 1.5, 0.0, 1.0, 1.5, 0.5, 1.5]
advantages = grpo_group_normalize(rewards)
# 归一化优势: [ 0.89 -1.48  0.89 -1.48  0.10  0.89 -0.69  0.89]
# 均值: 0.9375, 标准差: 0.634
思考题:GRPO 的组内归一化在什么情况下会失效?
  1. k 太小 时均值和标准差极不稳定,统计量不可靠。
  2. 奖励分布偏斜:大部分回答得零分时,少数高分回答主导梯度信号。
  3. 所有回答质量相同:方差为零,优势全部为零,无梯度信号——即训练后期"毕业"现象。
  4. 奖励信号不连续:只有 0/1 两个值时,归一化后的优势分布是离散的,梯度信号不够精细。

GRPO 通过 DAPO 的"动态采样"改进来缓解这些问题——过滤掉模型已经答对的题目,只保留有梯度信号的样本。

GRPO 与 PPO 全面对比

组件PPOGRPO
基线(Critic)独立的 网络组内均值
优势计算 或 GAE
模型数量4 个(Actor + Critic + Ref + RM)2 个(Actor + Ref)
裁剪机制PPO Clip同样的 PPO Clip
采样方式在线交互组采样(每个 prompt 采 k 个)
显存低 30-40%
基线质量依赖 Critic 训练质量依赖组大小
基线更新速度需要重新训练 Critic自动随 batch 更新

值得注意的是,GRPO 继承了 PPO 的裁剪机制,但没有继承 GAE。原因是 GRPO 的奖励通常只在序列末尾给出一个信号(答对/答错),而不是每个 token 都有奖励。在这种情况下,GAE 的多步 TD 退化为单步,和直接用最终奖励减去均值没有本质区别。

GRPO 通过组内归一化优雅地解决了 Critic 的问题。但这只是第一步——在策略端,DeepSeek-R1-Zero 证明了不需要 SFT 也能做纯 RL 训练,DAPO 进一步优化了 GRPO 的工程效率。让我们看看这些前沿进展——DeepSeek-R1 与 DAPO

现代强化学习实战课程