RRLforge
完整主线 路线14 / 27 · L2 手锻退出路线
实验预计 50 分钟

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 检查。

!log_softmax 前转 float32

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)

三个细节:

  1. unbiased=False:用总体标准差(除 G),和主流实现一致。用 unbiased=True(除 G−1)不算错,但和别人的数字对不上。
  2. std > 1e-4 的显式归零:这是最重要的一行。没有它,全对组会算出 0 / 1e-4 = 0(还好),但如果 std 是 1e-8 这种极小值,就会算出巨大的优势值,一步把模型炸掉。
  3. 同组样本必须连续排列view(-1, group_size) 依赖这个假设。上一节 rollout 返回的顺序天然满足(vLLM 的 n=G 会把同一个 prompt 的 G 条放在一起),但如果你中间做了 shuffle,这里就错了。
打印验证:优势总和应该恒为 0
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 的数值实现有细微差异,会引入偏差。
i这个「细微差异」工业框架专门处理过

TRL 默认开启 Truncated Importance Samplingvllm_importance_sampling_correction),slime 有 --use-tis 选项,就是在修正「生成引擎与训练引擎的概率不一致」。

手写实现建议用方案 A:多花一次前向,换来一个干净的 ratio。等你对齐了再考虑省这一步。

最重要的一个调试检查点

如果每批数据只更新一次(这是最常见的配置),那么在第一次也是唯一一次前向时,新旧参数完全相同,所以:

assert torch.allclose(ratio, torch.ones_like(ratio), atol=1e-3), ratio

ratio 必须恒等于 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_loweps_high 分开是有用的:非对称 clipeps_higheps_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
i为什么 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 论文
×做法 A 是「长度爆炸」的一个成因

如果按 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 ~ +2std 归零处理漏了
kl≥ 0用了朴素估计,或参考模型加载错了
mask.sum()等于 completion 的总 token 数mask 覆盖到 prompt 了
clip_frac 首步0ratio 不是 1(同上)
loss 首步接近 0ratio=1 且优势零均值时,min 项相消
rl@forge
目标 0/4
  1. 1.确认 ratio 初值恒为 1
  2. 2.确认优势零均值且退化组被归零
  3. 3.确认 KL 估计恒非负
  4. 4.确认首步 loss 接近 0、clip_frac 为 0
算法自检演练。这些检查全部通过之后才能开训。
goals 看目标,hint 要提示。
[rl@forge ~/rlforge]$
help 查看用法 · 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 序列平权不等价,显式选一个;新手用序列平权
  • 开训前把七项中间量逐个验证过

延伸资料