JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

PPO:業界主力的裁剪法

近端策略優化丟掉硬求解器,只用一個裁剪目標就把更新壓小——簡單、穩健、無所不在。

從 TRPO 的痛到 PPO 的簡單

近端策略優化(Proximal Policy Optimization,PPO) 問了一個問題:我們能不能只靠一階最佳化——對小批次做普通 SGD、不用共軛梯度、不用線搜尋——就拿到 TRPO 的穩定性?答案是把 替代目標函數 巧妙地重塑,讓它自我限制步伐大小。你用平常那套 Adam 就能最佳化它,還能對每批資料跑好幾個 epoch。

PPO 完全摒弃 TRPO 的二阶机制:它用朴素的一阶 SGD 改进策略——在小批量上做许多微小的梯度下降步骤。

梯度下降沿损失曲线逐步走向最小值,代表 PPO 的朴素一阶优化。

裁剪替代目標

令 r 是某個取樣動作的 重要性比值(importance ratio)——新策略機率除以舊策略——A 是它的優勢。裁剪替代目標(clipped surrogate objective) 取兩項的最小值:一是平常的 r·A,二是把 r 先裁剪到區間 [1−ε, 1+ε] 再乘上 A 的版本。取 min 就是挑較小(較悲觀)的那個。

L^{\mathrm{CLIP}}(\theta)=\mathbb{E}_t\!\left[\min\!\left(r_t(\theta)\,\hat{A}_t,\ \operatorname{clip}\!\big(r_t(\theta),\,1-\epsilon,\,1+\epsilon\big)\,\hat{A}_t\right)\right]

裁剪替代目标函数:对原始项与裁剪后的比率乘优势项取最小值。

精妙之處在於這個不對稱。當動作很好(A 為正)而比值升過 1+ε 時,裁剪會把目標壓平——再把該動作機率推得更高也得不到獎勵了,於是梯度消失、策略停止移動。當動作很差(A 為負)時,min 讓未裁剪那項主導,使策略仍能強力壓制它。整體效果模仿了信賴域:更新可以自由地修正錯誤,卻被勒住、不會對單一好樣本過度押注。

裁剪 vs 懲罰

PPO 其實有兩種版本,理解 裁剪 vs 懲罰 的區別很重要。裁剪版就是上面那個。懲罰版改成在目標裡加一個 KL 項,並用 自適應 KL 懲罰(adaptive KL penalty) 來縮放它:每次更新後,若量到的 KL 超出目標就調高係數、不足就調低。這正是 TRPO 硬約束、以及第 2 篇鏡像下降模板的「直接軟化版表親」。

在控制任務上,裁剪版贏得人氣,因為它不需要調 KL 目標、又簡單到不行。但懲罰版才是在語言模型對齊裡重新浮現的那一個——在那裡,「對參考策略維持受控的 KL」正是重點所在;第 5 篇你會再遇見它。

带 KL 惩罚的 PPO 变体在语言模型的 RLHF 策略微调阶段重新登场。

RLHF 流程:人类偏好训练奖励模型,再用其微调策略。

真正攸關成敗的細節

PPO 的名聲藏著一個不太舒服的真相:它測得的表現,有極大一部分不是來自裁剪目標,而是來自一疊不起眼的 實作細節。對這些做消融的研究發現,它們的重要性可能不亞於核心演算法。請把它們當成方法的一部分,而非可有可無的修飾。

  1. 對每個小批次的優勢做正規化(零均值、單位變異數),以穩定梯度尺度。
  2. 用同樣方式裁剪價值函數損失,並把它加權進共享的 actor-critic 損失。
  3. 加入熵獎勵以維持探索,防止策略過早崩塌。
  4. 對學習率做退火、裁剪全域梯度範數,並對每批資料跑數個 epoch。
  5. 使用正交權重初始化,以及獎勵/觀測正規化。
# One PPO update over a collected batch
for epoch in range(K):
    for mb in minibatches(batch):
        ratio = exp(logp_new(mb.a, mb.s) - mb.logp_old)
        adv   = normalize(mb.advantage)
        unclipped = ratio * adv
        clipped   = clip(ratio, 1 - eps, 1 + eps) * adv
        policy_loss = -mean(min(unclipped, clipped))
        value_loss  = mse(value(mb.s), mb.returns)
        loss = policy_loss + c1 * value_loss - c2 * entropy(mb.s)
        adam_step(loss)            # plain first-order; no CG, no line search
PPO-Clip 的精髓:對裁剪/未裁剪比值取 min,加上價值項與熵項,用普通 SGD 最佳化。