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 依赖策略梯度定理,用对数似然加权来估计梯度:
其中 是优势函数(Advantage),衡量在状态 下选择 相比平均水平「好多少」。这个公式有一个直读:梯度顺着 流,奖励信号只作为系数——采样本身不出现在梯度路径里。
把生成过程建模为 MDP:
| RL 概念 | LLM 对应物 |
|---|---|
| 状态 | 上下文 |
| 动作 | 下一个 token |
| 策略 | LLM 本身,输出词表上的分布 |
| 奖励 | 偏好对齐用 Reward Model;数学/代码推理用规则判题 |
二、PPO 的四模型架构
以标准 PPO 为例,系统涉及 4 个模型实体:
- Actor():待训练模型
- Critic():预测状态价值,初始化自 Reward Model
- Reference():冻结的初始 SFT 模型,提供 KL 约束,防 reward hacking
- Reward Model:打分系统(或确定性的判题器)
训练分两个阶段:
阶段 A —— Rollout(推理模式):Actor 自回归采样完整序列,记录每个 token 的 (即 old_log_probs),计算标量奖励,并把序列喂给 Reference 模型取 。
阶段 B —— Training(开梯度):把 KL 惩罚折进 token 级奖励——除最后一个 token 外都是 ,末位 token 额外加上 。Critic 前向得到 ,算 TD 误差 ,再做 GAE:
然后用新参数 对完整序列做一次并行前向(不再是逐 token 自回归),重新计算所有 token 的 。
Loss 由三部分组成:
其中概率比率 。RL 阶段学习率远小于预训练,通常 ,配 linear/cosine 衰减。
三、GRPO:去掉 Critic
GRPO(DeepSeek-Math / DeepSeek-R1 采用)的核心改动只有一个:扔掉与 Actor 同参数量的 Critic,改为对同一 prompt 采样 个回答,用组内相对得分当优势:
这是个序列级标量,直接广播到回答 的每个 token 位置。GRPO 目标函数:
KL 项用 Schulman 近似(k3 估计器):
3.1 对 logits 的梯度
设采样 token 为 ,未触发截断时目标 。由 softmax 导数 ,链式法则展开:
更新初期 、 时简化为 。直读:
- (好于组内均值)→ 该 token logit 的梯度为负 → 梯度下降后概率上升
- → 概率被压低,质量推给其他候选词
- 触发截断后目标变成常数 ,梯度恒为 0
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.py 的 compute_loss 逐行对齐):
import torchimport torch.nn as nnimport 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.gradtorch.nn.utils.clip_grad_norm_(actor.parameters(), max_norm=1.0) # 梯度裁剪optimizer.step() # AdamW 更新参数上文是教学版,与 HuggingFace TRL 当前 grpo_trainer.py 对比,有几处值得知道的差异:TRL 在优势归一化时分母用 std_rewards + 1e-4(不是 1e-8);KL 估计器与我们的一致(exp(ref - cur) - (ref - cur) - 1,即 k3);重要性采样支持 token 级与 sequence 级两种;当 num_iterations == 1 且 steps_per_generation <= gradient_accumulation_steps 时,TRL 直接用 per_token_logps.detach() 当 old_logps 跳过重复计算。
3.3 old_log_probs 与 current_log_probs 为什么不同
两者只在第一个 epoch 的第一个 mini-batch 前向时严格相等(此时 ,ratio = 1)。此后每一步更新都让参数偏离采样时刻的参数,两者必然不同。
为什么必须显式保留 old_log_probs?因为目标函数来自重要性采样:样本由 采样,但你在优化 ,必须除以采样概率密度修正偏差。若 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 应用 ,结果写入 x.grad。
每个节点内部存三类信息:
| 组成 | 内容 | 开销 |
|---|---|---|---|
| 拓扑指针 | 对象 + next_edges 引用 | 几乎可忽略(每层几十~几百字节) |
| Saved Tensors | 求导所需的前向输入/输出 | 主要开销, |
| apply 方法 | 反向的 C++ 算子 | 反向 FLOPs ≈ 2× 前向 |
哪些算子必须缓存:乘法要缓存 ();softmax 要缓存前向输出概率;自注意力 求导要乘回 Q、K。加法什么都不用存(导数恒为 1)。
两个工程对策由此而来:
- 推理/Rollout 阶段必须
with torch.no_grad():—— 直接不建 grad_fn,不存激活 - 长文本训练用 Activation Checkpointing(见第六节)
五、Optimizer:AdamW 四机制
优化器决定「根据梯度,下一步跨多大、往哪偏」。朴索 SGD 训大模型必死:病态曲率下陡壁震荡、鞍点卡死、不同参数梯度量级差数个数量级。AdamW 用四个机制解决:
① 一阶动量(惯性):()。历史梯度的指数加权平均,像有质量的重球——正负震荡被平均抵消,轨迹平滑。
② 二阶动量(逐参数油门):()。梯度剧烈的参数分母大→刹车;梯度微弱的参数(低频词 embedding)分母小→放大更新。每个参数拥有独立步长。
③ 偏差修正: 导致初期动量严重偏小,用 、 修正冷启动。
④ 解耦权重衰减:Adam 把 L2 惩罚混进动量会被 扭曲(大梯度参数反而少受罚)。AdamW 把正则从动量中剥离,更新后直接按比例缩减权重:
这四行构成训练闭环,值得逐行理解:
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 卡跑不下训练?」——因为显存的大头不是权重,优化器状态和激活值通常数倍于它。
( 为参数量,BF16 权重+梯度各 2 字节/参数,FP32 master weights + m + v 共 12 字节/参数。)
7B 模型的具体占用:
| 组成部分 | 格式 | 7B 占用 | 优化手段 |
|---|---|---|---|
| 权重 | BF16 | 14 GB | ZeRO-3 / FSDP 分片、LoRA |
| 梯度 | BF16 | 14 GB | ZeRO-2 分片 |
| 优化器状态 | FP32 ×3 | 84 GB | ZeRO-1、8-bit Adam、Paged AdamW |
| 激活值 | BF16, 动态 | 10~80+ GB | Activation Checkpointing、FlashAttention |
| 碎片/通信/kernel | 动态 | 5~10 GB | empty_cache()、通信桶调优 |
关键规律是缩放方向不同:
- 优化器状态只随参数量 走,与 batch、序列长度无关——静态占用
- 激活值随 线性走,注意力还要存 的 attention map——动态占用
短序列 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 碎片与缓冲
- Caching Allocator 碎片:
nvidia-smi显示占满 80GB 报 OOM,torch.cuda.memory_allocated()却只有 50GB——变长序列频繁申请/释放不同形状张量,物理显存被切碎,申请连续大张量失败。常浪费 10~25%。 - 通信缓冲:DDP/FSDP/ZeRO 的 AllReduce 按桶(bucket)打包梯度,几百 MB~数 GB 预分配;TP 每层 attention/MLP 内部都有聚合传输。
- Kernel workspace:cuBLAS/FlashAttention 执行时需要临时 workspace,几百 MB~2GB。
这就是 ZeRO、8-bit Adam、Paged AdamW 等技术存在的理由:分别切分/量化优化器状态与梯度,把「静态部分」拆到多卡与 CPU。
结语
从策略梯度到 GRPO,从 grad_fn 到 AdamW,再到显存的六个组成部分——LLM RL 系统的每一层都在做同一件事:把一个数学上干净的更新规则,塞进有限显存与有限带宽的物理约束里。理解了这条链条,再读 slime、OpenRLHF、veRL 这类框架的源码时,看到的就不是配置项的堆砌,而是每一项对应的具体问题。