WARNING
🧪 Beta公测版本提示:教程主体已完成,正在优化细节,欢迎大家提Issue反馈问题或建议。
GRPO:组相对策略优化 — demo.py 代码详解
运行方式
cd docs/nn-decision/rl/grpo/code
python demo.pyCPU 即可。36 道
代码逐段详解
第1步:导入与题库
N_A, N_B = 6, 6
N_ANS = N_A + N_B + 1 # 0 .. 12
N_Q = N_A * N_B # 36
GROUP = 8
CLIP_EPS = 0.2
KL_BETA = 0.02N_ANS=13:和到 的合法和一致,动作空间刚好盖住所有正确答案,没有「超出范围」的干扰项。 GROUP=8:同一 prompt 采 8 条。太小则估不准;太大则每步太贵。真 GRPO 往往 。 CLIP_EPS:和 PPO 同一套盒子。KL_BETA:把拴在参考策略旁,对应正文里「别离 SFT / 旧策略太远」。这里参考是初始化那一刻冻结的 logits。 Categorical:离散答案上的策略。.sample()/.log_prob/kl_divergence都靠它。
def build_bank():
for a in range(N_A):
for b in range(N_B):
qs.append(a * N_B + b)
gold.append(a + b)题目编号 AnswerPolicy 直接用整数下标查表。np.array 方便和 argmax 比正确率。
第2步:AnswerPolicy — 「按题查表的小 LLM」
class AnswerPolicy(nn.Module):
def __init__(self, n_q=N_Q, n_ans=N_ANS):
super().__init__()
self.logits = nn.Parameter(torch.zeros(n_q, n_ans))
def dist(self, q_idx):
return Categorical(logits=self.logits[q_idx])nn.Parameter:一张(36, 13)的表,每道题一组 logits。全零起步 = 对 13 个答案均匀。没有 MLP,梯度直接改这张表——对应「小 LLM 的最后一层分类头」,把表示学习剥掉。self.logits[q_idx]:q_idx是 Pythonint时取出长度为 13 的 1 维张量;Categorical在这 13 维上做 Softmax。- 必须
super().__init__():否则这张表进不了parameters(),SGD更新不到。
accuracy:argmax(dim=-1) 贪心取每题最可能的答案,和 gold 比均值。评估不算梯度(no_grad)。
第3步:group_advantages — 组内 z-score,没有
PPO 的
r = np.asarray(rewards, dtype=np.float64)
std = r.std()
if std < 1e-6:
return np.zeros_like(r, dtype=np.float32)
return ((r - r.mean()) / (std + eps)).astype(np.float32)- 比组内平均好 →
,提高这些答案的概率;差的压低。这就是「相对」:全员 0 分或全员 1 分时,没有谁比平均更好。 std < 1e-6:方差塌掉,整组优势为 0。后面grpo_step直接跳过backward——没信号就别更新,避免附近的噪声梯度。 float64算、float32回:和 PPO 里 GAE 的写法一样,先用更宽的浮点减均值。- 没有
compute_gae,没有self.v:长思维链上训又贵又不稳,组采样反正都要做,基线就用组均值。
同一组数,两种优势差在哪。 设
和 PPO 的 GAE 对照:那边 min 与 clamp 仍按样本各自生效:好答案的
第4步:grpo_step — 先用旧策略采样,再 clip + KL
dist_old = Categorical(logits=policy.logits[q_idx].detach())
for _ in range(group):
a = dist_old.sample()
answers.append(int(a.item()))
logp_old.append(dist_old.log_prob(a))
rewards.append(1.0 if answers[-1] == int(gold) else 0.0).detach():采样时的 logits 当常数。验证器是规则:猜中gold得 1,否则 0。这就是示意图里的「可验证奖励」,没有奖励模型。int(a.item()):张量 → Python 整数,才能和gold比、才能再torch.tensor(answers)。
adv = group_advantages(rewards)
if np.allclose(adv, 0):
return float(np.mean(rewards)), Trueallclose 把「全零优势」判定为跳过。返回的 True 累进 skipped,画右图。
有方差才更新:
dist = policy.dist(q_idx)
new_lp = dist.log_prob(answers_t)
ratio = torch.exp(new_lp - old_lp)
surr1 = ratio * adv_t
surr2 = torch.clamp(ratio, 1.0 - CLIP_EPS, 1.0 + CLIP_EPS) * adv_t
clip_loss = -torch.min(surr1, surr2).mean()和 PPO 章的 ppo_update 同一套:clamp 进 min 后取负(最大化 adv_t 的来源:这里是组内 z-score,不是 GAE。
old_lp = torch.stack(logp_old).detach():把 log_prob 仍连着 dist_old;再 detach 一次,保证比率的旧半边没有梯度。
ref = Categorical(logits=ref_logits[q_idx])
kl = torch.distributions.kl.kl_divergence(dist, ref)
loss = clip_loss + KL_BETA * kl离散 Categorical 的 KL 有闭式,不必 Monte Carlo。ref_logits 在 train 开头 detach().clone(),整段训练不更新——相当于「SFT 参考」。没有这项,clip 仍限制一步跨多远,但多步之后可以漂到只会背当前这几题。
opt 是 SGD 不是 Adam:表很小,固定步长更直观。zero_grad → backward → step 一次,对应论文里对一组样本的一拍(玩具没有再套 K 个 epoch)。
第5步:reinforce_step — 绝对 0/1,无基线、无 clip
for _ in range(group):
a = dist.sample()
r = 1.0 if int(a.item()) == int(gold) else 0.0
loss = loss - dist.log_prob(a) * r
loss = loss / group同一组采样次数,但权重是原始 clamp,一步可以离开旧策略很远。
注意:这里 dist 没有先 detach 再采样,采样与损失共用当前图,是经典 REINFORCE;GRPO 则显式分开
第6步:train — 每步随机抽一题
ref_logits = policy.logits.detach().clone()
opt = optim.SGD(policy.parameters(), lr=LR)
i = int(np.random.randint(0, len(qs)))120 步,每步均匀抽一道题。kind 用不同 seed 偏移。skip_frac = skipped / (step+1) 是累计跳过比例:训练后期更多题被学对,一组
draw_group / draw_vs_ppo / draw_verifier / draw_roadmap 是正文框图。训练图:左正确率,右跳过比例。
关键概念速查表
| 概念 | 数学 / 直觉 | 代码 |
|---|---|---|
| 组采样 | 同一 | for _ in range(group) |
| 组优势 | group_advantages | |
| 零方差 | 全对或全错 → 没梯度 | std<1e-6 / allclose → skip |
| 裁剪 | 与 PPO 同一 | clamp + min |
| KL | kl_divergence(dist, ref) | |
| 参考策略 | 初始化冻结的 logits | ref_logits = ...clone() |
| REINFORCE | 绝对 | reinforce_step |
| 可验证奖励 | 对=1 错=0 | answers[-1] == gold |
源码位置
clone 后打开(相对仓库根目录):
docs/nn-decision/rl/grpo/code/demo.py