大模型强化学习新范式:从 PPO 到 GRPO(群体相对策略优化)数学推导与 PyTorch 实现
大模型强化学习新范式:从 PPO 到 GRPO(群体相对策略优化)数学推导与 PyTorch 实现

在探讨大语言模型(LLM)的对齐与后训练(Post-training)时,由 OpenAI 在 ChatGPT 中率先带火的 基于人类反馈的强化学习(RLHF via PPO - Proximal Policy Optimization) 曾长期被奉为行业标准基座。
然而,几乎所有在生产线上实际训练过千亿参数大模型 PPO 的算法工程师,都深受其带来的**“四模型显存黑洞”与“价值网络(Critic)训练崩塌”**折磨:
- 极其沉重的“4 模型显存霸凌”:在标准的 PPO 训练中,单台 GPU 服务器必须同时容纳 4 个巨大的模型副本——Actor(待优化的当前策略网络)、Critic(状态价值评估网络)、Reference Model(冻结的初始参考基准) 以及 Reward Model(奖励打分模型)!即使使用 DeepSpeed ZeRO-3 极致分片,Critic 网络仍然占用了与 Actor 完全同等数量级的显存与通信带宽;
- Critic 网络的“价值估计深渊(Value Misestimation)”:在长文本推理(如 8000 个思考 Token)或复杂代码生成任务中,状态空间极其高维且高度非平稳。Critic 网络很难准确估计某个中间 Token 的绝对期望价值 $V(s)$。一旦 Critic 估计出现偏差,基于 GAE(Generalized Advantage Estimation)计算出的优势函数就会产生剧烈噪声,导致 Actor 策略梯度瞬间爆炸,训练直接走向不可逆的性能坍塌!
正是在这一痛点下,以 DeepSeekMath 与 DeepSeek-R1 为代表的 群体相对策略优化(Group Relative Policy Optimization - GRPO) 强势崛起,成为全球大模型推理强化学习领域的超级新星!
GRPO 是如何做到彻底抛弃 Critic 价值网络,实现显存减半且训练极其稳定的? 其“群体组内归一化相对优势”背后的数学推导严密性何在?
本文深入剖析 PPO vs GRPO 底层数学与显存架构代差、GRPO 目标函数严格推导,并给出生产级 PyTorch GRPO 损失计算与训练循环实现代码。
一、传统 PPO 架构 vs 现代化 GRPO 强化学习全景深度对比矩阵
| 强化学习对比维度 | 传统近端策略优化 (PPO 范式) | 群体相对策略优化 (GRPO / DeepSeek 范式) | 生产级核心收益 |
|---|---|---|---|
| 模型常驻显存开销 | 极高 (需同时常驻 Actor + Critic + Ref + Reward 4 个模型) | 🏆 极低 (彻底移除 Critic 模型,仅保留 Actor + Ref) | GPU 显存开销直接立减 50%! |
| 基线价值估计方式 | 依赖独立的 Critic 神经网络逼近绝对状态价值 $V(s)$ | 🏆 纯无参估计 (基于同 Prompt 采样 $G$ 个回答的均值与方差) | 彻底消灭 Critic 网络的拟合误差与梯度崩溃 |
| 训练稳定性与鲁棒性 | 较差 (易受 GAE 方差和 Critic 学习率不匹配影响) | 极强 (组内自适应相对排序,梯度方差天然受控) | 复杂长思维链推理训练几乎零掉线崩塌 |
| 通信与调度复杂度 | 极重 (Actor 与 Critic 之间频繁发生跨节点梯度同步) | 轻量 (仅需在 Actor 内部完成同组采样的轻量 Gather) | 分布式训练 Step 耗时缩短 30% ~ 40% |
| 适用任务场景 | 传统主观闲聊对话与多维度偏好对齐 | 🏆 形式化数学证明、竞赛代码、复杂逻辑长思考模型 | 是打造类 o1 深度推理模型的绝对首选基座 |
二、PPO 四模型笨重拓扑 vs GRPO 极简群体归一化拓扑时序图
1. 传统 PPO 训练拓扑(显存黑洞)
[Prompt $q$] ──+──> [Actor 策略模型] ─────────> [生成单条回答 $o$] ──> [Reward Model] ──> [标量奖励 $r$]
| |
+──> [Critic 价值模型 (占50%显存!)] ──> [预测状态价值 $V(s)$] ───────────────+
|
v
[更新 Actor 与 Critic 参数] <─── [计算 GAE 优势函数: $\hat{A} = r - V(s)$ (方差极大!)] <─────+
2. GRPO 群体相对策略优化拓扑(极简高效)
[Prompt $q$] ──> [Actor 策略模型] ──> 并行采样一组 $G$ 个候选回答 $\{o_1, o_2, \dots, o_G\}$ (例如 $G=8$)
|
v (使用规则验证器/奖励函数打分)
[获得一组标量奖励 $\{r_1, r_2, \dots, r_G\}$]
|
v (🌟 纯数学组内归一化,零 Critic 网络参与!)
[计算均值 $\mu$ 与标准差 $\sigma$: $A_i = \frac{r_i - \mu}{\sigma}$]
|
v
[直接更新 Actor 参数: 奖励高于均值的正向增强,低于均值的反向抑制!]
三、GRPO 目标函数的严密数学推导
对于给定的问题输入 $q$,GRPO 策略网络 $\pi_\theta$ 首先生成一组输出候选集 $\mathcal{O} = {o_1, o_2, \dots, o_G}$。
奖励模型或规则判定函数为每个输出赋予奖励值 ${r_1, r_2, \dots, r_G}$。
1. 群体相对优势函数(Group Relative Advantage)
GRPO 不依赖价值网络,而是直接通过该组内的样本统计量计算每个回答的相对优势 $A_i$:
$$\mu = \frac{1}{G} \sum_{i=1}^{G} r_i, \quad \sigma = \sqrt{\frac{1}{G} \sum_{i=1}^{G} (r_i - \mu)^2 + \epsilon}$$
$$A_i = \frac{r_i - \mu}{\sigma}$$
- 当 $A_i > 0$ 时:代表回答 $o_i$ 的质量在该组中处于前列水平,模型增加生成这些 Token 的概率;
- 当 $A_i < 0$ 时:代表回答 $o_i$ 的质量低于群体平均,模型降低生成这些 Token 的概率。
2. 最终优化目标(带重要性采样与 KL 散度约束)
GRPO 的总损失函数定义为:
$$\mathcal{J}{\text{GRPO}}(\theta) = \mathbb{E} \left[ \frac{1}{G} \sum{i=1}^{G} \frac{1}{|o_i|} \sum_{t=1}^{|o_i|} \min \left( \frac{\pi_\theta(o_{i,t} | q, o_{i,<t})}{\pi_{\text{old}}(o_{i,t} | q, o_{i,<t})} A_i, ; \text{clip}\left(\frac{\pi_\theta}{\pi_{\text{old}}}, 1-\epsilon, 1+\epsilon\right) A_i \right) - \beta \mathbb{D}{\text{KL}}(\pi\theta | \pi_{\text{ref}}) \right]$$
其中:
$$\mathbb{D}{\text{KL}}(\pi\theta | \pi_{\text{ref}}) = \frac{\pi_{\text{ref}}(o_{i,t} | \dots)}{\pi_\theta(o_{i,t} | \dots)} - \log \frac{\pi_{\text{ref}}(o_{i,t} | \dots)}{\pi_\theta(o_{i,t} | \dots)} - 1$$
这一无偏估计形式(Schulman KL 近似)保证了策略在优化过程中不会偏离基准模型 $\pi_{\text{ref}}$ 太远,从而避免策略崩溃。
四、生产级 PyTorch GRPO 损失计算与训练循环实战代码
下面的代码展示了在单机多卡或最小复现环境中,如何用纯 PyTorch 实现 GRPO 组内归一化、Clipped 策略梯度损失以及 Token 级 KL 散度约束。
"""
grpo_training_core_sim.py
群体相对策略优化 (GRPO) 核心算法推导与 PyTorch 训练循环实战
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import List, Dict, Tuple
def compute_group_relative_advantages(rewards: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
"""
计算 GRPO 组内相对优势: A_i = (r_i - mean(r)) / (std(r) + eps)
输入 rewards 维度: [BatchSize, GroupSize]
返回 优势矩阵 维度: [BatchSize, GroupSize]
"""
mean = rewards.mean(dim=-1, keepdim=True)
std = rewards.std(dim=-1, keepdim=True)
advantages = (rewards - mean) / (std + eps)
return advantages
def compute_grpo_loss(
log_probs: torch.Tensor, # 当前模型输出 Log 概率: [B, G, SeqLen]
old_log_probs: torch.Tensor, # 旧模型采样 Log 概率: [B, G, SeqLen]
ref_log_probs: torch.Tensor, # 参考模型 Log 概率: [B, G, SeqLen]
advantages: torch.Tensor, # 组内相对优势: [B, G]
mask: torch.Tensor, # 有效 Token 掩码: [B, G, SeqLen]
clip_eps: float = 0.2,
beta_kl: float = 0.04
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
计算 GRPO 策略梯度损失与 KL 散度惩罚
"""
# 1. 计算重要性采样概率比率 Ratio = exp(log_pi - log_pi_old)
ratios = torch.exp(log_probs - old_log_probs)
# 2. 将优势函数扩展到序列维度: [B, G, 1]
adv_expanded = advantages.unsqueeze(-1)
# 3. 计算 Clipped Surrogate 损失
surr1 = ratios * adv_expanded
surr2 = torch.clamp(ratios, 1.0 - clip_eps, 1.0 + clip_eps) * adv_expanded
policy_loss = -torch.min(surr1, surr2) # 最大化目标转为最小化 Loss
# 4. 计算无偏 KL 散度估计: D_KL = exp(log_ref - log_pi) - (log_ref - log_pi) - 1
# 采用 Schulman 2020 严格非负低方差估计形式
log_ratio_ref = ref_log_probs - log_probs
kl_div = torch.exp(log_ratio_ref) - log_ratio_ref - 1.0
# 5. 组合总损失 (仅对有效生成的 Token 应用 Mask)
total_token_loss = policy_loss + beta_kl * kl_div
masked_loss = (total_token_loss * mask).sum() / (mask.sum() + 1e-8)
avg_policy_loss = (policy_loss * mask).sum() / (mask.sum() + 1e-8)
avg_kl = (kl_div * mask).sum() / (mask.sum() + 1e-8)
return masked_loss, avg_policy_loss, avg_kl
if __name__ == "__main__":
torch.manual_seed(2026)
print("=================================================================")
print("🔬 醍醐实验室:群体相对策略优化(GRPO)数学推导与 PyTorch 训练模拟")
print("=================================================================\n")
batch_size = 2 # 每次输入 2 个不同的 Prompt
group_size = 4 # 每个 Prompt 采样 G=4 个回答进行组内对比
seq_len = 10 # 生成的序列长度
# 1. 模拟环境给出的原始标量奖励 (如由代码编译器或数学判题器给出)
# 假设 Prompt 1 下回答 0 和 3 完全正确 (得分 1.0),回答 1 和 2 错误 (得分 0.0)
raw_rewards = torch.tensor([
[1.0, 0.0, 0.0, 1.0], # Prompt 1
[0.0, 0.5, 1.0, 0.0] # Prompt 2
])
# 2. 计算组内相对优势 (Group Relative Advantages)
advantages = compute_group_relative_advantages(raw_rewards)
print(f"1. [原始奖励矩阵 (Raw Rewards)]:\n{raw_rewards}\n")
print(f"2. [🌟 GRPO 计算出的组内相对优势 (Advantages)]:\n{advantages}\n")
print(" 💡 说明: 组内高于平均分的样本获得正向优势 (+),低于平均分的获得负向抑制 (-)\n")
# 3. 模拟各模型输出的前向 Log 概率
log_probs = torch.randn(batch_size, group_size, seq_len, requires_grad=True)
old_log_probs = log_probs.detach() + torch.randn(batch_size, group_size, seq_len) * 0.05
ref_log_probs = log_probs.detach() + torch.randn(batch_size, group_size, seq_len) * 0.02
mask = torch.ones(batch_size, group_size, seq_len) # 简化假设所有 Token 均有效
# 4. 执行 GRPO Loss 计算与反向传播
loss, p_loss, kl = compute_grpo_loss(
log_probs, old_log_probs, ref_log_probs, advantages, mask, clip_eps=0.2, beta_kl=0.04
)
loss.backward()
print(f"3. [GRPO 损失计算结果]:")
print(f" - 总损失 (Total Loss) : {loss.item():.6f}")
print(f" - 策略损失 (Policy Loss) : {p_loss.item():.6f}")
print(f" - KL 散度惩罚 (KL Div) : {kl.item():.6f}")
print(f" - 梯度范数 (Grad Norm) : {log_probs.grad.norm().item():.6f}")
print("\n🎉 验证成功:GRPO 在零 Critic 网络参数下,完美完成组内策略梯度反传!")
print("=================================================================")
五、GRPO 训练调优实战与落地避坑指南
在将 GRPO 应用于真实大模型后训练时,必须坚守以下四项实战工程红线:
- 群体大小(Group Size $G$)推荐设置在 $8 \sim 16$ 之间:
若 $G < 4$,组内方差估计极其不稳定,相对优势容易失真;若 $G > 32$,采样阶段的吞吐开销会成倍增加。实践证明 $G=8$ 是显存利用与统计鲁棒性的最佳甜点区; - 在纯逻辑/数学/代码任务中,优先使用“硬规则奖励(Rule-based Reward)”替代神经奖励模型:
对于 LeetCode 或 AIME 题目,直接通过 Python 沙箱执行单元测试(通过 $=1.0$,报错 $=0.0$),100% 消除奖励模型自身的幻觉与对抗攻击(Reward Hacking); - 严格监控组内同质化(Group Collapse)问题:
如果模型在训练后期生成多样性骤降,导致同一组内的 $G$ 个输出完全一模一样(奖励全为 $1.0$ 或全为 $0.0$),此时标准差 $\sigma = 0$,优势函数 $A_i$ 将全部退化为零导致梯度失效。必须维持适当的采样温度(如 $T = 0.7 \sim 0.9$)并适度调大 KL 散度惩罚系数 $\beta$。
通过彻底剔除笨重易崩溃的 Critic 价值网络,依靠纯粹而优雅的组内相对归一化优势估计,GRPO 为大模型深度推理与强化学习后训练铺设了一条高吞吐、极低显存与极其平稳的全新高速公路。
更多推荐


所有评论(0)