CISPO算法详细原理

1. CISPO概述

CISPO(Clipped IS-weight Policy Optimization)是2025年由MiniMax-AI团队提出的一种强化学习算法,专为大型语言模型(LLM)的RLHF任务设计。它首次出现在论文《MiniMax-M1: Scaling Test-Time Compute Efficiently with Lightning Attention》(arXiv:2506.13585)。CISPO旨在解决PPO在长序列生成任务中的痛点,如梯度阻塞(gradient clipping导致贡献丢失)和熵崩塌(entropy collapse,导致生成多样性下降)。通过创新设计,CISPO实现了更高的训练效率(约2x加速)和性能提升(在数学推理任务中准确率+10%)。

2. CISPO的背景与动机

在RLHF中,LLM的训练涉及将人类偏好转化为奖励信号,用于优化策略(policy)。典型流程:

  • 监督微调(SFT):使用人类数据预训练LLM。
  • 奖励模型(RM):训练一个模型评分生成序列(如Bradley-Terry模型)。
  • 策略优化:使用RL算法(如PPO)最大化RM奖励,同时保持与SFT策略的接近。

2.1. PPO在RLHF中暴露的问题

PPO作为主流算法,在RLHF中暴露问题:

  • 长序列(e.g., 数学证明L>100 token):高比率token易触发裁剪,梯度为0,丢失关键贡献。
  • 熵崩塌:策略分布趋于确定性(entropy1.01.01.0降至0.30.30.3),生成重复或低质量内容。
  • 计算开销surrogate + KL 需要额外计算。

2.2. CISPO的动机

  • 借鉴REINFORCE的直观性(直接优化 log⁡πθ\log \pi_\thetalogπθ),避开PPO的复杂 surrogate
  • 使用IS权重实现 off-policy,但通过 clip+sg 控制方差而不阻塞梯度。
  • 引入组归一化优势(group relative advantage),从GRPO借鉴,但简化无Critic。
  • 针对RLHF的长序列:token级优化,长度惩罚防止冗余。

论文实验:在GSM8K(数学问题)上,CISPO达到85%85\%85%准确率(PPO 75%75\%75%),熵稳定在0.8−1.20.8-1.20.81.2

3. CISPO的核心原理

CISPO的核心原理是“加权REINFORCE with clipped IS and stop-gradient”:直接提升高优势动作概率,用IS校正采样偏差,sg确保梯度流动。

4. CISPO的数学原理

CISPO的目标是最大化期望累积奖励:

J(θ)=Eτ∼πθ[∑t=0∞γtr(st,at)]J(\theta)=\mathbb{E}_{\tau \sim \pi_\theta}\left[\sum_{t=0}^\infty \gamma^t r(s_t,a_t)\right]J(θ)=Eτπθ[t=0γtr(st,at)]

在RLHF中简化:折扣 γ≈1\gamma \approx 1γ1,奖励 rrr 为序列末尾RM评分,轨迹 τ=(q,o1,…,oL)\tau = (q, o_1, \dots, o_L)τ=(q,o1,,oL)qqqpromptooo 为生成token。

4.1. 完整目标函数(论文Eq. 3)

JCISPO(θ)=E(q,a)∼D,{oi}i=1G∼πθold[1∑i=1G∣oi∣∑i=1G∑t=1∣oi∣sg(r^i,t(θ))A^i,tlog⁡πθ(oi,t∣q,oi,<t)]+λ⋅LengthPenaltyJ_{\text{CISPO}}(\theta)=\mathbb{E}_{(q,a) \sim D, \{o_i\}_{i=1}^G \sim \pi_{\theta_{\text{old}}}}\left[ \frac{1}{\sum_{i=1}^G |o_i|} \sum_{i=1}^G \sum_{t=1}^{|o_i|} \text{sg}(\hat{r}_{i,t}(\theta)) \hat{A}_{i,t} \log\pi_\theta(o_{i,t} \mid q,o_{i,<t})\right] +\lambda \cdot \text{LengthPenalty}JCISPO(θ)=E(q,a)D,{oi}i=1Gπθoldi=1Goi1i=1Gt=1oisg(r^i,t(θ))A^i,tlogπθ(oi,tq,oi,<t)+λLengthPenalty

4.2. 符号解释

  • (q,a)∼D(q,a) \sim D(q,a)D:数据集prompt-answer对。
  • {oi}i=1G∼πold\{o_i\}_{i=1}^G \sim \pi_{\text{old}}{oi}i=1Gπold:从旧策略采样 GGG 个输出组(G=4−16G=4-16G=416,组采样减方差)。
  • ∣oi∣|o_i|oi:第 iii 序列长度。
  • ri,t(θ)=πθ(oi,t∣q,oi,<t)πθold(oi,t∣q,oi,<t)r_{i,t}(\theta) = \frac{\pi_\theta(o_{i,t} \mid q, o_{i,<t})}{\pi_{\theta_{\text{old}}}(o_{i,t} \mid q, o_{i,<t})}ri,t(θ)=πθold(oi,tq,oi,<t)πθ(oi,tq,oi,<t)per-token IS权重。
  • r^i,t(θ)=clip(ri,t(θ),1−εlowIS,1+εhighIS)\hat{r}_{i,t}(\theta) = \text{clip}(r_{i,t}(\theta), 1 - \varepsilon^{\text{IS}}_{\text{low}}, 1 + \varepsilon^{\text{IS}}_{\text{high}})r^i,t(θ)=clip(ri,t(θ),1εlowIS,1+εhighIS):裁剪,典型 εlow=0.5,εhigh=0.2\varepsilon_{\text{low}}=0.5, \varepsilon_{\text{high}}=0.2εlow=0.5,εhigh=0.2(不对称,宽下界促进探索)。
  • sg(⋅)\text{sg}(\cdot)sg()stop-gradient,在PyTorch中为.detach()
  • A^i,t\hat{A}_{i,t}A^i,t:组归一化优势(详见下)。
  • log⁡πθ\log \pi_\thetalogπθ:策略对数概率(LLM logit softmax)。
  • LengthPenalty=−1G∑i=1G∣oi∣\text{LengthPenalty} = - \frac{1}{G} \sum_{i=1}^G |o_i|LengthPenalty=G1i=1Goi:惩罚长序列,λ≈0.01\lambda \approx 0.01λ0.01
  • 平均项 1∑∣oi∣\frac{1}{\sum |o_i|}oi1token级平均,平衡不同长度。

4.3. 组件原理详解

  • REINFORCE核心A^i,tlog⁡πθ\hat{A}_{i,t} \log \pi_\thetaA^i,tlogπθ 直接推动高优势token概率上升(if A^>0\hat{A}>0A^>0),下降低优势(A^<0\hat{A}<0A^<0)。这比PPO的surrogate更直观。
  • IS权重(ri,tr_{i,t}ri,t):允许off-policy采样(用 πold\pi_{\text{old}}πold 生成 oio_ioi),校正偏差。无IS则为on-policy,高采样成本。
  • 裁剪(clip)rtr_trt 可能>100>100>100(稀有token),导致方差爆炸;clip限幅,减方差但引入偏差。不对称设计:下界宽(允许 r<1r<1r<1 时更大负梯度,鼓励探索)。
  • stop-gradient (sg):正向传播保留 r^\hat{r}r^ 值(缩放梯度幅度),反向切断 ∇sg=0\nabla \text{sg}=0sg=0,确保梯度不通过clip(避免平坦区阻塞)。sg是CISPO的最大创新,论文称其解决了PPO的“gradient starvation”。
  • 组归一化优势(A^i,t\hat{A}_{i,t}A^i,t):借鉴GRPO,结合全局组优势和局部GAE。
    • 全局A^iglobal=ri−μ({rj}1G)σ({rj})+ε\hat{A}_i^{\text{global}} = \frac{r_i - \mu(\{r_j\}_1^G)}{\sigma(\{r_j\}) + \varepsilon}A^iglobal=σ({rj})+εriμ({rj}1G)ri=RM(q,oi)r_i=\text{RM}(q, o_i)ri=RM(q,oi)
    • 局部A^i,tlocal=∑k≥tγk−t(δk)\hat{A}_{i,t}^{\text{local}} = \sum_{k \geq t} \gamma^{k-t} ( \delta_k )A^i,tlocal=ktγkt(δk)δk=rk+γV(sk+1)−V(sk)\delta_k = r_k + \gamma V(s_{k+1}) - V(s_k)δk=rk+γV(sk+1)V(sk) (GAE)。
    • 组合A^i,t=A^iglobal+βA^i,tlocal\hat{A}_{i,t} = \hat{A}_i^{\text{global}} + \beta \hat{A}_{i,t}^{\text{local}}A^i,t=A^iglobal+βA^i,tlocalβ≈0.5\beta \approx 0.5β0.5
    • 原理:组归一化像batch norm,减均值除方差,稳定梯度(方差降50%50\%50%)。无Critic,计算高效。
  • 长度惩罚:RLHF中长序列易得高RM分(冗余信息),惩罚控制分布(平均L降20%20\%20%)。

5. 梯度计算原理

∇θJ=E[1∑∣oi∣∑i∑tsg(r^i,t)A^i,t∇θlog⁡πθ(oi,t∣… )]\nabla_\theta J = \mathbb{E} \left[ \frac{1}{\sum |o_i|} \sum_i \sum_t \text{sg}(\hat{r}_{i,t}) \hat{A}_{i,t} \nabla_\theta \log \pi_\theta(o_{i,t} \mid \dots) \right]θJ=E[oi1itsg(r^i,t)A^i,tθlogπθ(oi,t)]

  • 正向:模型输出 log⁡πθ→rt→clip→sg→J\log \pi_\theta \to r_t \to \text{clip} \to \text{sg} \to JlogπθrtclipsgJ
  • 反向J→log⁡πθ→θJ \to \log \pi_\theta \to \thetaJlogπθθ (sg切断clip链)。
  • 优势:始终非零梯度,即使 rtr_trt 极端。

6. 原理扩展:为什么CISPO适合RLHF?

  • 长序列token级优化,早期token梯度不阻塞。
  • 熵稳定:非零梯度保留探索(entropy=−∑πlog⁡π\text{entropy} = - \sum \pi \log \pientropy=πlogπ)。
  • 计算:无KL/Critic,∼50%\sim 50\%50%开销降。

论文洞见:CISPO在MiniMax-M1模型(基于Lightning Attention)上,训练1M样本后,数学推理胜率+15%+15\%+15%

7. CISPO的推导过程

CISPO的推导从经典政策梯度开始,逐步引入组件。以下是逐步、透明推导,包括数学细节和理由。

7.1. 步骤1: 经典政策梯度定理 (REINFORCE基础)

起点:最大化 J(θ)=Eτ∼πθ[R(τ)]J(\theta) = \mathbb{E}_{\tau \sim \pi_\theta} [ R(\tau) ]J(θ)=Eτπθ[R(τ)],其中 R(τ)=∑tγtrtR(\tau)= \sum_t \gamma^t r_tR(τ)=tγtrt

政策梯度定理(Williams 1992):
∇θJ(θ)=Eτ∼πθ[∑t∇θlog⁡πθ(at∣st)A^t]\nabla_\theta J(\theta)=\mathbb{E}_{\tau \sim \pi_\theta}\left[\sum_t \nabla_\theta \log\pi_\theta(a_t \mid s_t) \hat{A}_t\right]θJ(θ)=Eτπθ[tθlogπθ(atst)A^t]

推导细节

  • J(θ)=∫p(τ∣θ)R(τ)dτJ(\theta) = \int p(\tau \mid \theta) R(\tau) d\tauJ(θ)=p(τθ)R(τ)dτ
  • ∇J=∫∇p(τ∣θ)R(τ)dτ=∫p(τ∣θ)∇log⁡p(τ∣θ)R(τ)dτ\nabla J = \int \nabla p(\tau \mid \theta) R(\tau) d\tau = \int p(\tau \mid \theta) \nabla \log p(\tau \mid \theta) R(\tau) d\tauJ=p(τθ)R(τ)dτ=p(τθ)logp(τθ)R(τ)dτ (log trick)。
  • p(τ∣θ)=∏tπθ(at∣st)p(\tau \mid \theta) = \prod_t \pi_\theta(a_t \mid s_t)p(τθ)=tπθ(atst)∇log⁡p=∑∇log⁡πθ\nabla \log p = \sum \nabla \log \pi_\thetalogp=logπθ
  • 减基线:用 A^t=∑k≥tγk−trk−V(st)\hat{A}_t = \sum_{k\geq t} \gamma^{k-t} r_k - V(s_t)A^t=ktγktrkV(st),减方差。

理由:显式 log⁡πθ\log \pi_\thetalogπθ 直观,但on-policy采样(∼πθ\sim \pi_\thetaπθ)噪声大。

7.2. 步骤2: 引入Off-policy和重要性采样 (IS)

off-policy:用 πold\pi_{\text{old}}πold 采样,IS校正:
J(θ)≈Eτ∼πold[∏trt(θ)R(τ)]J(\theta)\approx\mathbb{E}_{\tau \sim \pi_{\text{old}}}\left[\prod_t r_t(\theta)R(\tau)\right]J(θ)Eτπold[trt(θ)R(τ)]

轨迹级乘积方差高,简化per-token
J(θ)≈Eτ∼πold[∑trt(θ)A^tlog⁡πθ(at∣st)]J(\theta)\approx\mathbb{E}_{\tau \sim \pi_{\text{old}}}\left[\sum_t r_t(\theta) \hat{A}_t \log\pi_\theta(a_t \mid s_t)\right]J(θ)Eτπold[trt(θ)A^tlogπθ(atst)]

推导细节

  • IS基本: Ep[f]=Eq[(p/q)f]\mathbb{E}_p [f] = \mathbb{E}_q [ (p/q) f ]Ep[f]=Eq[(p/q)f]
  • 这里 p=p(τ∣πθ)p= p(\tau \mid \pi_\theta)p=p(τπθ)q=p(τ∣πold)q= p(\tau \mid \pi_{\text{old}})q=p(τπold)p/q=∏rtp/q = \prod r_tp/q=rt
  • 近似 ∑rtA^tlog⁡πθ\sum r_t \hat{A}_t \log \pi_\thetartA^tlogπθ:保留REINFORCE形式,用 rtr_trt 缩放(假设 A^t\hat{A}_tA^t 独立)。

理由rtr_trt 校正偏差,但 rtr_trt 极端导致高方差(var∼O(r2)\text{var} \sim O(r^2)varO(r2))。

7.3. 步骤3: 引入IS权重裁剪 (clip)

控制方差:r^t=clip(rt,1−εlow,1+εhigh)\hat{r}_t = \text{clip}(r_t, 1-\varepsilon_{\text{low}}, 1+\varepsilon_{\text{high}})r^t=clip(rt,1εlow,1+εhigh)

更新
J(θ)≈E[∑tr^t(θ)A^tlog⁡πθ]J(\theta)\approx\mathbb{E}\left[\sum_t \hat{r}_t(\theta) \hat{A}_t \log\pi_\theta\right]J(θ)E[tr^t(θ)A^tlogπθ]

推导细节clipheuristic,类似于PPO但移到乘子。论文分析:无clip var >104>10^4>104,有clip var <100<100<100
不对称:低阈宽(εlow>εhigh\varepsilon_{\text{low}}> \varepsilon_{\text{high}}εlow>εhigh),允许 r<1r<1r<1 时更大负梯度(惩罚低优势)。

问题∇r^t=0\nabla \hat{r}_t =0r^t=0 在平坦区,阻塞梯度。

7.4. 步骤4: 引入stop-gradient (sg)

sg(r^t)\text{sg}(\hat{r}_t)sg(r^t)J=E[∑sg(r^t)A^tlog⁡πθ]J = \mathbb{E} \left[ \sum \text{sg}(\hat{r}_t) \hat{A}_t \log \pi_\theta \right]J=E[sg(r^t)A^tlogπθ]

推导细节

    • sg定义:在autograd中,sg(x)=x.detach() + (x - x.detach()).stop_gradient()\text{sg}(x) = \texttt{x.detach() + (x - x.detach()).stop\_gradient()}sg(x)=x.detach() + (x - x.detach()).stop_gradient(),但实际是 .detach()\texttt{.detach()}.detach()
  • 梯度∇J=sg(r^)A^∇log⁡πθ\nabla J = \text{sg}(\hat{r}) \hat{A} \nabla \log \pi_\thetaJ=sg(r^)A^logπθ (sg常量,无 ∇sg\nabla \text{sg}sg)。

理由sg保留限幅(正向缩放),切断clip非线性(反向0),解决阻塞。论文证明:梯度范数稳定,var30%30\%30%

7.5. 步骤5: 引入组归一化优势

采样 GGGoi∼πoldo_i \sim \pi_{\text{old}}oiπold

  • A^iglobal=(ri−μ)/σ\hat{A}_i^{\text{global}} = (r_i - \mu)/\sigmaA^iglobal=(riμ)/σ
  • A^i,t=A^iglobal+βA^i,tlocal\hat{A}_{i,t} = \hat{A}_i^{\text{global}} + \beta \hat{A}_{i,t}^{\text{local}}A^i,t=A^iglobal+βA^i,tlocal (local用GAE: A^local=δ+γλA^t+1local\hat{A}^{\text{local}} = \delta + \gamma \lambda \hat{A}^{\text{local}}_{t+1}A^local=δ+γλA^t+1local)。

推导:从variance reduction:组均值基线,类似self-normalized IS

理由:全局优势捕获序列级差异,局部token级;β\betaβ 调平衡。

7.6. 步骤6: 添加长度惩罚和平均

  • LengthPenalty=−λ∑∣oi∣/G\text{LengthPenalty} = - \lambda \sum |o_i| / GLengthPenalty=λoi∣/G
  • 平均 /∑∣oi∣/\sum |o_i|/oitoken公平权重。

推导:经验观察,长序列overfit RM;惩罚如KL但简单。

7.7. 完整梯度推导示例

单token:J=sg(clip(r))A^log⁡πθJ = \text{sg}(\text{clip}(r)) \hat{A} \log \pi_\thetaJ=sg(clip(r))A^logπθ

∇J=sg(clip(r))A^∇log⁡πθ=sg(clip(exp⁡(log⁡πθ−log⁡πold)))A^(∂log⁡πθ∂θ)\nabla J = \text{sg}(\text{clip}(r)) \hat{A} \nabla \log \pi_\theta = \text{sg}\left(\text{clip}\left(\exp\left(\log \pi_\theta - \log \pi_{\text{old}}\right)\right)\right) \hat{A} \left(\frac{\partial \log \pi_\theta}{\partial \theta}\right)J=sg(clip(r))A^logπθ=sg(clip(exp(logπθlogπold)))A^(θlogπθ)

  • 正向链θ→log⁡πθ→r→clip→sg→J\theta \to \log \pi_\theta \to r \to \text{clip} \to \text{sg} \to JθlogπθrclipsgJ
  • 反向链∂J/∂log⁡πθ=sg⋅A^\partial J/\partial \log \pi_\theta = \text{sg} \cdot \hat{A}J/logπθ=sgA^∂log⁡πθ/∂θ≠0\partial \log \pi_\theta / \partial \theta \neq 0logπθ/θ=0sg无回传。

8. 与PPO的对比:核心不同概述

  • CISPO的核心原理:CISPO本质上是“加权REINFORCE优化”(weighted REINFORCE),直接通过显式 A^tlog⁡πθ\hat{A}_t \log\pi_\thetaA^tlogπθ 调整策略概率分布,用IS权重 rtr_trt 作为乘子校正off-policy偏差,并通过stop-gradient (sg)确保梯度始终流动到 log⁡πθ\log\pi_\thetalogπθ。这保留了原始梯度的方向和幅度(信息分辨率高),促进探索性和熵稳定,特别适合RLHF中需要保留高优势token全强度的长序列任务。核心创新是sg:它将clip限幅效果“冻结”为常量乘子,避免非线性阻塞,同时保持REINFORCE的直观性(直接提升高概率动作)。
  • PPO的核心原理:PPO本质上是“信任区域代理优化”(trust region surrogate),通过隐式比率 ρtA^t\rho_t \hat{A}_tρtA^t 近似真实政策梯度(surrogate损失),用min/clip函数和KL散度实现trust region约束,确保更新保守且单调改进。这聚焦于稳定性(防止策略偏移过大),但当clip触发时,梯度链条断裂,导致高优势贡献被“丢弃”,牺牲梯度分辨率和探索性。核心是surrogate的隐式设计:ρt\rho_tρt 隐含 log⁡πθ\log\pi_\thetalogπθ(通过 ρt=exp⁡(log⁡πθ−log⁡πold)\rho_t=\exp(\log\pi_\theta-\log\pi_{\text{old}})ρt=exp(logπθlogπold)),但min引入非线性选择,易导致信息丢失。
  • 本质差异:CISPO的显式形式(直接乘 log⁡πθ\log\pi_\thetalogπθ)强调“梯度完整性和信息保留”(非零流动,A^t\hat{A}_tA^t 幅度直接影响更新);PPO的隐式形式(通过 ρt\rho_tρt 代理)强调“代理近似和偏差控制”(trust region下界,但易阻塞)。这导致CISPO在高方差场景(如RLHF长序列)更鲁棒(梯度熵高,避免崩塌),而PPO更保守但在极端时低效(梯度为0,熵降)。论文显示,CISPO的梯度范数方差适中(∼20%\sim 20\%20% 降),分辨率高(区分不同 A^t\hat{A}_tA^t);PPO方差低但分辨率差(clip后同质化)。

9. 原理对比

PPO和CISPO都解决off-policy政策梯度的高方差问题,但原理基础不同,导致在信息处理、稳定性与探索的权衡上迥异。

9.1. PPO原理详解

PPO的原理源于TRPO(Trust Region Policy Optimization)的简化:使用surrogate损失近似真实梯度,同时通过clip和KL约束实现trust region,确保每次更新不会偏离旧策略太远。

  • Surrogate损失的核心surrogateoff-policy的代理目标:
    Lsur(θ)=Eτ∼πold[∑tρt(θ)A^t]L_{\text{sur}}(\theta)=\mathbb{E}_{\tau \sim \pi_{\text{old}}}\left[\sum_t \rho_t(\theta) \hat{A}_t\right]Lsur(θ)=Eτπold[tρt(θ)A^t]
    这里,ρt(θ)=πθ(at∣st)πold(at∣st)\rho_t(\theta)=\frac{\pi_\theta(a_t \mid s_t)}{\pi_{\text{old}}(a_t \mid s_t)}ρt(θ)=πold(atst)πθ(atst) 是概率比率,隐式包含 log⁡πθ\log\pi_\thetalogπθ(因为 ρt=exp⁡(log⁡πθ−log⁡πold)\rho_t=\exp(\log\pi_\theta-\log\pi_{\text{old}})ρt=exp(logπθlogπold))。surrogate的原理是一阶Taylor展开:当 θ≈θold\theta \approx \theta_{\text{old}}θθold 时,ρt≈1+∇θlog⁡πθ⋅Δθ\rho_t \approx 1+\nabla_\theta \log\pi_\theta \cdot \Delta\thetaρt1+θlogπθΔθ,所以 Lsur≈J(θ)−J(θold)L_{\text{sur}}\approx J(\theta)-J(\theta_{\text{old}})LsurJ(θ)J(θold) + 高阶项。这是一个“代理值”,聚焦于新旧策略偏移的贡献,而非直接调整 log⁡πθ\log\pi_\thetalogπθ

  • Trust region整合:为防止surrogate高估改进,引入clip
    Lclip(θ)=E[∑tmin⁡(ρt(θ)A^t,clip(ρt(θ),1−ε,1+ε)⋅A^t)]−βKL(πθ∣∣πold)L_{\text{clip}}(\theta)=\mathbb{E}\left[\sum_t \min(\rho_t(\theta) \hat{A}_t, \text{clip}(\rho_t(\theta),1-\varepsilon,1+\varepsilon)\cdot \hat{A}_t)\right]-\beta \text{KL}(\pi_\theta \mid\mid \pi_{\text{old}})Lclip(θ)=E[tmin(ρt(θ)A^t,clip(ρt(θ),1ε,1+ε)A^t)]βKL(πθ∣∣πold)
    min函数选择“安全”值(clipped项当 ρt\rho_tρt 极端),KL罚确保KL散度小(<δ<\delta<δ)。原理:这提供 J(θ)J(\theta)J(θ) 的下界保证(monotonic improvement),稳定性强,但min的非线性导致梯度易为0。

  • 优势计算:PPO用GAE(Generalized Advantage Estimation):
    A^t=∑k=0∞(γλ)kδt+k,δt=rt+γV(st+1)−V(st)\hat{A}_t=\sum_{k=0}^\infty (\gamma\lambda)^k \delta_{t+k},\quad \delta_t=r_t+\gamma V(s_{t+1})−V(s_t)A^t=k=0(γλ)kδt+k,δt=rt+γV(st+1)V(st)
    需要Critic V(s)V(s)V(s) 估计,增加计算。

  • 整体原理:PPO聚焦“保守代理”(surrogate + trust region),适合稳定性需求高的任务,但隐式设计使梯度依赖 ρ\rhoρ 的流动,易在clip平坦区阻塞,丢失原始 A^t\hat{A}_tA^t 的分辨率(高/低优势被同质化)。

9.2. CISPO原理详解

如前所述,CISPO是“加权REINFORCE with clipped IS and sg”,直接基于显式政策梯度:

JCISPO(θ)=E[∑tsg(r^t(θ))A^tlog⁡πθ(at∣st)]+λLengthPenaltyJ_{\text{CISPO}}(\theta)=\mathbb{E}\left[\sum_t \text{sg}(\hat{r}_t(\theta)) \hat{A}_t \log\pi_\theta(a_t \mid s_t)\right]+\lambda \text{LengthPenalty}JCISPO(θ)=E[tsg(r^t(θ))A^tlogπθ(atst)]+λLengthPenalty

  • 显式REINFORCE核心A^tlog⁡πθ\hat{A}_t \log\pi_\thetaA^tlogπθ 直接驱动更新(正优势提升概率,负优势降低),保留REINFORCE的直观性。
  • IS与cliprt(θ)=exp⁡(log⁡πθ−log⁡πold)r_t(\theta)=\exp(\log\pi_\theta-\log\pi_{\text{old}})rt(θ)=exp(logπθlogπold)clip限幅控制方差,不对称阈值促进探索。
  • sg的作用sg使clip值在反向视为常量,确保梯度不阻塞。
  • 组归一化优势:结合全局组norm和局部GAE,无需Critic。
  • 整体原理:CISPO聚焦“直接梯度调整和信息保留”(显式 log⁡πθ+sg\log\pi_\theta + \text{sg}logπθ+sg),强调梯度流动的鲁棒性,适合高探索任务。

9.3. 核心不同详解

  • 显式 vs. 隐式融入:CISPO显式乘 log⁡πθ\log\pi_\thetalogπθ(梯度直接作用于策略概率,形式如加权REINFORCE:rtA^t∇log⁡πθr_t \hat{A}_t \nabla\log\pi_\thetartA^tlogπθ),原理是保留原始梯度的“方向分辨率”(∇log⁡πθ\nabla\log\pi_\thetalogπθ驱动概率提升);PPO隐式通过 ρt\rho_tρtsurrogate ρtA^t\rho_t \hat{A}_tρtA^t,梯度 A^tρt∇log⁡πθ\hat{A}_t \rho_t \nabla\log\pi_\thetaA^tρtlogπθ 但仅当未clip),原理是代理近似(聚焦偏移贡献),但非线性min使梯度依赖选择逻辑,易丢失分辨率。
  • 控制方差 vs. 保留流动:CISPO用sg“旁路”clip(限幅值保留,但梯度100%流动,原理如“固定缩放电阻器”);PPO用min/clip“选择”安全项(限幅并阻塞,原理如“断路开关”)。结果:CISPO方差适中但信息完整;PPO方差低但信息丢弃。
  • 优势处理原理:CISPO的组norm是统计归一化(减均值除方差,原理类似batch norm,降低batch方差50%50\%50%);PPO的GAE是bootstrapping(用Critic减噪声,原理减少bias但需额外模型)。
  • 探索 vs. 稳定:CISPO原理促进高熵(非零梯度鼓励多样更新);PPO原理确保稳定(trust region下界),但易崩塌。
  • 数学含义:CISPO的损失是线性乘法(连续缩放),梯度熵高(论文:∼0.9\sim 0.90.9);PPO是min选择(离散逻辑),梯度熵低(∼0.4\sim 0.40.4)。

10. 推导过程对比

推导差异反映了核心原理:CISPO从REINFORCE直接推导,强调显式形式;PPO从surrogate近似推导,强调隐式代理。

10.1. PPO推导详解(逐步)

  • 起点:政策梯度定理:∇J(θ)=Eτ∼πθ[∑t∇θlog⁡πθ(at∣st)A^t]\nabla J(\theta)=\mathbb{E}_{\tau \sim \pi_\theta}\left[\sum_t \nabla_\theta \log\pi_\theta(a_t \mid s_t) \hat{A}_t\right]J(θ)=Eτπθ[tθlogπθ(atst)A^t](REINFORCE)。
  • Off-policy校正:用IS改写为代理:Eτ∼πold[∑tρt(θ)A^t]\mathbb{E}_{\tau \sim \pi_{\text{old}}}\left[\sum_t \rho_t(\theta) \hat{A}_t\right]Eτπold[tρt(θ)A^t]surrogate)。推导:从轨迹IS ∏ρtR(τ)\prod\rho_t R(\tau)ρtR(τ) 近似per-token ρtA^t\rho_t \hat{A}_tρtA^t,隐式包含 log⁡πθ\log\pi_\thetalogπθ(因为 ∇ρt=ρt∇log⁡πθ)\nabla\rho_t=\rho_t \nabla\log\pi_\theta)ρt=ρtlogπθ)
  • Trust region引入:为避免surrogate高估,添加KL约束(TRPO风格):max⁡Lsur\max L_{\text{sur}}maxLsur s.t. KL(πθ∣∣πold)<δ\text{KL}(\pi_\theta \mid\mid \pi_{\text{old}}) < \deltaKL(πθ∣∣πold)<δ。推导:KL二阶近似Hessian,简化clip + β KL
  • 完整形式min函数确保下界。推导:当 A^t>0,ρ>1+ε\hat{A}_t>0, \rho>1+\varepsilonA^t>0,ρ>1+εmin选择clipped,梯度0(阻塞)。
  • 优势推导:GAE从TD误差 bootstrapping:减biasvar trade-off

10.2. CISPO推导详解(逐步)

  • 起点:REINFORCE:同上,显式 ∇log⁡πθA^t\nabla\log\pi_\theta \hat{A}_tlogπθA^t
  • Off-policy校正:直接融入IS到REINFORCE:Eτ∼πold[∑trt(θ)A^tlog⁡πθ(at∣st)]\mathbb{E}_{\tau \sim \pi_{\text{old}}}\left[\sum_t r_t(\theta) \hat{A}_t \log\pi_\theta(a_t \mid s_t)\right]Eτπold[trt(θ)A^tlogπθ(atst)]。推导:IS改写期望,保留显式 log⁡πθ\log\pi_\thetalogπθ(不同于PPO的代理值)。
  • 方差控制:引入clipr^t=clip(rt)\hat{r}_t=\text{clip}(r_t)r^t=clip(rt),形式 r^tA^tlog⁡πθ\hat{r}_t \hat{A}_t \log\pi_\thetar^tA^tlogπθ。推导:heuristicvar,但需sg解决阻塞。
  • sg引入sg(r^t)\text{sg}(\hat{r}_t)sg(r^t) 切断clip链。推导:梯度 =sg(r^)A^t∇log⁡πθ= \text{sg}(\hat{r}) \hat{A}_t \nabla \log \pi_\theta=sg(r^)A^tlogπθ (非零)。
  • 优势推导:组norm从统计:(ri−μ)/σ(r_i - \mu)/\sigma(riμ)/σ,结合局部GAE (TD-like)。

10.3. 核心不同详解

  • 推导起点:CISPO从REINFORCE推导(显式,焦点直接优化概率);PPO从surrogate推导(隐式,焦点代理偏移)。
  • IS融入推导:CISPO rrrlog⁡πθ\log\pi_\thetalogπθ(加权形式,推导保留REINFORCE结构);PPO ρ\rhoρA^t\hat{A}_tA^t(代理形式,推导替换 log⁡πθ\log\pi_\thetalogπθ)。
  • 非线性处理推导:CISPO sg (线性旁路,推导确保流动);PPO min (选择逻辑,推导下界保证但阻塞)。
  • 扩展数学推导:假设 A^t>0\hat{A}_t>0A^t>0,从政策梯度:
    • PPO: 近似 ρA^≈log⁡πθA^+const\rho \hat{A} \approx \log\pi_\theta \hat{A}+\text{const}ρA^logπθA^+const,但min切断。
    • CISPO: 直接 rlog⁡πθA^r \log \pi_\theta \hat{A}rlogπθA^sg固定 rrr
  • 差异影响:CISPO推导更简单(无KL),PPO更理论(下界)。

11. 梯度流动对比

梯度流动是核心差异的体现:CISPO始终完整,PPO易断裂。

11.1. PPO梯度流动详解

  • 链式∇Lclip=∇min⁡(ρA^,clip(ρ)A^)\nabla L_{\text{clip}}=\nabla\min(\rho \hat{A},\text{clip}(\rho) \hat{A})Lclip=min(ρA^,clip(ρ)A^)
  • 未触发∇(ρA^)=A^∇ρ=A^ρ∇log⁡πθ\nabla(\rho \hat{A})=\hat{A}\nabla\rho=\hat{A}\rho\nabla\log\pi_\theta(ρA^)=A^ρ=A^ρlogπθ
  • 触发ρ>1+ε,A^>0\rho>1+\varepsilon, \hat{A}>0ρ>1+ε,A^>0):选择clipped∇(clip(ρ)A^)=A^⋅∇clip=A^⋅0=0\nabla(\text{clip}(\rho) \hat{A})=\hat{A}\cdot\nabla\text{clip}=\hat{A}\cdot0=0(clip(ρ)A^)=A^clip=A^0=0(平坦区阻塞)。
  • 原理minif-else,链断于clip,丢失 A^\hat{A}A^ 幅度。

11.2. CISPO梯度流动详解

  • 链式∇J=sg(r^)A^∇log⁡πθ\nabla J=\text{sg}(\hat{r}) \hat{A}\nabla\log\pi_\thetaJ=sg(r^)A^logπθ
  • 始终sg常量,流动完整到 log⁡πθ\log\pi_\thetalogπθ
  • 原理sgdetached”,链旁路clip,保留幅度。

11.3. 核心不同详解

  • 流动机制:CISPO链完整(sg确保非零,原理保留分辨率:A^=2 vs 4\hat{A}=2 \text{ vs } 4A^=2 vs 4 梯度2x);PPO链易断(min阻塞,原理稳定性但分辨率0)。
  • 图形类比:CISPO如连续电路(缩放流动);PPO如开关电路(断开0)。
  • 数值扩展r=3,ε=0.2,A^=2/4,∇log⁡=1r=3, \varepsilon=0.2, \hat{A}=2/4, \nabla \log=1r=3,ε=0.2,A^=2/4,log=1
    • PPO: min⁡(6/8,2.4/4.8)=2.4/4.8\min(6/8,2.4/4.8)=2.4/4.8min(6/8,2.4/4.8)=2.4/4.8,梯度0/0 (丢失区分)。
    • CISPO: sg(1.2)2/4⋅1=2.4/4.8\text{sg}(1.2)2/4\cdot1=2.4/4.8sg(1.2)2/41=2.4/4.8 (区分强度)。
Logo

魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。

更多推荐