大模型CISPO算法详细原理万字解析
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,丢失关键贡献。 - 熵崩塌:策略分布趋于确定性(
entropy从1.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.8−1.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),qqq 为 prompt,ooo 为生成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∼πθold∑i=1G∣oi∣1i=1∑Gt=1∑∣oi∣sg(r^i,t(θ))A^i,tlogπθ(oi,t∣q,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=4−16,组采样减方差)。
- ∣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,t∣q,oi,<t)πθ(oi,t∣q,oi,<t):
per-tokenIS权重。 - 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=−G1∑i=1G∣oi∣:惩罚长序列,λ≈0.01\lambda \approx 0.01λ≈0.01。
- 平均项 1∑∣oi∣\frac{1}{\sum |o_i|}∑∣oi∣1:
token级平均,平衡不同长度。
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}=0∇sg=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=∑k≥tγk−t(δ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[∑∣oi∣1i∑t∑sg(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πθ→rt→clip→sg→J。
- 反向:J→logπθ→θJ \to \log \pi_\theta \to \thetaJ→logπθ→θ (
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πθ(at∣st)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(τ∣θ)∇logp(τ∣θ)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\tau∇J=∫∇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πθ(at∣st),∇logp=∑∇logπθ\nabla \log p = \sum \nabla \log \pi_\theta∇logp=∑∇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=∑k≥tγk−trk−V(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[t∏rt(θ)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[t∑rt(θ)A^tlogπθ(at∣st)]
推导细节:
- 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_\theta∑rtA^tlogπθ:保留REINFORCE形式,用 rtr_trt 缩放(假设 A^t\hat{A}_tA^t 独立)。
理由:rtr_trt 校正偏差,但 rtr_trt 极端导致高方差(var∼O(r2)\text{var} \sim O(r^2)var∼O(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[t∑r^t(θ)A^tlogπθ]
推导细节:clip是heuristic,类似于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 =0∇r^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_\theta∇J=sg(r^)A^∇logπθ (
sg常量,无 ∇sg\nabla \text{sg}∇sg)。
理由:sg保留限幅(正向缩放),切断clip非线性(反向0),解决阻塞。论文证明:梯度范数稳定,var降30%30\%30%。
7.5. 步骤5: 引入组归一化优势
采样 GGG 组 oi∼π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|/∑∣oi∣:
token公平权重。
推导:经验观察,长序列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πθ→r→clip→sg→J。
- 反向链:∂J/∂logπθ=sg⋅A^\partial J/\partial \log \pi_\theta = \text{sg} \cdot \hat{A}∂J/∂logπθ=sg⋅A^,∂logπθ/∂θ≠0\partial \log \pi_\theta / \partial \theta \neq 0∂logπθ/∂θ=0;
sg无回传。
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损失的核心:surrogate是off-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(at∣st)πθ(at∣st) 是概率比率,隐式包含 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ρt≈1+∇θlogπθ⋅Δθ,所以 Lsur≈J(θ)−J(θold)L_{\text{sur}}\approx J(\theta)-J(\theta_{\text{old}})Lsur≈J(θ)−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[t∑min(ρ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[t∑sg(r^t(θ))A^tlogπθ(at∣st)]+λLengthPenalty
- 显式REINFORCE核心:A^tlogπθ\hat{A}_t \log\pi_\thetaA^tlogπθ 直接驱动更新(正优势提升概率,负优势降低),保留REINFORCE的直观性。
- IS与
clip:rt(θ)=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^t∇logπθ),原理是保留原始梯度的“方向分辨率”(∇logπθ\nabla\log\pi_\theta∇logπθ驱动概率提升);PPO隐式通过 ρt\rho_tρt(
surrogateρtA^t\rho_t \hat{A}_tρtA^t,梯度 A^tρt∇logπθ\hat{A}_t \rho_t \nabla\log\pi_\thetaA^tρt∇logπθ 但仅当未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.9∼0.9);PPO是
min选择(离散逻辑),梯度熵低(∼0.4\sim 0.4∼0.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πθ(at∣st)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=ρt∇logπθ)。Trust region引入:为避免surrogate高估,添加KL约束(TRPO风格):maxLsur\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:减bias增var trade-off。
10.2. CISPO推导详解(逐步)
- 起点:REINFORCE:同上,显式 ∇logπθA^t\nabla\log\pi_\theta \hat{A}_t∇logπθ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πθ(at∣st)]。推导:IS改写期望,保留显式 logπθ\log\pi_\thetalogπθ(不同于PPO的代理值)。- 方差控制:引入
clip:r^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πθ。推导:heuristic减var,但需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^t∇logπθ (非零)。- 优势推导:组
norm从统计:(ri−μ)/σ(r_i - \mu)/\sigma(ri−μ)/σ,结合局部GAE (TD-like)。
10.3. 核心不同详解
- 推导起点:CISPO从REINFORCE推导(显式,焦点直接优化概率);PPO从
surrogate推导(隐式,焦点代理偏移)。 - IS融入推导:CISPO rrr 乘 logπθ\log\pi_\thetalogπθ(加权形式,推导保留REINFORCE结构);PPO ρ\rhoρ 乘 A^t\hat{A}_tA^t(代理形式,推导替换 logπθ\log\pi_\thetalogπθ)。
- 非线性处理推导:CISPO
sg(线性旁路,推导确保流动);PPOmin(选择逻辑,推导下界保证但阻塞)。 - 扩展数学推导:假设 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。
- PPO: 近似 ρA^≈logπθA^+const\rho \hat{A} \approx \log\pi_\theta \hat{A}+\text{const}ρA^≈logπθA^+const,但
- 差异影响: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(平坦区阻塞)。 - 原理:
min如if-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_\theta∇J=sg(r^)A^∇logπθ。
- 始终:
sg常量,流动完整到 logπθ\log\pi_\thetalogπθ。 - 原理:
sg“detached”,链旁路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/4⋅1=2.4/4.8 (区分强度)。
魔乐社区(Modelers.cn) 是一个中立、公益的人工智能社区,提供人工智能工具、模型、数据的托管、展示与应用协同服务,为人工智能开发及爱好者搭建开放的学习交流平台。社区通过理事会方式运作,由全产业链共同建设、共同运营、共同享有,推动国产AI生态繁荣发展。
更多推荐


所有评论(0)