Skip to content
alarik.me
Go back

从 PPO 到 GRPO:LLM 强化学习的原理与 PyTorch 实现

LLM 的 RL 后训练(RLHF / 推理强化)常被当作黑盒:loss、ratio、advantage、KL 这些词人人都说,但从「为什么不能直接交叉熵」到「loss.backward() 那一刻到底发生了什么」之间的链条,很少有人完整讲清。本文把这条链条拆开:从策略梯度出发,推到 GRPO 的 PyTorch 实现,再下沉到 grad_fn、AdamW 优化器与训练显存的真实构成——每一层都对应到可以直接运行的 torch 代码。

Table of contents

Open Table of contents

一、为什么 LLM 的 RL 不能直接反向传播

标准 SFT 有 ground truth token 作为标签,交叉熵全程可导。RL 的奖励则要等整段文本生成结束才给出,且生成过程包含离散采样(top-p / multinomial),采样不可导,链式法则无法穿透到参数上。

所以 LLM RL 依赖策略梯度定理,用对数似然加权来估计梯度:

θEτπθ[R(τ)]=Eτπθ[t=1Tθlogπθ(ytx,y<t)A^t]\nabla_\theta \mathbb{E}_{\tau \sim \pi_\theta}[R(\tau)] = \mathbb{E}_{\tau \sim \pi_\theta}\left[\sum_{t=1}^{T} \nabla_\theta \log \pi_\theta(y_t \mid x, y_{<t}) \cdot \hat{A}_t\right]

其中 A^t\hat{A}_t 是优势函数(Advantage),衡量在状态 sts_t 下选择 yty_t 相比平均水平「好多少」。这个公式有一个直读:梯度顺着 logπθ\log \pi_\theta 流,奖励信号只作为系数——采样本身不出现在梯度路径里。

把生成过程建模为 MDP:

RL 概念LLM 对应物
状态 sts_t上下文 (x,y<t)(x, y_{<t})
动作 ata_t下一个 token yty_t
策略 πθ(atst)\pi_\theta(a_t \mid s_t)LLM 本身,输出词表上的分布
奖励 rr偏好对齐用 Reward Model;数学/代码推理用规则判题

二、PPO 的四模型架构

以标准 PPO 为例,系统涉及 4 个模型实体:

训练分两个阶段:

阶段 A —— Rollout(推理模式):Actor 自回归采样完整序列,记录每个 token 的 logπθold(ytst)\log \pi_{\theta_{old}}(y_t \mid s_t)(即 old_log_probs),计算标量奖励,并把序列喂给 Reference 模型取 logπref\log \pi_{ref}

阶段 B —— Training(开梯度):把 KL 惩罚折进 token 级奖励——除最后一个 token 外都是 β(logπθlogπref)-\beta(\log \pi_\theta - \log \pi_{ref}),末位 token 额外加上 R(x,y)R(x,y)。Critic 前向得到 V(st)V(s_t),算 TD 误差 δt=rt+γV(st+1)V(st)\delta_t = r_t + \gamma V(s_{t+1}) - V(s_t),再做 GAE:

A^t=l=0Tt1(γλ)lδt+l\hat{A}_t = \sum_{l=0}^{T-t-1} (\gamma\lambda)^l \delta_{t+l}

然后用新参数 θ\theta 对完整序列做一次并行前向(不再是逐 token 自回归),重新计算所有 token 的 logπθ\log \pi_\theta

Loss 由三部分组成:

Lactor(θ)=1Tt=1Tmin(rt(θ)A^t, clip(rt(θ),1ϵ,1+ϵ)A^t)\mathcal{L}_{actor}(\theta) = -\frac{1}{T}\sum_{t=1}^{T} \min\left(r_t(\theta)\hat{A}_t,\ \mathrm{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon)\hat{A}_t\right) Lcritic(ϕ)=1Tt=1T12(Vϕ(st)Rttarget)2,Ltotal=Lactor+c1Lcriticc2Lentropy\mathcal{L}_{critic}(\phi) = \frac{1}{T}\sum_{t=1}^{T}\frac{1}{2}\left(V_\phi(s_t) - R_t^{target}\right)^2, \quad \mathcal{L}_{total} = \mathcal{L}_{actor} + c_1\mathcal{L}_{critic} - c_2\mathcal{L}_{entropy}

其中概率比率 rt(θ)=exp(logπθlogπold)r_t(\theta) = \exp(\log \pi_\theta - \log \pi_{old})。RL 阶段学习率远小于预训练,通常 5×1072×1065\times10^{-7} \sim 2\times10^{-6},配 linear/cosine 衰减。

三、GRPO:去掉 Critic

GRPO(DeepSeek-Math / DeepSeek-R1 采用)的核心改动只有一个:扔掉与 Actor 同参数量的 Critic,改为对同一 prompt 采样 GG 个回答,用组内相对得分当优势:

A^i=rimean({r1,,rG})std({r1,,rG})+ϵ\hat{A}_i = \frac{r_i - \mathrm{mean}(\{r_1, \dots, r_G\})}{\mathrm{std}(\{r_1, \dots, r_G\}) + \epsilon}

这是个序列级标量,直接广播到回答 oio_i 的每个 token 位置。GRPO 目标函数:

LGRPO(θ)=1Gi=1G1oit=1oi[min(πθ(oi,tq,oi,<t)πold(oi,tq,oi,<t)A^i, clip()A^i)βDKL(πθπref)]\mathcal{L}_{GRPO}(\theta) = -\frac{1}{G}\sum_{i=1}^{G}\frac{1}{|o_i|}\sum_{t=1}^{|o_i|}\left[\min\left(\frac{\pi_\theta(o_{i,t} \mid q, o_{i,<t})}{\pi_{old}(o_{i,t} \mid q, o_{i,<t})}\hat{A}_i,\ \mathrm{clip}(\cdot)\hat{A}_i\right) - \beta D_{KL}(\pi_\theta \| \pi_{ref})\right]

KL 项用 Schulman 近似(k3 估计器):

DKL(πθπref)exp(logπreflogπθ)(logπreflogπθ)1D_{KL}(\pi_\theta \| \pi_{ref}) \approx \exp(\log \pi_{ref} - \log \pi_\theta) - (\log \pi_{ref} - \log \pi_\theta) - 1

3.1 对 logits 的梯度

设采样 token 为 k=otk = o_t,未触发截断时目标 t=rt(θ)A^i\ell_t = -r_t(\theta)\hat{A}_i。由 softmax 导数 π(k)zj=π(k)(I(j=k)π(j))\frac{\partial \pi(k)}{\partial z_j} = \pi(k)(\mathbb{I}(j{=}k) - \pi(j)),链式法则展开:

tzt,j=A^irt(θ)(I(j=k)πθ(jst))\frac{\partial \ell_t}{\partial z_{t,j}} = -\hat{A}_i \cdot r_t(\theta) \cdot \left(\mathbb{I}(j{=}k) - \pi_\theta(j \mid s_t)\right)

更新初期 πθπold\pi_\theta \approx \pi_{old}rt1r_t \approx 1 时简化为 A^i(eytπθ)-\hat{A}_i(e_{y_t} - \pi_\theta)。直读:

3.2 PyTorch 实现

核心算子对照表:

算法步骤PyTorch 算子
组优势归一化torch.mean(), torch.std(), 广播除法
log-prob 提取F.log_softmax() + torch.gather()
概率比值torch.exp(logps - old_logps)
PPO 截断torch.clamp() + torch.min()
KL 估计torch.exp() + 算术运算
mask 归一化(loss * mask).sum() / mask.sum()
反向传播loss.backward() + clip_grad_norm_() + step()

完整核心代码(与 TRL grpo_trainer.pycompute_loss 逐行对齐):

import torch
import torch.nn as nn
import torch.nn.functional as F
def compute_group_advantage(rewards: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
"""rewards: (batch_size, group_size) —— 同一 prompt 的 G 个回答"""
mean = rewards.mean(dim=-1, keepdim=True)
std = rewards.std(dim=-1, keepdim=True)
return (rewards - mean) / (std + eps) # (B, G)
def get_per_token_logps(model, input_ids, attention_mask):
"""input_ids: (B, S) —— prompt + completion 拼接"""
logits = model(input_ids, attention_mask=attention_mask).logits # (B, S, V)
shift_logits = logits[:, :-1, :].contiguous() # 自回归对齐: 位置 t 预测 t+1
shift_labels = input_ids[:, 1:].contiguous()
log_probs = F.log_softmax(shift_logits, dim=-1)
per_token_logps = torch.gather(log_probs, dim=-1, index=shift_labels.unsqueeze(-1)).squeeze(-1)
return per_token_logps # (B, S-1)
def grpo_loss(model, ref_model, input_ids, attention_mask, completion_mask,
old_log_probs, advantages, clip_eps=0.2, beta=0.04):
# 1. 当前策略 logps(开梯度)
current_log_probs = get_per_token_logps(model, input_ids, attention_mask)
# 2. 参考模型 logps(锁梯度)
with torch.no_grad():
ref_log_probs = get_per_token_logps(ref_model, input_ids, attention_mask)
# 3. 重要性采样比值, exp(log_a - log_b) 数值稳定
ratio = torch.exp(current_log_probs - old_log_probs)
# 4. PPO-Clip
surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1.0 - clip_eps, 1.0 + clip_eps) * advantages
policy_loss = -torch.min(surr1, surr2)
# 5. Schulman KL 近似
kl_div = torch.exp(ref_log_probs - current_log_probs) - (ref_log_probs - current_log_probs) - 1
# 6. 合并 + 7. mask 平均
per_token_loss = policy_loss + beta * kl_div
loss = (per_token_loss * completion_mask).sum() / completion_mask.sum().clamp(min=1.0)
return loss

训练主循环:

optimizer.zero_grad()
loss = grpo_loss(model=actor, ref_model=ref, input_ids=batch_ids,
attention_mask=batch_mask, completion_mask=batch_completion_mask,
old_log_probs=batch_old_logps, advantages=flat_advantages,
clip_eps=0.2, beta=0.01)
loss.backward() # 反向传播, 梯度自动累加进 param.grad
torch.nn.utils.clip_grad_norm_(actor.parameters(), max_norm=1.0) # 梯度裁剪
optimizer.step() # AdamW 更新参数
与 TRL 实现的细节差异

上文是教学版,与 HuggingFace TRL 当前 grpo_trainer.py 对比,有几处值得知道的差异:TRL 在优势归一化时分母用 std_rewards + 1e-4(不是 1e-8);KL 估计器与我们的一致(exp(ref - cur) - (ref - cur) - 1,即 k3);重要性采样支持 token 级与 sequence 级两种;当 num_iterations == 1steps_per_generation <= gradient_accumulation_steps 时,TRL 直接用 per_token_logps.detach() 当 old_logps 跳过重复计算。

3.3 old_log_probs 与 current_log_probs 为什么不同

两者只在第一个 epoch 的第一个 mini-batch 前向时严格相等(此时 θ=θold\theta = \theta_{old},ratio = 1)。此后每一步更新都让参数偏离采样时刻的参数,两者必然不同。

为什么必须显式保留 old_log_probs?因为目标函数来自重要性采样:样本由 πθold\pi_{\theta_{old}} 采样,但你在优化 πθ\pi_\theta,必须除以采样概率密度修正偏差。若 old_log_probs 每步重算,ratio 恒为 1,clip 机制与 off-policy 校正全部失效——这就是 Rollout 和 Training 必须是两个阶段的原因。

四、grad_fn:自动微分在做什么

调用 loss.backward() 时,梯度如何知道往哪流?

PyTorch 前向时,只要参与运算的张量有 requires_grad=True,引擎就会为每个运算生成一个可调用对象挂在输出的 .grad_fn 上——「前向运算的反函数记录器」:记录用了哪种运算、上游是谁、反向该用哪条导数公式。

x = torch.tensor([2.0], requires_grad=True)
y = x ** 3 # y.grad_fn → <PowBackward0>
z = y + 5 # z.grad_fn → <AddBackward0>

z.backward() 时引擎沿 DAG 逆向遍历:AddBackward 传梯度给 y,PowBackward 应用 3x23x^2,结果写入 x.grad

每个节点内部存三类信息:

| 组成 | 内容 | 开销 | |---|---|---|---| | 拓扑指针 | 对象 + next_edges 引用 | 几乎可忽略(每层几十~几百字节) | | Saved Tensors | 求导所需的前向输入/输出 | 主要开销O(B×S×L×H)O(B \times S \times L \times H) | | apply 方法 | 反向的 C++ 算子 | 反向 FLOPs ≈ 2× 前向 |

哪些算子必须缓存:乘法要缓存 yyz/x=y\partial z/\partial x = y);softmax 要缓存前向输出概率;自注意力 QKTQK^T 求导要乘回 Q、K。加法什么都不用存(导数恒为 1)。

两个工程对策由此而来:

五、Optimizer:AdamW 四机制

优化器决定「根据梯度,下一步跨多大、往哪偏」。朴索 SGD θθηgt\theta \leftarrow \theta - \eta g_t 训大模型必死:病态曲率下陡壁震荡、鞍点卡死、不同参数梯度量级差数个数量级。AdamW 用四个机制解决:

① 一阶动量(惯性)mt=β1mt1+(1β1)gtm_t = \beta_1 m_{t-1} + (1-\beta_1)g_tβ1=0.9\beta_1 = 0.9)。历史梯度的指数加权平均,像有质量的重球——正负震荡被平均抵消,轨迹平滑。

② 二阶动量(逐参数油门)vt=β2vt1+(1β2)gt2v_t = \beta_2 v_{t-1} + (1-\beta_2)g_t^2β2=0.999\beta_2 = 0.999)。梯度剧烈的参数分母大→刹车;梯度微弱的参数(低频词 embedding)分母小→放大更新。每个参数拥有独立步长。

③ 偏差修正m0=v0=0m_0 = v_0 = 0 导致初期动量严重偏小,用 m^t=mt/(1β1t)\hat{m}_t = m_t/(1-\beta_1^t)v^t=vt/(1β2t)\hat{v}_t = v_t/(1-\beta_2^t) 修正冷启动。

④ 解耦权重衰减:Adam 把 L2 惩罚混进动量会被 vt\sqrt{v_t} 扭曲(大梯度参数反而少受罚)。AdamW 把正则从动量中剥离,更新后直接按比例缩减权重:

θt=θt1η(m^tv^t+ϵ+λθt1)\theta_t = \theta_{t-1} - \eta\left(\frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} + \lambda\theta_{t-1}\right)

这四行构成训练闭环,值得逐行理解:

optimizer.zero_grad() # 梯度是累加(param.grad += dL/dparam)不是覆盖, 不清会混入上一轮
loss.backward() # 沿 DAG 逆遍历, 偏导写入 param.grad, 权重本身未变
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # 全局 L2 范数超限则等比缩放
optimizer.step() # 按 param.grad + 动量状态真正修改权重

clip_grad_norm_ 在 RL 里尤其关键:RL 梯度方差天然高于监督学习,离群奖励会瞬时产生巨大梯度,一步失控的更新足以毁掉预训练学来的表征。裁剪保留方向、约束最大步长。

zero_grad 推荐用 zero_grad(set_to_none=True):把 .grad 直接设 None 而非写零,省显存带宽。梯度累积(gradient accumulation)正是利用「不清零」实现的——小卡攒多步梯度再 step。

六、训练显存全景

「7B 模型只有 14GB,为什么单张 80GB 卡跑不下训练?」——因为显存的大头不是权重,优化器状态和激活值通常数倍于它。

显存=权重2Φ+梯度2Φ+优化器状态12Φ+激活值O(BSLH)+碎片+通信/kernel 工作区\text{显存} = \underbrace{\text{权重}}_{2\Phi} + \underbrace{\text{梯度}}_{2\Phi} + \underbrace{\text{优化器状态}}_{12\Phi} + \underbrace{\text{激活值}}_{O(B{\cdot}S{\cdot}L{\cdot}H)} + \text{碎片} + \text{通信/kernel 工作区}

Φ\Phi 为参数量,BF16 权重+梯度各 2 字节/参数,FP32 master weights + m + v 共 12 字节/参数。)

7B 模型的具体占用:

组成部分格式7B 占用优化手段
权重BF1614 GBZeRO-3 / FSDP 分片、LoRA
梯度BF1614 GBZeRO-2 分片
优化器状态FP32 ×384 GBZeRO-1、8-bit Adam、Paged AdamW
激活值BF16, 动态10~80+ GBActivation Checkpointing、FlashAttention
碎片/通信/kernel动态5~10 GBempty_cache()、通信桶调优

关键规律是缩放方向不同

短序列 SFT 时优化器状态是绝对主力;RL 长链推理(S = 16k+)时激活值轻易数倍超越权重与优化器状态。

6.1 Activation Checkpointing

「有的算子求导必须依赖前向输入/输出,那怎么敢丢?」——重计算不违反数学,它是用时间换空间:

前向:以 Transformer 层为单位,层内照常算,但结束时丢弃全部中间激活,只保留层输入张量。 反向:梯度回传到该层时,用保留的层输入在 torch.no_grad()现场重跑一遍层内前向,把重算出的中间值挂回 grad_fn 立即求导,用完即弃。

标准: 前向存 Layer1[A→B→C 全部中间值] → Layer2[...]
CKPT: 前向仅存 Layer1 输入, A→B→C 中间值全丢
反向: 到达 Layer1 → 用输入重跑 A→B→C → 立即求导 → 丢弃

代价与收益:激活显存降 6080%,整体耗时增约 2030%。

6.2 碎片与缓冲

这就是 ZeRO、8-bit Adam、Paged AdamW 等技术存在的理由:分别切分/量化优化器状态与梯度,把「静态部分」拆到多卡与 CPU。

结语

从策略梯度到 GRPO,从 grad_fn 到 AdamW,再到显存的六个组成部分——LLM RL 系统的每一层都在做同一件事:把一个数学上干净的更新规则,塞进有限显存与有限带宽的物理约束里。理解了这条链条,再读 slime、OpenRLHF、veRL 这类框架的源码时,看到的就不是配置项的堆砌,而是每一项对应的具体问题。


Share this post:

Next Post
GLM-5.3 系统优化深度解读:slime 框架如何支撑长周期 RL 扩展