GRPO 的核心 20 行:优势、ratio、clip、KL
整个算法真正的实现只有二十行。这一节逐行写出来,并用打印验证每个中间量。
学完这节你能做到
- 写出组内归一化的优势计算,并处理 std 为 0 的情况
- 写出带 clip 的策略损失,并解释每个符号对应 L1 的哪个概念
- 写出低方差 KL 估计,并知道 KL 系数设 0 意味着什么
整个算法只有二十行
前面两节做的都是准备工作。GRPO 本身——优势、ratio、clip、KL——加起来大约二十行 PyTorch。
这一节逐行写出来,并且每一步都给一个可以打印验证的中间量。RL 的 bug 大多不报错,只能靠对中间量的预期来发现。
第一步:算 logprob
给定一批序列,取出实际采样出的那些 token 的对数概率。
# scripts/grpo_scratch/grpo.py
import torch
import torch.nn.functional as F
def sequence_logprobs(model, input_ids, attention_mask, completion_mask):
"""
返回每个 completion token 的 logprob,形状与 completion_mask 相同。
"""
out = model(input_ids=input_ids, attention_mask=attention_mask)
logits = out.logits # [B, T, V]
# 关键的错一位:位置 t 的 logits 预测的是位置 t+1 的 token
logits = logits[:, :-1, :] # [B, T-1, V]
targets = input_ids[:, 1:] # [B, T-1]
mask = completion_mask[:, 1:] # [B, T-1]
logprobs = torch.log_softmax(logits.float(), dim=-1)
token_logprobs = torch.gather(
logprobs, dim=-1, index=targets.unsqueeze(-1)
).squeeze(-1) # [B, T-1]
return token_logprobs, mask
位置 t 的 logits 是「看过 0..t 之后对 t+1 的预测」。所以要拿位置 t 的 token 的概率,得用位置 t-1 的 logits。
代码里就是 logits[:, :-1] 配 input_ids[:, 1:]。
验证方法:拿一条已知序列,用同一个模型算两次 logprob,然后手工用 model(...) 逐 token 检查其中一个位置。或者更简单:见下面的 ratio 检查。
logits.float() 这一步不能省。bf16 的有效位数只有 8 位,log_softmax 在 bf16 下的数值误差会直接污染 ratio。
代价是这一步的显存会翻倍(float32 的 logits)——这也是上一阶段显存账本里 logits 那一行为什么按 2 份算的原因。
第二步:算优势
def group_advantages(rewards: torch.Tensor, group_size: int) -> torch.Tensor:
"""
rewards: [B],其中 B = n_prompts * group_size,同组的样本连续排列
返回: [B]
"""
r = rewards.view(-1, group_size) # [n_groups, G]
mean = r.mean(dim=1, keepdim=True)
std = r.std(dim=1, unbiased=False, keepdim=True) # 总体标准差,除 G
adv = (r - mean) / (std + 1e-4)
# 关键:一组全对或全错时 std=0,显式归零,不要让它变成 NaN 或巨大值
adv = torch.where(std > 1e-4, adv, torch.zeros_like(adv))
return adv.view(-1)
三个细节:
unbiased=False:用总体标准差(除 G),和主流实现一致。用unbiased=True(除 G−1)不算错,但和别人的数字对不上。std > 1e-4的显式归零:这是最重要的一行。没有它,全对组会算出0 / 1e-4 = 0(还好),但如果 std 是 1e-8 这种极小值,就会算出巨大的优势值,一步把模型炸掉。- 同组样本必须连续排列:
view(-1, group_size)依赖这个假设。上一节 rollout 返回的顺序天然满足(vLLM 的n=G会把同一个 prompt 的 G 条放在一起),但如果你中间做了 shuffle,这里就错了。
adv = group_advantages(rewards, group_size=8)
print(adv.view(-1, 8).sum(dim=1)) # 每一行都应该接近 0如果不是 0,说明分组错了(大概率是样本顺序被打乱)。
第三步:ratio
ratio = torch.exp(logp_new - logp_old)
一行。但有两件事要说清。
logp_old 从哪来? 是采样时那个模型给出的 logprob。两种拿法:
- 方案 A:采样后立刻用训练模型(此时参数还没更新)跑一次前向,存下来。多一次前向的开销,但准确。
- 方案 B:直接用 vLLM 返回的 logprobs。省一次前向,但 vLLM 和 transformers 的数值实现有细微差异,会引入偏差。
TRL 默认开启 Truncated Importance Sampling(vllm_importance_sampling_correction),slime 有 --use-tis 选项,就是在修正「生成引擎与训练引擎的概率不一致」。
手写实现建议用方案 A:多花一次前向,换来一个干净的 ratio。等你对齐了再考虑省这一步。
如果每批数据只更新一次(这是最常见的配置),那么在第一次也是唯一一次前向时,新旧参数完全相同,所以:
assert torch.allclose(ratio, torch.ones_like(ratio), atol=1e-3), ratioratio 必须恒等于 1。
这一个断言能同时抓住三类 bug:错一位对齐错了、mask 错了、logp_old 来源不对。先让这个断言过,再往下写。
第四步:clip
def policy_loss(logp_new, logp_old, advantages, mask, eps_low=0.2, eps_high=0.2):
"""
logp_new/old: [B, T] advantages: [B] mask: [B, T]
"""
# 序列级优势广播到 token 级:「按序列打分,按 token 求梯度」
adv = advantages.unsqueeze(1) # [B, 1] → 广播
ratio = torch.exp(logp_new - logp_old) # [B, T]
clipped = torch.clamp(ratio, 1 - eps_low, 1 + eps_high)
# min 让它对「好回答提太多」和「坏回答压太狠」都生效
per_token = -torch.min(ratio * adv, clipped * adv) # 负号:要最小化
# 统计有多少 token 被截断了,作为监控指标
with torch.no_grad():
clip_frac = (((ratio * adv) > (clipped * adv)).float() * mask).sum() / mask.sum()
return per_token, clip_frac
eps_low 和 eps_high 分开是有用的:非对称 clip(eps_high 比 eps_low 大,如 0.2 / 0.28)能鼓励探索,DAPO 提出、slime 的入门配方里就是这么设的。
第五步:KL
朴素做法是 logp_new - logp_ref,但它的方差很大,而且可能为负(KL 应该恒非负)。用 k3 估计:
def kl_k3(logp_new, logp_ref):
"""
Schulman 的 k3 低方差估计:exp(d) - d - 1,其中 d = logp_ref - logp_new
性质:恒非负;期望等于真实 KL;方差远小于朴素差值
"""
d = logp_ref - logp_new
return torch.exp(d) - d - 1
对任意实数 d,有 exp(d) ≥ d + 1(指数函数在 d=0 处的切线是 d+1,而 exp 是凸的,永远在切线之上)。所以 exp(d) - d - 1 ≥ 0,且只在 d = 0 时取等。
朴素估计 logp_new - logp_ref 的期望也是 KL,但单个样本上可以是负数——画出来的曲线会在 0 附近抖动,看不出趋势。
第六步:mask 与归一化 —— 这里有个经典陷阱
把 per-token 的损失汇总成一个标量,有两种做法:
# 做法 A:per-token 平均(所有 token 平权)
loss = (per_token * mask).sum() / mask.sum()
# 做法 B:per-sequence 平均(每条序列先内部平均,再对序列平均)
loss = ((per_token * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1)).mean()
两者不等价,而且差别有实际后果:
| 做法 A(token 平权) | 做法 B(序列平权) | |
|---|---|---|
| 长回答的权重 | 更大(token 多) | 与短回答相同 |
| 已知副作用 | 变相鼓励写长 | 长短公平 |
| 谁在用 | DAPO 推荐(配长度惩罚) | 原始 GRPO 论文 |
如果按 token 平权,一条 500 token 的回答对损失的贡献是 50 token 回答的 10 倍。当长回答恰好 reward 略高时,模型会得到「写长有好处」的信号,然后一路写到 max_tokens。
L4 会讲长度爆炸的完整诊断。这里先记住:这两种写法要显式选一个,并且知道自己选了哪个。 框架里通常有个开关(slime 叫 --calculate-per-token-loss)。
新手建议先用做法 B(序列平权),它的行为更符合直觉。
拼起来
def grpo_loss(
logp_new, logp_old, logp_ref, advantages, mask,
eps_low=0.2, eps_high=0.28, beta=0.0,
):
per_token, clip_frac = policy_loss(
logp_new, logp_old, advantages, mask, eps_low, eps_high
)
if beta > 0.0 and logp_ref is not None:
per_token = per_token + beta * kl_k3(logp_new, logp_ref)
# 序列平权
seq_loss = (per_token * mask).sum(dim=1) / mask.sum(dim=1).clamp(min=1)
loss = seq_loss.mean()
with torch.no_grad():
kl = (kl_k3(logp_new, logp_ref) * mask).sum() / mask.sum() \
if logp_ref is not None else torch.tensor(0.0)
return loss, {
"loss": loss.item(),
"clip_frac": clip_frac.item(),
"kl": kl.item(),
"adv_std": advantages.std().item(),
"ratio_mean": torch.exp(logp_new - logp_old).mean().item(),
}
逐个中间量检查
写完不要直接开训。按这张表逐项验证:
| 中间量 | 期望值 | 不对说明什么 |
|---|---|---|
ratio 初值 | 恒等于 1.0 | 错一位、mask 错、logp_old 来源不对 |
每组 advantages 之和 | 0 | 分组错了(样本顺序被打乱) |
advantages 的量级 | 大致在 −2 ~ +2 | std 归零处理漏了 |
kl | ≥ 0 | 用了朴素估计,或参考模型加载错了 |
mask.sum() | 等于 completion 的总 token 数 | mask 覆盖到 prompt 了 |
clip_frac 首步 | 0 | ratio 不是 1(同上) |
loss 首步 | 接近 0 | ratio=1 且优势零均值时,min 项相消 |
- 1.确认 ratio 初值恒为 1
- 2.确认优势零均值且退化组被归零
- 3.确认 KL 估计恒非负
- 4.确认首步 loss 接近 0、clip_frac 为 0
算法自检演练。这些检查全部通过之后才能开训。 goals 看目标,hint 要提示。
手写 GRPO 时最有价值的一个调试断言是什么?
`group_advantages` 里为什么要写 `torch.where(std > 1e-4, adv, 0)`?
per-token 平均与 per-sequence 平均的损失归一化,差别的实际后果是什么?
这节课的落点
- logprob:
logits[:, :-1]配input_ids[:, 1:],log_softmax前转 float32 - 优势:总体标准差(
unbiased=False),std 接近 0 时显式归零,同组样本必须连续 - ratio:
logp_old建议多跑一次前向拿(方案 A),初值必须恒为 1 - clip:
torch.min让它对两个方向都生效;非对称 eps(0.2/0.28)鼓励探索 - KL:用 k3 估计
exp(d)-d-1,恒非负且方差小 - 归一化:token 平权 vs 序列平权不等价,显式选一个;新手用序列平权
- 开训前把七项中间量逐个验证过
延伸资料
- ·DeepSeekMath 论文(GRPO 原始出处) ↗
- ·rlforge 配套脚本
scripts/grpo_scratch/grpo.py