PPO 深度解析:信赖域优化的艺术与 LLM 对齐实践

5457 字
27 分钟
PPO 深度解析:信赖域优化的艺术与 LLM 对齐实践

1. 引言:PPO 为什么会成为 LLM 对齐的主流?#

1.1 强化学习的三代算法#

在 PPO 之前,Policy Gradient 算法经历了三代演进,每一代都在解决上一代的核心问题:

第一代: REINFORCE (Williams, 1992)
∇J = E[∇_θ log π_θ(a|s) · R]
问题:更新步长难以控制,方差极大,训练极不稳定
↓ 解决方差问题
第二代: Actor-Critic (Sutton & Barto, 1999)
用 Value Network 估计 baseline,减少方差
问题:仍无步长保证,可能一步跨太大 → 崩溃
↓ 加入信赖域约束
第三代: TRPO (Schulman et al., 2015)
用 KL 约束限制新旧策略差异
问题:需要共轭梯度求解二阶优化,工程复杂
↓ 简化 KL 约束(不用二阶优化)
第四代: PPO (Schulman et al., 2017)
用 Clip 机制一阶优化,兼顾稳定性与简单性
成为 LLM 对齐的默认选择

1.2 为什么 LLM 对齐选择 PPO?#

PPO 相比 DPO 有不可替代的优势:

维度DPOPPO
在线探索❌ 纯离线✅ 在线采样
多轮迭代改进❌ 一次性✅ 持续改进
复杂奖励⚠️ 单标量✅ 任意奖励函数
工程复杂度
显存需求~2× 模型~4× 模型
代表应用Llama 3, MistralGPT-4, Claude 2, ChatGPT

核心洞察:DPO 是离线的直接偏好优化,PPO 是在线的策略优化。当你的目标是:

  • 追赶 SOTA(最大性能)→ PPO
  • 快速迭代(速度和简单性)→ DPO
Tip

2024 年 DeepSeek-R1 的出现打破了这一格局——GRPO 以 PPO 的在线探索优势 + DPO 的简洁性,在数学推理任务上取得了 SOTA。这将在后文详细讨论。


2. 策略梯度基础:从 REINFORCE 到 Actor-Critic#

2.1 策略梯度定理#

对于一个策略 πθ(as)\pi_\theta(a|s),目标是最小化负期望回报(对于 LLM 对齐是最大化期望奖励):

J(θ)=Eτπθ[t=0T1r(st,at)]J(\theta) = \mathbb{E}_{\tau \sim \pi_\theta}\Big[\sum_{t=0}^{T-1} r(s_t, a_t)\Big]

其中 τ=(s0,a0,s1,a1,)\tau = (s_0, a_0, s_1, a_1, \ldots) 是轨迹(对于 LLM,ss 是 token 序列的前缀,aa 是下一个 token)。

策略梯度定理给出:

θJ(θ)=Eτπθ[t=0T1θlogπθ(atst)  Gt]\nabla_\theta J(\theta) = \mathbb{E}_{\tau\sim\pi_\theta}\Big[\sum_{t=0}^{T-1}\nabla_\theta\log\pi_\theta(a_t|s_t)\;G_t\Big]

其中 Gt=t=tT1rtG_t = \sum_{t'=t}^{T-1} r_{t'} 是从时刻 tt 开始的累计回报。

2.2 REINFORCE:基本策略梯度#

def reinforce_loss(log_probs, rewards, dones, gamma=1.0):
"""
REINFORCE 损失函数。
log_probs: (B, T) 每个 token 的 log π(a_t | s_t)
rewards: (B, T) 每个 token 的奖励
"""
T = log_probs.size(1)
# 计算 G_t(从后往前累计)
G = torch.zeros_like(rewards)
running_return = torch.zeros(rewards.size(0), device=rewards.device)
for t in reversed(range(T)):
running_return = rewards[:, t] + (1 - dones[:, t]) * gamma * running_return
G[:, t] = running_return
# 策略梯度: E[∇log π · G]
# 对每条轨迹的 log prob 求和,乘以 G
loss = -(log_probs * G).sum(dim=1).mean()
return loss

2.3 方差问题的根源#

REINFORCE 的方差极大——因为 GtG_t 是从随机策略采样的路径上计算的,不同采样路径的 GtG_t 差异巨大。

两种降低方差的方法

  1. 增加采样数量:理论上有效,但计算成本爆炸
  2. 引入 baseline:从 GtG_t 中减去一个基准 b(st)b(s_t),不改变期望但降低方差

2.4 Actor-Critic:引入 Value Network#

引入一个 Value Network Vϕ(st)V_\phi(s_t) 作为 baseline:

At=GtVϕ(st)A_t = G_t - V_\phi(s_t)

AtA_t 称为优势函数(Advantage),衡量在状态 sts_t 下采取动作 ata_t 比”平均水平”好多少。

Actor-Critic 架构:

class ActorCritic(torch.nn.Module):
"""
简化的 Actor-Critic。
对于 LLM:状态 s 是 token 序列,动作 a 是下一个 token。
"""
def __init__(self, llm_backbone):
super().__init__()
self.llm = llm_backbone
hidden_size = llm_backbone.config.hidden_size
self.value_head = nn.Linear(hidden_size, 1) # Critic
self.action_head = nn.Linear(hidden_size, llm_backbone.config.vocab_size) # Actor
def forward(self, input_ids, attention_mask):
outputs = self.llm(input_ids=input_ids, attention_mask=attention_mask)
hidden = outputs.hidden_states[-1]
values = self.value_head(hidden).squeeze(-1) # (B, L)
logits = self.action_head(hidden) # (B, L, V)
return logits, values
Actor-Critic 的根本问题

Actor-Critic 没有新旧策略比率的约束,梯度可能步子太大,导致策略崩溃。PPO 的 clip 机制正是为了解决这个问题。


3. TRPO:信赖域优化的理论基础#

3.1 信赖域的核心思想#

直觉:在参数空间中,新旧策略不能离得太远,否则策略会崩溃。

TRPO 的约束形式:

maxθ Et[πθ(atst)πθold(atst)  At]\max_\theta\ \mathbb{E}_t\Big[\frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{\text{old}}}(a_t|s_t)}\;A_t\Big]s.t.Et[KL(πθold(st)πθ(st))]δ\text{s.t.}\quad \mathbb{E}_t\Big[\text{KL}\big(\pi_{\theta_{\text{old}}}(\cdot|s_t)\,\|\,\pi_\theta(\cdot|s_t)\big)\Big] \leq \delta

δ\delta 是信任域半径(trust region radius)。

3.2 从约束到惩罚:共轭梯度法#

TRPO 需要求解带约束的优化问题,实际实现中:

  1. 共轭梯度法(Conjugate Gradient)近似求解 KL 约束的二阶近似
  2. 线性搜索确定步长
def trpo_update(policy_model, old_policy, advantages, kl_target=0.01):
"""
TRPO 更新步骤(简化的核心逻辑)。
1. 计算策略梯度
2. 共轭梯度法求搜索方向
3. 线搜索确定步长
"""
# === 1. 计算 KL 约束的 Hessian-vector product ===
# H = ∇² KL / ∇θ²,复杂度 O(θ²)
# 共轭梯度法只需要 H·v,不需要显式构造 H
# === 2. 计算自然梯度方向 ===
# ∇_natural = H⁻¹ · ∇_policy
search_direction = conjugate_gradient(
hvp_func=lambda v: kl_hessian_vector_product(old_policy, policy_model, v),
b=policy_grad,
max_iter=20,
)
# === 3. 线搜索(在信赖域内找最优)===
for alpha in [1.0, 0.5, 0.25, 0.125]:
candidate_theta = theta + alpha * search_direction
if kl_divergence(theta, candidate_theta) < kl_target:
theta = candidate_theta
break

3.3 TRPO 的工程痛点#

痛点描述
共轭梯度开销大需要计算 Hessian-Vector Product,O(θ2)O(\theta^2)
内存占用高需要存储 Fisher 信息矩阵或其近似
超参敏感δ\delta 的选择对结果影响很大
难以与大规模 LLM 结合上述问题在亿级参数上被放大

PPO 用 clip 机制一阶近似了这个二阶约束,完美解决了这些问题。


4. PPO:Clip 机制的数学推导#

4.1 重要性比率#

定义重要性比率(Importance Sampling Ratio):

rt(θ)=πθ(atst)πθold(atst)r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{\text{old}}}(a_t|s_t)}

对于 LLM,这意味着每个 token 的概率比:

def compute_importance_ratio(
log_probs_new: torch.Tensor, # π_θ(a_t | s_t)
log_probs_old: torch.Tensor, # π_{θ_old}(a_t | s_t)
) -> torch.Tensor:
"""r_t(θ) = π_θ / π_{θ_old}"""
return torch.exp(log_probs_new - log_probs_old) # (B, T)

4.2 原始策略梯度损失#

如果用原始的策略梯度:

LPG(θ)=Et[rt(θ)  At]\mathcal{L}^{\text{PG}}(\theta) = \mathbb{E}_t\Big[r_t(\theta)\;A_t\Big]

问题:当 rt(θ)r_t(\theta) 变得很大(超过 1)时,更新步长无法控制——策略可能在一次更新中完全偏离。

4.3 Clipped Surrogate 目标#

PPO 的核心创新是用 clip 操作限制 rtr_t 的变化范围:

LCLIP(θ)=Et[min(rt(θ)  At, clip(rt(θ),1ϵ,1+ϵ)  At)]\mathcal{L}^{\text{CLIP}}(\theta) = \mathbb{E}_t\Big[\min\big(r_t(\theta)\;A_t,\ \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon)\;A_t\big)\Big]

其中 ϵ\epsilon 是超参数(通常 ϵ=0.2\epsilon = 0.2)。

4.4 Clip 机制的直觉#

假设 A > 0(好的动作,正优势):
r_t(θ)
2.0│ ┌────────── ← clip 上界
│ /
1.8│ /
1.6│ / ← 超出上界部分被 clip
1.4│ /
1.2│ /
1.0│─────────●───────────────────
│ ↑
└─────────┼─────────────── θ_old
1 - ε 1.0 1 + ε
情况分析:
θ 在 [θ_old - ε, θ_old + ε] 内 → r_t 正常,无 clip
θ 超出信赖域(|r_t - 1| > ε)→ r_t 被 clip 到 [1-ε, 1+ε]
→ 梯度被限制,不会一步更新太大

4.5 完整的数学推导#

分情况讨论

情况 1:A > 0(正优势,鼓励这个动作)

  • 如果 rt<1+ϵr_t < 1 + \epsilon:无 clip,L=rtA\mathcal{L} = r_t \cdot A(正常增加概率)
  • 如果 rt1+ϵr_t \geq 1 + \epsilon:被 clip 到 (1+ϵ)A(1+\epsilon)\cdot A,但 min 操作使得额外增大不再被奖励

情况 2:A < 0(负优势,抑制这个动作)

  • 如果 rt>1ϵr_t > 1 - \epsilon:无 clip,L=rtA\mathcal{L} = r_t \cdot A(负值,减小概率)
  • 如果 rt1ϵr_t \leq 1 - \epsilon:被 clip 到 (1ϵ)A(1-\epsilon)\cdot A,但 min 操作使得过度减小不再被惩罚
def ppo_clipped_loss(
log_probs_new: torch.Tensor, # (B, T)
log_probs_old: torch.Tensor, # (B, T)
advantages: torch.Tensor, # (B, T)
clip_eps: float = 0.2,
) -> torch.Tensor:
"""
PPO Clipped Surrogate Loss。
L = min(r(θ)·A, clip(r(θ), 1-ε, 1+ε)·A)
物理含义:
- min 操作确保 clip 不增加 loss(不会过度优化)
- 当 A > 0 时,防止 r 变得太大(过度增加好动作的概率)
- 当 A < 0 时,防止 r 变得太小(过度惩罚坏动作)
"""
ratio = torch.exp(log_probs_new - log_probs_old) # (B, T)
# 原始目标
policy_loss_1 = ratio * advantages
# Clipped 目标
ratio_clipped = torch.clamp(ratio, 1 - clip_eps, 1 + clip_eps)
policy_loss_2 = ratio_clipped * advantages
# 取较小的那个
policy_loss = -torch.min(policy_loss_1, policy_loss_2)
# 通常对序列求平均(也可以 token-wise)
return policy_loss.mean()

4.6 为什么 min 而不是 max?#

这是一个精妙的设计:保证在 clip 区域外不会过度优化

考虑 A>0A > 0rt(θ)>1+ϵr_t(\theta) > 1+\epsilon(策略已经大幅增加了这个动作的概率):

  • L1=rtA\mathcal{L}_1 = r_t \cdot A 很小(因为 rtr_t 大导致梯度小)
  • L2=(1+ϵ)A\mathcal{L}_2 = (1+\epsilon)\cdot A 更小(是上界)

min(L1,L2)=L2\min(\mathcal{L}_1, \mathcal{L}_2) = \mathcal{L}_2不再增加这个动作的概率——因为已经过度了。


5. GAE:广义优势估计#

5.1 TD(λ) 的优势函数版本#

对于 LLM 的序列决策(只有一个终态 reward),优势函数的估计是一个关键问题。

TD(0) 估计

A^tTD(0)=rt+γV(st+1)V(st)\hat{A}_t^{\text{TD}(0)} = r_t + \gamma\,V(s_{t+1}) - V(s_t)

n步回报

A^t(n)=l=0n1γlrt+l+γnV(st+n)V(st)\hat{A}_t^{(n)} = \sum_{l=0}^{n-1}\gamma^l r_{t+l} + \gamma^n V(s_{t+n}) - V(s_t)

**GAE(Generalized Advantage Estimation)**将这些 n步估计做加权平均:

A^tGAE(λ,γ)=(1λ)n=0(λγ)n  A^t(n)\hat{A}_t^{\text{GAE}(\lambda,\gamma)} = (1-\lambda)\sum_{n=0}^{\infty}(\lambda\gamma)^n\;\hat{A}_t^{(n)}

5.2 GAE 的递推形式#

GAE 可以用递归方式高效计算:

A^t=δt+γλ  A^t+1\hat{A}_t = \delta_t + \gamma\,\lambda\;\hat{A}_{t+1}

其中 δt=rt+γV(st+1)V(st)\delta_t = r_t + \gamma\,V(s_{t+1}) - V(s_t) 是 TD 误差。

def compute_gae(
rewards: torch.Tensor, # (B, T) 序列奖励
values: torch.Tensor, # (B, T) Value 网络估计
dones: torch.Tensor, # (B, T) 是否结束
gamma: float = 1.0,
lam: float = 0.95,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
GAE 计算。
递推公式: A_t = δ_t + γλA_{t+1}
其中 δ_t = r_t + γV(s_{t+1}) - V(s_t)
参数:
gamma (γ): 折扣因子。对于 LLM 序列(终态任务),通常 γ=1.0
lam (λ): GAE 平滑参数。
λ=0 → TD(0),方差低但高偏差
λ=1 →蒙特卡洛,无偏差但高方差
λ≈0.95 → 平衡两者
"""
B, T = rewards.shape
advantages = torch.zeros_like(rewards)
last_adv = torch.zeros(B, device=rewards.device)
for t in reversed(range(T)):
if t == T - 1:
next_value = 0.0 # 序列结束,无 future value
else:
next_value = values[:, t + 1]
nonterminal = 1.0 - dones[:, t]
delta = rewards[:, t] + gamma * nonterminal * next_value - values[:, t]
# A_t = δ_t + γλ A_{t+1}
last_adv = delta + gamma * lam * nonterminal * last_adv
advantages[:, t] = last_adv
returns = advantages + values # 用于训练 Value 网络
return advantages, returns

5.3 GAE 超参数的影响#

λ\lambda偏差方差适用场景
0(TD(0))极低在线学习,快速估计
0.9 ~ 0.95大多数场景,推荐值
1(MC)0(无偏)短序列,高采样数
LLM 对齐中的特殊处理

在 LLM 对齐中,reward 通常只在序列最后一个 token 给出(RM 分数),中间 token 的 reward 为 0。这种情况下 GAE 的 λ\lambda 参数影响更加显著——λ\lambda 越接近 1,奖励信号越能传播到所有 token。

5.4 Value 网络损失#

Value 网络(Critic)的任务是预测期望回报,其损失为均方误差:

LV=E[(R^tVϕ(st))2]\mathcal{L}_V = \mathbb{E}\big[(\hat{R}_t - V_\phi(s_t))^2\big]

其中 R^t\hat{R}_t 是 GAE 返回值(returns)。

通常使用 clip Value 预测来防止 Value 网络更新过大(PPO-Penalty 的做法):

def value_loss(
values_new: torch.Tensor,
values_old: torch.Tensor,
returns: torch.Tensor,
clip_eps: float = 0.2,
):
"""
Clipped Value Loss(PPO-Penalty 的 Value 损失)。
限制 Value 预测的变化幅度,防止 Value 更新过大导致不稳定。
"""
clipped_values = torch.clamp(
values_new,
values_old - clip_eps,
values_old + clip_eps,
)
loss1 = (values_new - returns) ** 2
loss2 = (clipped_values - returns) ** 2
return torch.max(loss1, loss2).mean()

6. PPO 用于 LLM 对齐:完整实现#

6.1 四模型架构#

在 LLM 对齐中,PPO 需要同时运行四个模型:

┌─────────────────────────────────────────────────────────────┐
│ PPO 四模型架构 │
├─────────────────────────────────────────────────────────────┤
│ │
│ ① Policy Model (π_θ) │
│ → 待优化的 LLM,梯度更新 │
│ │
│ ② Reference Model (π_ref) │
│ → SFT 模型,参数冻结,计算 KL 惩罚 │
│ │
│ ③ Reward Model (R_φ) │
│ → 偏好数据训练的奖励模型,参数冻结 │
│ → 只在 rollout 阶段使用一次打分 │
│ │
│ ④ Value Network (V_ψ) │
│ → 与 policy 共享 LLM backbone,仅 value head 可学习 │
│ → 在 rollout 和 update 阶段都运行 │
│ │
│ 显存占用 ≈ 4 × LLM 参数 + embeddings │
│ → 通常需要对 policy/ref/value 做 LoRA 或量化 │
│ │
└─────────────────────────────────────────────────────────────┘

6.2 Rollout:在线采样#

def ppo_rollout(
policy_model,
ref_model,
reward_model,
critic_model,
prompts: list[str],
tokenizer,
max_new_tokens: int = 256,
temperature: float = 1.0,
top_p: float = 0.9,
num_return_sequences: int = 4, # GRPO 风格的 G
) -> dict:
"""
PPO Rollout 阶段:
1. 用当前策略采样多个回答
2. 用 Reward Model 评分
3. 计算 KL 惩罚
4. 收集用于 PPO 更新的数据
"""
# === 1. 用策略采样(可以多个 response per prompt)===
all_responses = []
all_log_probs = []
all_ref_log_probs = []
all_values = []
all_rewards = []
for prompt in prompts:
prompt_tensors = tokenizer([prompt], return_tensors="pt").to(policy_model.device)
# 如果 num_return_sequences > 1,就是 GRPO-style
if num_return_sequences > 1:
# 批量生成多个 response
prompt_batch = [prompt] * num_return_sequences
else:
prompt_batch = [prompt]
with torch.no_grad():
outputs = policy_model.generate(
input_ids=prompt_batch,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=temperature,
top_p=top_p,
pad_token_id=tokenizer.pad_token_id,
)
# === 2. 提取 response 部分 ===
response_tensors = extract_response(prompt_tensors, outputs, tokenizer)
# === 3. 计算 log probs ===
# Policy log probs
policy_log_probs = compute_response_log_probs(policy_model, response_tensors)
# Reference log probs(用于 KL)
with torch.no_grad():
ref_log_probs = compute_response_log_probs(ref_model, response_tensors)
# === 4. Value 估计 ===
with torch.no_grad():
_, values_pred = critic_model(
input_ids=response_tensors["input_ids"],
attention_mask=response_tensors["attention_mask"],
)
values_pred = values_pred[:, :-1].mean(dim=-1) # 标量 per sequence
# === 5. Reward Model 评分 ===
with torch.no_grad():
rewards = reward_model(
input_ids=response_tensors["input_ids"],
attention_mask=response_tensors["attention_mask"],
)
# === 6. 构建 KL 奖励塑形 ===
# KL penalty per token:-β * (log π_θ - log π_ref)
kl_penalty = -(policy_log_probs - ref_log_probs) # (B, T)
# 只有最后一个 token 获得真正的 RM 分数
T = response_tensors["input_ids"].size(1)
token_rewards = kl_penalty * 0.0 # 初始化为 0
token_rewards[:, -1] = rewards # 最后一个 token 获得 RM 分数
# 或者更常用的:token_rewards = -beta * kl_penalty, token_rewards[:, -1] += rewards
all_rewards.append(token_rewards)
all_log_probs.append(policy_log_probs)
all_ref_log_probs.append(ref_log_probs)
all_values.append(values_pred)
return {
"rewards": torch.cat(all_rewards, dim=0),
"log_probs": torch.cat(all_log_probs, dim=0),
"ref_log_probs": torch.cat(all_ref_log_probs, dim=0),
"values": torch.cat(all_values, dim=0),
}

6.3 PPO 更新#

def ppo_update_step(
policy_model,
critic_model,
rollout_data: dict,
beta: float = 0.1, # KL 系数
clip_eps: float = 0.2, # PPO clip epsilon
gamma: float = 1.0, # 折扣因子
lam: float = 0.95, # GAE lambda
entropy_coef: float = 0.01, # 熵正则
vf_coef: float = 0.5,
max_grad_norm: float = 1.0,
ppo_epochs: int = 4, # 同一个 rollout 数据更新 PPO epochs 次
minibatch_size: int = 8,
):
"""
PPO 更新步骤。
"""
rewards = rollout_data["rewards"] # (B, T)
log_probs_old = rollout_data["log_probs"] # (B, T)
ref_log_probs = rollout_data["ref_log_probs"]
values_old = rollout_data["values"] # (B,)
# === 1. GAE 计算优势 ===
dones = torch.zeros_like(rewards)
dones[:, -1] = 1.0 # 序列结束标记
advantages, returns = compute_gae(rewards, values_old.unsqueeze(1).repeat(1, rewards.size(1)), dones, gamma, lam)
# 标准化优势
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
# === 2. 每个 token 的 log prob(用于策略梯度)===
# 对于 LLM,优势是 per-sequence 的,但 apply 时展平到 token
advantages_flat = advantages.reshape(-1) # (B*T,)
# === 3. PPO Epoch 循环 ===
B, T = rewards.shape
indices = torch.randperm(B * T) # 打乱所有 token
for epoch in range(ppo_epochs):
for start in range(0, B * T, minibatch_size):
end = start + minibatch_size
idx = indices[start:end]
# === 3.1 Policy forward ===
# 需要重新 forward 以获取当前策略的 log probs
# (简化版:这里假设 rollout_data 存储了 old_log_probs)
log_probs_new = log_probs_old # placeholder,实际需要重算
# === 3.2 PPO Clipped Loss ===
policy_loss = ppo_clipped_loss(
log_probs_new.reshape(-1)[idx].unsqueeze(0),
log_probs_old.reshape(-1)[idx].unsqueeze(0),
advantages_flat[idx].unsqueeze(0),
clip_eps=clip_eps,
)
# === 3.3 Value Loss ===
# 从 rollout 数据中取出对应样本的 returns
values_new = values_old # placeholder,实际需要重算
vf_loss = ((values_new - returns.reshape(-1)[idx]) ** 2).mean()
# === 3.4 熵正则(鼓励探索)===
# entropy = compute_entropy(logits) # placeholder
# ent_loss = -entropy_coef * entropy.mean()
# === 3.5 总损失 ===
loss = policy_loss + vf_coef * vf_loss # + ent_loss
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(policy_model.parameters(), max_grad_norm)
optimizer.step()

6.4 KL 约束的奖励塑形形式#

在 LLM 对齐中,更常用的是奖励塑形(Reward Shaping)版本,而非 PPO-Penalty 的 KL 惩罚项:

rtshaped={R(x,y)βKL(πθ(xt)πref(xt))t=TβKL(πθ(xt)πref(xt))t<Tr_t^{\text{shaped}} = \begin{cases} R(x, y) - \beta\,\text{KL}\big(\pi_\theta(\cdot|x_t)\,\|\,\pi_{\text{ref}}(\cdot|x_t)\big) & t = T \\ -\beta\,\text{KL}\big(\pi_\theta(\cdot|x_t)\,\|\,\pi_{\text{ref}}(\cdot|x_t)\big) & t < T \end{cases}
def compute_kl_reward_shaping(
policy_log_probs: torch.Tensor, # (B, T)
ref_log_probs: torch.Tensor, # (B, T)
reward_scores: torch.Tensor, # (B,) RM 分数
beta: float = 0.1,
) -> torch.Tensor:
"""
KL 奖励塑形。
每个 token 获得 -β·KL 惩罚,只有最后一个 token 获得 RM 分数。
"""
kl_per_token = policy_log_probs - ref_log_probs # (B, T)
token_rewards = -beta * kl_per_token
token_rewards[:, -1] += reward_scores # 最后一个 token 加上 RM 分数
return token_rewards

7. 进阶技巧:让 PPO 训练更稳定#

7.1 Reward 归一化#

RM 的原始分数方差可能很大(不同回答的 RM 分数可能差 10 倍),导致 GAE 估计不稳定。

def normalize_rewards(rewards: torch.Tensor, eps: float = 1e-8) -> torch.Tensor:
"""对每个 batch 的 reward 做 z-score 归一化"""
return (rewards - rewards.mean()) / (rewards.std() + eps)

7.2 Reward Clipping#

将 reward 限制在一个范围内,防止极端 reward 破坏训练:

def clip_rewards(rewards: torch.Tensor, lo: float = -5.0, hi: float = 5.0) -> torch.Tensor:
return torch.clamp(rewards, lo, hi)

7.3 KL Annealing#

训练初期可以允许策略偏离参考模型多一些(快速学习),后期收紧 KL(稳定收敛):

def get_beta(step: int, total_steps: int, beta_init: float = 0.1, beta_final: float = 0.2) -> float:
"""
KL 系数预热。
训练初期 β 较小,后期逐渐增大。
"""
progress = step / total_steps
return beta_init + (beta_final - beta_init) * progress

7.4 Value 网络预热#

在 PPO 更新之前,先单独训练 Value 网络几个 epoch,减少初始 Value 预测的巨大误差:

def pretrain_critic(critic_model, rollout_data, num_warmup_epochs=3):
"""
Value 网络预热:用 rollout 数据的 returns 训练几个 epoch。
"""
returns = rollout_data["returns"]
values_old = rollout_data["values"]
for _ in range(num_warmup_epochs):
loss = ((values_old - returns) ** 2).mean()
loss.backward()
optimizer.step()

7.5 早停机制#

监控 KL 散度,当策略偏离参考模型太远时早停:

def check_early_stop(
policy_log_probs: torch.Tensor,
ref_log_probs: torch.Tensor,
kl_threshold: float = 0.05,
):
"""如果平均 KL 超过阈值,跳过此次更新"""
mean_kl = (policy_log_probs - ref_log_probs).mean()
if mean_kl > kl_threshold:
return True # 早停
return False

8. PPO vs GRPO:两条技术路线#

8.1 GRPO 的核心洞察#

GRPO(DeepSeek-R1)提出:对于 LLM 对齐,不需要 critic 网络。

原因:

  • LLM 的 action space 是离散的 token(A=310|\mathcal{A}| = 3-10万)
  • 每个 prompt 可以采样 GG 个回答,用组内排名作为优势

8.2 GRPO vs PPO 详细对比#

维度PPOGRPO
Value Network需要(估计 baseline)不需要
优势估计GAE(需 Value 网络)组内归一化
模型数量4(Policy + Ref + RM + Critic)3(Policy + Ref + RM)
显存占用~4× LLM~3× LLM
在线探索是(ε-greedy 或 sample)是(sample)
稳定性中(需调 clip、Value loss 等)较高(无需 Value 网络)
性能SOTA(大量调参后)SOTA(DeepSeek-R1 证明)

8.3 GRPO 实现#

def grpo_update(
policy_model,
ref_model,
reward_model,
prompts: list[str],
G: int = 8, # 每个 prompt 采样的回答数
beta: float = 0.04,
clip_ratio: float = 0.2,
):
"""
GRPO 更新(DeepSeek-R1 风格)。
"""
# === 1. 采样 G 个回答 ===
responses = []
policy_log_probs_all = []
ref_log_probs_all = []
for prompt in prompts:
prompt_batch = [prompt] * G
# 生成
with torch.no_grad():
outputs = policy_model.generate(
prompt_batch, max_new_tokens=256, do_sample=True, temperature=0.8
)
# 提取 response
resp = extract_responses(prompt_batch, outputs)
# 计算 log probs
policy_lps = compute_response_log_probs(policy_model, resp)
with torch.no_grad():
ref_lps = compute_response_log_probs(ref_model, resp)
responses.extend(resp)
policy_log_probs_all.append(policy_lps)
ref_log_probs_all.append(ref_lps)
# === 2. RM 评分 ===
rewards = reward_model.batch_score(responses) # (len(prompts), G)
# === 3. 组内归一化得到优势 ===
mean_r = rewards.mean(dim=-1, keepdim=True)
std_r = rewards.std(dim=-1, keepdim=True) + 1e-8
advantages = (rewards - mean_r) / std_r # (num_prompts, G)
# === 4. GRPO 损失 ===
policy_lps_flat = torch.cat(policy_log_probs_all, dim=0) # (num_prompts*G,)
ref_lps_flat = torch.cat(ref_log_probs_all, dim=0)
# 展平优势(每个 response 一个标量优势)
adv_flat = advantages.flatten() # (num_prompts*G,)
# PPO-style loss
grpo_loss = compute_grpo_loss(
policy_lps_flat, ref_lps_flat, adv_flat,
beta=beta, clip_ratio=clip_ratio,
)
return grpo_loss

8.4 何时选 PPO,何时选 GRPO?#

选择 PPO:
✅ 需要追赶最强性能(GPT-4 级别)
✅ 有足够的 GPU 显存(支持 4 模型)
✅ 有能力做细致的超参数调优
✅ 任务需要复杂的多步探索
选择 GRPO:
✅ 数学推理、代码生成(DeepSeek-R1 证明)
✅ 显存受限(3 模型 vs 4 模型)
✅ 希望简化工程复杂度
✅ 快速验证对齐效果

9. 超参数调优指南#

9.1 关键超参数表#

超参数推荐范围说明
clip ϵ\epsilon0.1 ~ 0.3(默认 0.2)太小→学习慢;太大→不稳定
KL 系数 β\beta0.01 ~ 0.2太小→偏离原始模型;太大→学习慢
GAE λ\lambda0.9 ~ 0.95方差-偏差权衡
折扣 γ\gamma1.0LLM 终态任务不需要折扣
PPO epochs4 ~ 10多 epoch 用同一 rollout 数据
Mini-batch 数4 ~ 32分多批更新更稳定
学习率1e-6 ~ 5e-6比 SFT 小 10 倍
熵系数0 ~ 0.01鼓励探索,太大→策略退化
Value loss 系数0.5 ~ 1.0平衡 Policy 和 Value 学习
Gradient clip1.0防止梯度爆炸
Rollout 采样数 GG4 ~ 16(GRPO)越大优势估计越稳定

9.2 调参优先级#

第一优先(影响最大):
1. β(KL 系数)—— 直接控制策略偏离程度
2. clip ε —— 控制每次更新步长
3. 学习率 —— 整体更新幅度
第二优先:
4. GAE λ —— 优势估计质量
5. Value loss 系数 —— Critic 训练速度
第三优先(微调):
6. 熵系数
7. PPO epochs
8. Batch size

9.3 训练不稳定信号与对策#

监控指标危险信号对策
**KL(ππ_ref)**> 0.3 / token
RM reward突然跌到 0 或负回滚 checkpoint,检查 RM
Value loss持续增大预热 Value 网络,增大 vf_coef
Clip fraction> 20% 被 clip增大 clip ε
Policy entropy快速下降→0增大熵系数
Gradient norm> 5减小学习率,增大 gradient clip
def log_training_metrics(step, metrics):
"""关键监控指标"""
print(f"Step {step}: "
f"KL={metrics['mean_kl']:.4f} | "
f"Reward={metrics['mean_reward']:.3f} | "
f"ClipFrac={metrics['clip_fraction']:.2%} | "
f"VF={metrics['value_loss']:.4f} | "
f"Ent={metrics['entropy']:.3f} | "
f"GradNorm={metrics['grad_norm']:.2f}")

10. 核心公式汇总#

10.1 策略梯度定理#

θJ(θ)=Eτπθ[t=0T1θlogπθ(atst)  Gt]\nabla_\theta J(\theta) = \mathbb{E}_{\tau\sim\pi_\theta}\Big[\sum_{t=0}^{T-1}\nabla_\theta\log\pi_\theta(a_t|s_t)\;G_t\Big]

10.2 重要性比率#

rt(θ)=πθ(atst)πθold(atst)r_t(\theta) = \frac{\pi_\theta(a_t|s_t)}{\pi_{\theta_{\text{old}}}(a_t|s_t)}

10.3 PPO Clipped Surrogate#

LCLIP(θ)=Et[min(rt(θ)At, clip(rt(θ),1ϵ,1+ϵ)At)]\mathcal{L}^{\text{CLIP}}(\theta) = \mathbb{E}_t\Big[\min\big(r_t(\theta)\,A_t,\ \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon)\,A_t\big)\Big]

10.4 GAE 递推#

A^t=δt+γλ(1dt)A^t+1\hat{A}_t = \delta_t + \gamma\,\lambda\,(1-d_t)\,\hat{A}_{t+1}

其中 δt=rt+γV(st+1)V(st)\delta_t = r_t + \gamma\,V(s_{t+1}) - V(s_t)dtd_t 是结束信号。

10.5 Value 损失#

LV=E[(RtVϕ(st))2]\mathcal{L}_V = \mathbb{E}\Big[\big(R_t - V_\phi(s_t)\big)^2\Big]

10.6 KL 奖励塑形#

rtshaped=rtβKL(πθ(st)πref(st))r_t^{\text{shaped}} = r_t - \beta\,\text{KL}\big(\pi_\theta(\cdot|s_t)\,\|\,\pi_{\text{ref}}(\cdot|s_t)\big)

10.7 GRPO 组内归一化优势#

Ai=rimean(r)std(r)A_i = \frac{r_i - \text{mean}(\mathbf{r})}{\text{std}(\mathbf{r})}

10.8 TRPO 信赖域约束#

Et[KL(πθold(st)πθ(st))]δ\mathbb{E}_t\Big[\text{KL}\big(\pi_{\theta_{\text{old}}}(\cdot|s_t)\,\|\,\pi_\theta(\cdot|s_t)\big)\Big] \leq \delta

11. 总结#

11.1 PPO 的核心价值#

┌─────────────────────────────────────────────────────────────┐
│ PPO 的三个关键创新 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 1. Clip 机制:无需二阶优化即可实现信赖域约束 │
│ 比 TRPO 简单 10 倍,比 Actor-Critic 稳定 10 倍 │
│ │
│ 2. 置信域思想:限制新旧策略比率的变化范围 │
│ 防止策略在单次更新中崩溃 │
│ │
│ 3. GAE:高效的方差-偏差权衡优势估计 │
│ 一个参数 λ 控制整个谱的估计质量 │
│ │
└─────────────────────────────────────────────────────────────┘

11.2 算法演进图谱#

REINFORCE (1992)
↓ 方差极大
Actor-Critic (1999)
↓ 无步长保证
TRPO (2015)
↓ 二阶优化太复杂
PPO (2017) ←──────────────────────────────┐
↓ │
├── PPO-Penalty(KL 奖励塑形) │
├── PPO with Value Clip │
└── Adaptive KL Penalty │
┌─────────────────┘
GRPO (DeepSeek-R1, 2024)
→ 去掉 Critic,用组内排名替代
→ 在数学推理任务上达到 SOTA

11.3 工程实践建议#

Warning

PPO 是 LLM 对齐中最复杂的训练范式。在决定使用 PPO 之前,请确认:

  1. 你有足够的 GPU 显存(至少 2×80G A100 用于 7B 模型,LoRA 优化)
  2. 你有足够的工程能力处理 4 模型协调、early stopping、gradient checkpointing
  3. 你的目标是追赶最强性能,而非快速验证

如果不确定,先用 DPO 或 GRPO,在验证有效后再升级到 PPO。

11.4 未来方向#

  • Adaptive Clip:根据训练动态自动调整 clip ε
  • V-MPO:基于最大后验的策略优化,比 PPO 更适合离策略
  • Dreamer/Model-based PPO:结合世界模型,减少样本复杂度
  • On-policy → Off-policy:PPO 的 off-policy 变体,减少采样浪费
推荐阅读
  1. PPO 原论文(Schulman et al., 2017)—— 核心思想
  2. TRPO 论文(Schulman et al., 2015)—— 理论基础
  3. GAE 论文(Schulman et al., 2016)—— 优势估计
  4. DeepSeek-R1(2024)—— GRPO 实战
  5. InstructGPT(Ouyang et al., 2022)—— PPO 用于 LLM 对齐的工程实践

参考资料#

  1. Schulman, J., et al. (2017). “Proximal Policy Optimization Algorithms.” arXiv.
  2. Schulman, J., et al. (2015). “Trust Region Policy Optimization.” ICML.
  3. Schulman, J., et al. (2016). “High-Dimensional Continuous Control Using Generalized Advantage Estimation.” ICLR.
  4. Ouyang, L., et al. (2022). “Training language models to follow instructions with human feedback.” NeurIPS.
  5. Bai, Y., et al. (2022). “Training a Helpful and Harmless Assistant with Reinforcement Learning from Human Feedback.” arXiv.
  6. Shao, Z., et al. (2024). “DeepSeekMath: Pushing the Limits of Mathematical Reasoning in Open Language Models.” arXiv (GRPO).
  7. Guo, D., et al. (2025). “DeepSeek-R1: Incentivizing Reasoning Capability in LLMs via Reinforcement Learning.” arXiv.
  8. Engstrom, L., et al. (2020). “Implementation Matters in Deep RL: A Case Study on PPO and TRPO.” ICLR.
  9. Irving, G., et al. (2018). “AI safety via debate.” arXiv.
  10. Lambert, N., et al. (2022). “Tune: A Research Platform for Model Tuning.” GitHub.

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

PPO 深度解析:信赖域优化的艺术与 LLM 对齐实践
https://aiattnstudio.link/posts/ppo/
作者
Federico
发布于
2026-07-16
许可协议
CC BY-NC-SA 4.0
Profile Image of the Author

Federico

AI Research Lab

Hello, I'm Federico.

关于实验室 / About
公告

欢迎来到Federico的个人博客

分类
标签
站点统计
57文章
7分类
404标签
1
1. 引言:PPO 为什么会成为 LLM 对齐的主流?
1.1 强化学习的三代算法
1.2 为什么 LLM 对齐选择 PPO?
2
2. 策略梯度基础:从 REINFORCE 到 Actor-Critic
2.1 策略梯度定理
2.2 REINFORCE:基本策略梯度
2.3 方差问题的根源
2.4 Actor-Critic:引入 Value Network
3
3. TRPO:信赖域优化的理论基础
3.1 信赖域的核心思想
3.2 从约束到惩罚:共轭梯度法
3.3 TRPO 的工程痛点
4
4. PPO:Clip 机制的数学推导
4.1 重要性比率
4.2 原始策略梯度损失
4.3 Clipped Surrogate 目标
4.4 Clip 机制的直觉
4.5 完整的数学推导
4.6 为什么 min 而不是 max?
5
5. GAE:广义优势估计
5.1 TD(λ) 的优势函数版本
5.2 GAE 的递推形式
5.3 GAE 超参数的影响
5.4 Value 网络损失
6
6. PPO 用于 LLM 对齐:完整实现
6.1 四模型架构
6.2 Rollout:在线采样
6.3 PPO 更新
6.4 KL 约束的奖励塑形形式
7
7. 进阶技巧:让 PPO 训练更稳定
7.1 Reward 归一化
7.2 Reward Clipping
7.3 KL Annealing
7.4 Value 网络预热
7.5 早停机制
8
8. PPO vs GRPO:两条技术路线
8.1 GRPO 的核心洞察
8.2 GRPO vs PPO 详细对比
8.3 GRPO 实现
8.4 何时选 PPO,何时选 GRPO?
9
9. 超参数调优指南
9.1 关键超参数表
9.2 调参优先级
9.3 训练不稳定信号与对策
10
10. 核心公式汇总
10.1 策略梯度定理
10.2 重要性比率
10.3 PPO Clipped Surrogate
10.4 GAE 递推
10.5 Value 损失
10.6 KL 奖励塑形
10.7 GRPO 组内归一化优势
10.8 TRPO 信赖域约束
11
11. 总结
11.1 PPO 的核心价值
11.2 算法演进图谱
11.3 工程实践建议
11.4 未来方向
12
参考资料