Skip to content
alarik.me
Go back

OPD 深度解析(三):知识蒸馏与 OPD 梯度推导

本文转载自:微信公众号「菜菜Vibe Coding」

这是本系列的第三篇文章,讲述知识蒸馏的内容,以及 OPD 和强化学习的关系;后面几章的内容有些简略,后续应该还会有几篇来详细的讲解,用来讲述学习 OPD 中遇到的问题。

本文档将分为 3 个部分来介绍:

  1. 数学基础、信息论基础和 KL 散度系列第一篇
  2. 大模型微调(SFT)、强化学习策略梯度系列第二篇
  3. OPD 的介绍和强化学习的关系(本篇)

Table of contents

Open Table of contents

第七章:知识蒸馏(Off-policy KD)

7.1 什么是知识蒸馏 (Knowledge Distillation, KD)?

知识蒸馏的核心思想是:将大模型(Teacher)学到的”知识”迁移到小模型(Student)中,使小模型在保持较低计算成本的同时,逼近甚至达到大模型的性能。知识蒸馏也是模型压缩的一种方法。

7.2 Off-policy KD Loss

在自回归语言模型中,Off-policy(异策) 意味着 Student 模型是在固定的数据 prefix(通常由 Teacher 生成或来自真实数据集)上进行训练,而不是基于 Student 自己之前生成的 token(On-policy)进行训练。

其 Logits 级别的蒸馏 Loss 通常定义为:

LKD,t=vVPT(vctdata)logPS(vctdata)\mathcal{L}_{\mathrm{KD},t} = -\sum_{v \in \mathcal{V}} P_T(v|c_t^{data}) \log P_S(v|c_t^{data})

根据信息论,交叉熵可以分解为:

LKD,t=H(PT)+DKL(PTPS)\mathcal{L}_{\mathrm{KD},t} = H(P_T) + D_{\mathrm{KL}}(P_T \| P_S)
Important

等价于最小化 Forward KL:因为 Teacher 的分布 PTP_T 是固定的,其熵 H(PT)H(P_T) 对 Student 的参数 θ\theta 来说是常数。因此,最小化上述 Loss 严格等价于最小化 Forward KL 散度 DKL(PTPS)D_{\mathrm{KL}}(P_T \| P_S)

7.3 温度系数 (Temperature, TT) 及其作用

在计算 Soft Label 时,通常会引入一个温度系数 TT 来对 Teacher 的 Logits 进行平滑处理:

PT(vc)=exp(zv/T)uVexp(zu/T)P_T(v|c) = \frac{\exp(z_v / T)}{\sum_{u \in \mathcal{V}} \exp(z_u / T)}

为什么要加入温度系数 TT

7.4 为什么 KD Loss 需要乘以 T2T^2(详细数学推导)

在 Hinton 等人的经典论文中,计算 Soft Label Loss 时会乘以 T2T^2。其核心原因是:为了抵消温度系数 TT 变大带来的梯度衰减,保持梯度的量级(Scale)与 T=1T=1 时一致,防止模型学不到东西。

以下是变量严格一致的详细推导过程:

1. 变量定义

2. 第一步:求 Loss 对 Student logit zkz_k 的偏导数

由于 Teacher 的分布 pip_i 是固定的(视为常数),我们对 L\mathcal{L} 关于 zkz_k 求导:

Lzk=i=1Vpilogqizk\frac{\partial \mathcal{L}}{\partial z_k} = - \sum_{i=1}^V p_i \frac{\partial \log q_i}{\partial z_k}

根据链式法则,logqizk=1qiqizk\frac{\partial \log q_i}{\partial z_k} = \frac{1}{q_i} \frac{\partial q_i}{\partial z_k}。对于 Softmax 函数 qiq_i,其对 zkz_k 的导数为标准形式:

qizk=1Tqi(δikqk)\frac{\partial q_i}{\partial z_k} = \frac{1}{T} q_i (\delta_{ik} - q_k)

(其中 δik\delta_{ik} 是 Kronecker delta,当 i=ki=k 时为 1,否则为 0)

代入上式:

logqizk=1qi1Tqi(δikqk)=1T(δikqk)\frac{\partial \log q_i}{\partial z_k} = \frac{1}{q_i} \cdot \frac{1}{T} q_i (\delta_{ik} - q_k) = \frac{1}{T} (\delta_{ik} - q_k)

将其代回 Loss 的导数公式:

Lzk=i=1Vpi[1T(δikqk)]=1T(i=1Vpiδiki=1Vpiqk)\frac{\partial \mathcal{L}}{\partial z_k} = - \sum_{i=1}^V p_i \left[ \frac{1}{T} (\delta_{ik} - q_k) \right] = - \frac{1}{T} \left( \sum_{i=1}^V p_i \delta_{ik} - \sum_{i=1}^V p_i q_k \right)

因为 i=1Vpiδik=pk\sum_{i=1}^V p_i \delta_{ik} = p_k,且概率之和 i=1Vpi=1\sum_{i=1}^V p_i = 1,所以:

Lzk=1T(pkqk)=1T(qkpk)\frac{\partial \mathcal{L}}{\partial z_k} = - \frac{1}{T} (p_k - q_k) = \frac{1}{T} (q_k - p_k)

(注意:此时梯度中已经出现了一个 1T\frac{1}{T} 的衰减因子)

3. 第二步:对概率差值 (qkpk)(q_k - p_k) 进行泰勒展开

TT 较大时,uiT\frac{u_i}{T}ziT\frac{z_i}{T} 都趋近于 0。我们可以利用一阶泰勒展开 ex1+xe^x \approx 1 + x

展开 Student 的 probability qkq_k

qk=exp(zk/T)j=1Vexp(zj/T)1+zk/Tj=1V(1+zj/T)=1+zk/TV+1Tj=1Vzjq_k = \frac{\exp(z_k / T)}{\sum_{j=1}^V \exp(z_j / T)} \approx \frac{1 + z_k / T}{\sum_{j=1}^V (1 + z_j / T)} = \frac{1 + z_k / T}{V + \frac{1}{T}\sum_{j=1}^V z_j}

关键假设:在 Softmax 中,对所有 logits 减去同一个常数不会改变概率分布。因此,我们可以假设 logits 已经过均值中心化 (mean-centered),即 j=1Vzj=0\sum_{j=1}^V z_j = 0j=1Vuj=0\sum_{j=1}^V u_j = 0

在此假设下,分母简化为 VV,因此:

qk1+zk/TV=1V+zkVTq_k \approx \frac{1 + z_k / T}{V} = \frac{1}{V} + \frac{z_k}{V \cdot T}

同理,Teacher 的 probability pkp_k 也可以展开为:

pk1V+ukVTp_k \approx \frac{1}{V} + \frac{u_k}{V \cdot T}

计算两者的差值:

qkpk(1V+zkVT)(1V+ukVT)=zkukVTq_k - p_k \approx \left( \frac{1}{V} + \frac{z_k}{V \cdot T} \right) - \left( \frac{1}{V} + \frac{u_k}{V \cdot T} \right) = \frac{z_k - u_k}{V \cdot T}

(注意:这里又出现了一个 1T\frac{1}{T} 的衰减因子)

4. 第三步:代回梯度公式,得出 T2T^2 的必要性

将第二步的结果代入第一步的梯度公式中:

Lzk1T(zkukVT)=1VT2(zkuk)\frac{\partial \mathcal{L}}{\partial z_k} \approx \frac{1}{T} \left( \frac{z_k - u_k}{V \cdot T} \right) = \frac{1}{V \cdot T^2} (z_k - u_k)
Important

结论分析:从最终公式可以看出,当 TT 较大时,原始交叉熵 Loss 产生的梯度与 1T2\frac{1}{T^2} 成正比。如果 T=5T=5,梯度会缩小 25 倍;如果 T=10T=10,梯度会缩小 100 倍。这会导致 Student 模型的参数更新停滞。

补偿方案:如果我们在 Loss 前面乘以 T2T^2,定义缩放后的 Loss 为 Lscaled=T2L\mathcal{L}_{scaled} = T^2 \mathcal{L},那么新的梯度为:

Lscaledzk=T2LzkT21VT2(zkuk)=1V(zkuk)\frac{\partial \mathcal{L}_{scaled}}{\partial z_k} = T^2 \frac{\partial \mathcal{L}}{\partial z_k} \approx T^2 \cdot \frac{1}{V \cdot T^2} (z_k - u_k) = \frac{1}{V} (z_k - u_k)

乘以 T2T^2 后,梯度恰好近似为 1V(zkuk)\frac{1}{V} (z_k - u_k)。这等价于最小化 Student logits zkz_k 和 Teacher logits uku_k 之间的均方误差 (MSE) 的梯度。这保证了无论 TT 取何值,梯度的量级都能与 T=1T=1 时保持一致,从而维持稳定的学习率。

5. 工程实践避坑指南 (PyTorch)

# 正确示范
T = 5.0
teacher_probs = F.softmax(teacher_logits / T, dim=-1)
student_log_probs = F.log_softmax(student_logits / T, dim=-1)
# 1. 计算 KL 散度 (或 Cross Entropy)
kl_loss = F.kl_div(student_log_probs, teacher_probs, reduction='batchmean')
# 2. 关键:乘以 T^2 来补偿梯度衰减!
loss_kd = (T ** 2) * kl_loss

7.5 基于特征的蒸馏 (Feature-based Distillation)

除了上述基于输出 Logits 的蒸馏(Response-level KD),基于特征的蒸馏是另一种极其重要的 Off-policy 方法。它要求 Student 的中间层表示(Representation)去拟合 Teacher 的中间层表示。

由于 Teacher 的特征是预先计算好并固定的,这同样属于 Off-policy 范式。常见的特征蒸馏包括:

Lhidden=1Ni=1NWhS(i)hT(i)22\mathcal{L}_{\mathrm{hidden}} = \frac{1}{N} \sum_{i=1}^{N} \left\| \mathbf{W} \mathbf{h}_S^{(i)} - \mathbf{h}_T^{(i)} \right\|_2^2

(其中 W\mathbf{W} 是用于对齐维度的投影矩阵,NN 为序列长度)

Lattn=1LNl=1Li=1NAS(l,i)AT(l,i)22\mathcal{L}_{\mathrm{attn}} = \frac{1}{L \cdot N} \sum_{l=1}^{L} \sum_{i=1}^{N} \left\| \mathbf{A}_S^{(l, i)} - \mathbf{A}_T^{(l, i)} \right\|_2^2

(其中 LL 为层数,A\mathbf{A} 为注意力权重矩阵)

为什么需要基于特征的蒸馏?

7.6 总结:Off-policy KD 的核心优势

将 Logits 蒸馏与特征蒸馏结合,构成了现代 LLM 压缩的标准 Off-policy 范式:

第八章:Exposure Bias 问题

8.1 核心矛盾

训练 prefix = 数据/teacher 的 (x,y<tdata)(x, y_{<{t}}^{data})

推理 prefix = student 自己的 (x,y<tstudent)(x, y_{<{t}}^{student})

一旦 student 前面走偏,后续进入训练未见的区域 → 错误连锁放大。

8.2 生活比喻

学开车:教练总带你走标准路线(SFT训练)。自己开车时第3个路口就偏离了,从第4个路口开始从未练习过 → 开始迷路。

8.3 OPD 解决 Exposure Bias

让训练也发生在 student 自己生成的 prefix 上 → 训练和推理的前缀分布一致。

第九章:OPD 的核心思想

9.1 OPD 的一句话定义

On-Policy Distillation(OPD,同策略蒸馏):

学生模型先用自己的当前策略生成回答,再让教师模型在这些学生自己生成的轨迹上提供监督信号,学生据此更新。

形式上:

9.2 OPD 与 SFT/KD 的根本区别

SFT / Off-policy KDOPD
Prefix 来源数据/teacher(固定的)Student 自己生成(动态的)
训练 prefixct=(x,y<tdata)c_t = (x, y_{<{t}}^{data})ct=(x,y<tstudent)c_t = (x, y_{<{t}}^{student})
是否有 exposure bias有(训练≠推理)无(训练=推理)

9.3 OPD 的三个核心维度

维度选项关键问题
Prefix 来源dataset / teacher / student是 off-policy 还是 on-policy?
Teacher 信号粒度sampled-token / top-k / full-vocab看 1 个、K 个还是全部 token?
优化方式direct loss / policy gradient直接反传还是把 KL 当 advantage?

第十章:Sequence-level OPD 梯度推导

10.1 目标函数

最自然的 OPD 目标——reverse KL(mode-seeking):

J(θ)=yπθ(y)[logπθ(y)logq(y)]J(\theta) = \sum_y \pi_\theta(y)\left[\log\pi_\theta(y) - \log q(y)\right]

写成期望:

Jx(θ)=Eyπθ(x)[logπθ(yx)logq(yx)]J_x(\theta) = \mathbb{E}_{y \sim \pi_\theta(\cdot|x)}\left[\log\pi_\theta(y|x) - \log q(y|x)\right]
Note

为什么是 reverse KL?因为期望在 student 分布上取——我们只能看到 student 会生成哪些轨迹,在这些轨迹上评估差距。

10.2 梯度推导——逐步详解

这一步是为了介绍 OPD 和 RL 的关系。

第一步:乘积求导法则

J(θ)J(\theta)πθ(y)\pi_\theta(y)[logπθ(y)logq(y)][\log\pi_\theta(y) - \log q(y)] 的乘积之和。两者都依赖 θ\thetalogq(y)\log q(y) 不依赖 θ\theta),所以用乘积求导法则:

θ[AB]=(A)B+A(B)\nabla_\theta [A \cdot B] = (\nabla A) \cdot B + A \cdot (\nabla B)

应用到求和内部:

θ[πθ(y)(logπθ(y)logq(y))]\nabla_\theta\left[\pi_\theta(y) \cdot \left(\log\pi_\theta(y) - \log q(y)\right)\right] =πθ(y)(logπθ(y)logq(y))+πθ(y)(logπθ(y))= \nabla\pi_\theta(y) \cdot \left(\log\pi_\theta(y) - \log q(y)\right) + \pi_\theta(y) \cdot \nabla\left(\log\pi_\theta(y)\right)

注意:logq(y)=0\nabla\log q(y) = 0(teacher 不依赖 θ\theta

对所有 yy 求和:

J=yπθ(y)[logπθ(y)logq(y)]+yπθ(y)logπθ(y)\nabla J = \sum_y \nabla\pi_\theta(y)\left[\log\pi_\theta(y) - \log q(y)\right] + \sum_y \pi_\theta(y) \nabla\log\pi_\theta(y)

第二步:Log-derivative trick 处理第一项

πθ(y)=πθ(y)logπθ(y)\nabla\pi_\theta(y) = \pi_\theta(y) \cdot \nabla\log\pi_\theta(y)

推导πθ(y)=elogπθ(y)\pi_\theta(y) = e^{\log\pi_\theta(y)}elogπθ(y)=elogπθ(y)logπθ(y)=πθ(y)logπθ(y)\nabla e^{\log\pi_\theta(y)} = e^{\log\pi_\theta(y)} \cdot \nabla\log\pi_\theta(y) = \pi_\theta(y) \cdot \nabla\log\pi_\theta(y)

代入第一项:

yπθ(y)[logπθ(y)logq(y)]=yπθ(y)logπθ(y)[logπθ(y)logq(y)]\sum_y \nabla\pi_\theta(y)\left[\log\pi_\theta(y) - \log q(y)\right] = \sum_y \pi_\theta(y) \cdot \nabla\log\pi_\theta(y) \cdot \left[\log\pi_\theta(y) - \log q(y)\right]

写成期望(yπθ(y)()=Eyπθ[]\sum_y \pi_\theta(y) \cdot (\cdots) = \mathbb{E}_{y \sim \pi_\theta}[\cdots]):

=Eyπθ[(logπθ(y)logq(y))logπθ(y)]= \mathbb{E}_{y \sim \pi_\theta}\left[\left(\log\pi_\theta(y) - \log q(y)\right) \cdot \nabla\log\pi_\theta(y)\right]

第三步:第二项为什么等于 0?

Eyπθ[logπθ(y)]=yπθ(y)logπθ(y)\mathbb{E}_{y \sim \pi_\theta}\left[\nabla\log\pi_\theta(y)\right] = \sum_y \pi_\theta(y) \nabla\log\pi_\theta(y)

用 log-derivative trick 反向:πθ(y)logπθ(y)=πθ(y)\pi_\theta(y) \nabla\log\pi_\theta(y) = \nabla\pi_\theta(y)

yπθ(y)logπθ(y)=yπθ(y)=(yπθ(y))=1=0\sum_y \pi_\theta(y) \nabla\log\pi_\theta(y) = \sum_y \nabla\pi_\theta(y) = \nabla\left(\sum_y \pi_\theta(y)\right) = \nabla 1 = 0
Important

关键:概率分布归一性 yπθ(y)=1\sum_y \pi_\theta(y) = 11=0\nabla 1 = 0。每个时间步的 softmax 保证概率和为 1,联合分布也归一化,所以总和对参数的梯度恒为 0。

第四步:合并

J=Eyπθ[(logπθ(y)logq(y))logπθ(y)]+0\nabla J = \mathbb{E}_{y \sim \pi_\theta}\left[\left(\log\pi_\theta(y) - \log q(y)\right) \cdot \nabla\log\pi_\theta(y)\right] + 0 J=Eyπθ[(logπθ(y)logq(y))logπθ(y)]\nabla J = \mathbb{E}_{y \sim \pi_\theta}\left[\left(\log\pi_\theta(y) - \log q(y)\right) \cdot \nabla\log\pi_\theta(y)\right]

10.3 完整推导链路图

J(θ) = Σ_y π_θ(y)[logπ_θ(y) - logq(y)] ← reverse KL展开
│ 乘积求导法则:∇(A·B) = (∇A)·B + A·(∇B)
│ 注意:logq(y) 不依赖θ,所以 ∇logq(y) = 0
∇J = Σ_y ∇π_θ(y)[logπ_θ(y) - logq(y)] ← 第一项
+ Σ_y π_θ(y) ∇logπ_θ(y) ← 第二项
│ Log-derivative trick:∇π_θ(y) = π_θ(y)·∇logπ_θ(y)
│ 因为 π = e^(logπ),所以 ∇e^(logπ) = e^(logπ)·∇logπ = π·∇logπ
第一项 = Σ_y π_θ(y)·∇logπ_θ(y)·[logπ_θ(y) - logq(y)]
= E_{y~π_θ}[(logπ_θ(y) - logq(y))·∇logπ_θ(y)] ← 写成期望
│ 第二项证明 = 0:
│ E_{y~π_θ}[∇logπ_θ(y)]
│ = Σ_y π_θ(y)·∇logπ_θ(y)
│ = Σ_y ∇π_θ(y) ← 再次用 log-derivative trick 反向
│ = ∇(Σ_y π_θ(y)) ← 求和与梯度交换
│ = ∇1 ← 概率分布求和 = 1(归一性)
│ = 0 ✓
∇J = E_{y~π_θ}[(logπ_θ(y) - logq(y))·∇logπ_θ(y)] ← 最终 policy-gradient 形式

10.4 结果的含义

J=Eyπθ[(logπθ(y)logq(y))reward/权重logπθ(y)score function]\nabla J = \mathbb{E}_{y \sim \pi_\theta}\left[\underbrace{\left(\log\pi_\theta(y) - \log q(y)\right)}_{\text{reward/权重}} \cdot \underbrace{\nabla\log\pi_\theta(y)}_{\text{score function}}\right]

这就是经典的 REINFORCE / policy-gradient 结构:score function × reward。因此就可以将 OPD 和 RL 进行比较。

第十一章:Return-to-Go 与方差分析

11.1 Log-ratio 的定义

rt=logπθ(ytct)q(ytct)=logπθ(ytct)logq(ytct)r_t = \log\frac{\pi_\theta(y_t|c_t)}{q(y_t|c_t)} = \log\pi_\theta(y_t|c_t) - \log q(y_t|c_t) gt=θlogπθ(ytct)g_t = \nabla_\theta \log\pi_\theta(y_t|c_t)

log-ratio 就是 student 和 teacher 对同一个 token 的 logprob 之差。

rtr_t 的值含义
rt>0r_t > 0student 比 teacher 更支持这个 token
rt=0r_t = 0完全一致
rt<0r_t < 0teacher 比 student 更支持

整条轨迹的 log-ratio = 所有 token 的 log-ratio 之和:

logπθ(yx)q(yx)=t=1Trt\log\frac{\pi_\theta(y|x)}{q(y|x)} = \sum_{t'=1}^{T} r_{t'}

11.2 直接的 Sequence-level Estimator

在 on-policy 蒸馏(OPD)框架下,学生策略 πθ\pi_\theta 与环境交互并收集轨迹,同时以教师策略提供的信号作为奖励。直接的序列级梯度估计器定义为:

g^seq=t=1T(t=1Trt)gt\hat{g}_{\mathrm{seq}} = \sum_{t=1}^{T}\left(\sum_{t'=1}^{T} r_{t'}\right) g_t

该估计器用整条轨迹的总回报 R=t=1TrtR = \sum_{t'=1}^T r_{t'} 同时加权所有时间步的梯度,以此提升高回报轨迹的出现概率。

问题:因果性违背与方差增大

从公式可以明显看出,第 tt 步的梯度 gtg_t所有时间步的奖励 rtr_{t'} 加权,其中包含了 t<tt' < t 的奖励。也就是说,动作发生之前已经获得的奖励被用来评价这个动作的好坏。这违背了时间因果性:当前动作不可能影响过去,过去奖励的大小不应成为强化或抑制当前动作的依据。

将总回报拆分为”过去奖励”与”未来奖励”:

R=Rpast(t)+Rfuture(t),Rpast(t)=t=1t1rt,Rfuture(t)=t=tTrtR = R_{\text{past}}^{(t)} + R_{\text{future}}^{(t)}, \qquad R_{\text{past}}^{(t)} = \sum_{t'=1}^{t-1} r_{t'}, \quad R_{\text{future}}^{(t)} = \sum_{t'=t}^{T} r_{t'}

对第一项求期望。给定 sts_t 时,过去奖励 Rpast(t)R_{\text{past}}^{(t)} 与当前动作 ata_t 无关,因此:

Eatπθ[Rpast(t)gt]=Rpast(t)Eat[θlogπθ(atst)]=0\mathbb{E}_{a_t \sim \pi_\theta}\big[ R_{\text{past}}^{(t)} g_t \big] = R_{\text{past}}^{(t)} \,\mathbb{E}_{a_t}\big[\nabla_\theta \log \pi_\theta(a_t \mid s_t)\big] = 0
Important

这说明包含过去奖励的项在期望下为零,不贡献梯度信号,只凭空增加方差。尽管 g^seq\hat{g}_{\mathrm{seq}} 仍是真实策略梯度的无偏估计,但其方差远大于只使用未来奖励的版本,会导致训练收敛缓慢且不稳定。

正确的因果估计器:Reward-to-Go

应将估计器修正为只让每个动作对其之后发生的奖励负责,即使用”奖励到结束”(reward-to-go):

g^causal=t=1T(t=tTrt)gt\hat{g}_{\mathrm{causal}} = \sum_{t=1}^{T}\left(\sum_{t'=t}^{T} r_{t'}\right) g_t

这正是经典 REINFORCE 算法中不带折扣的形式。它保留了无偏性,并因移除了与动作无关的过去奖励而显著降低了方差。

进一步的方差缩减

即使采用 reward-to-go,面对长序列时方差仍可能很大。常用改进手段包括:

11.4 方差分析

Token-level:O(T2)O(T^2)

g^tok=t=1Trtgt\hat{g}_{\mathrm{tok}} = \sum_{t=1}^{T} r_t \cdot g_t

TT 项求和,每项方差 O(1)O(1),协方差最坏 O(T2)O(T^2) → 总方差 O(T2)O(T^2)

Sequence-level:O(T4)O(T^4)

g^seq=t=1TRtgt\hat{g}_{\mathrm{seq}} = \sum_{t=1}^{T} R_t \cdot g_t

权重 RtR_t 本身是 O(T)O(T)rtr_{t'} 的求和 → RtO(T)R|R_t| \le O(T) \cdot R → 每项方差 O(T2)O(T^2) → 总方差 O(T4)O(T^4)

实际影响

T=50T=50T=4000T=4000
token-level O(T2)O(T^2)~2,500~16,000,000
sequence-level O(T4)O(T^4)~6,250,000~256,000,000,000,000

这就是为什么长 reasoning 任务中 sequence-level OPD “容易炸”。

11.5 折扣形式:γ\gamma 连接两者

g^γ=t=1T(t=tTγttrt)gt\hat{g}_\gamma = \sum_{t=1}^{T}\left(\sum_{t'=t}^{T} \gamma^{t'-t} r_{t'}\right) g_t
γ\gamma含义方差
γ=0\gamma = 0token-level OPDO(T2)O(T^2),有偏但稳定
γ=1\gamma = 1sequence-level OPDO(T4)O(T^4),无偏但高方差
0<γ<10 < \gamma < 1折中介于两者

γ<1\gamma < 1 时权重 RtγR_t^\gamma 的上界变为常数 R1γ\frac{R}{1-\gamma},总方差回到 O(T2)O(T^2)

第十二章:Token-level OPD

12.1 定义

只保留当前 token 的 rtr_t

g^tok=t=1Trtgt\hat{g}_{\mathrm{tok}} = \sum_{t=1}^{T} r_t \cdot g_t

相对 sequence-level 有偏(丢掉了未来项 t=t+1Trt\sum_{t'=t+1}^{T} r_{t'}),但方差从 O(T4)O(T^4) 降到 O(T2)O(T^2)

12.2 Advantage 形式

At=logq(ytct)logπθ(ytct)=rtA_t = \log q(y_t|c_t) - \log\pi_\theta(y_t|c_t) = -r_t

Policy-gradient 更新:

JAtlogπθ(ytct)\nabla J \approx A_t \cdot \nabla\log\pi_\theta(y_t|c_t)
Important

这就是 RL-style OPD:只需 teacher 对 sampled token 的一个 logprob。

第十三章:三种粒度——sampled-token、top-k、full-vocab

13.1 为什么粒度选择只在 on-policy 时有意义?

Off-policy KD 天然 full-vocab:prefix 固定,teacher/student 同步 forward,full logits 直接可用。

On-policy OPD 需要讨论粒度:student 先生成 prefix → teacher 需要在新 prefix 上额外 forward → teacher 信号返回量成为工程瓶颈。

Off-policyOn-policy
Prefix固定,预先就有Student 实时生成
Teacher forward和 student 一起跑需要额外在新 prefix 上跑
Full logits天然可用需要额外获取,代价高

13.2 Sampled-token OPD

Student rollout yπθy \sim \pi_\theta,只比较 sampled token:

At=logq(ytct)logπθ(ytct)A_t = \log q(y_t|c_t) - \log\pi_\theta(y_t|c_t)

13.3 Top-K OPD

Teacher 在 prefix ctc_t 上返回词表中概率最高的 KK 个 token:

St=TopKq(ct)VS_t = \mathrm{TopK}_q(c_t) \subset \mathcal{V}
Note

Top-K 是从词表中选,不是从 prefix 中选。Prefix 是已生成的历史(固定输入),Top-K 选择的是”下一个 token”的候选集。

在支持集内重归一化:

q^(vct)=q(vct)uStq(uct),π^(vct)=πθ(vct)uStπθ(uct)\hat{q}(v|c_t) = \frac{q(v|c_t)}{\sum_{u \in S_t} q(u|c_t)}, \quad \hat{\pi}(v|c_t) = \frac{\pi_\theta(v|c_t)}{\sum_{u \in S_t} \pi_\theta(u|c_t)}
Note

为什么重归一化? 只看 KK 个 token 时概率总和不到 1,重归一化让它们在子集上重新变成合法分布(总和=1),才能算有意义的 KL。

局部 reverse KL:

LtopK(ct)=vStπ^(vct)logπ^(vct)q^(vct)\mathcal{L}_{\mathrm{topK}}(c_t) = \sum_{v \in S_t} \hat{\pi}(v|c_t) \log \frac{\hat{\pi}(v|c_t)}{\hat{q}(v|c_t)}

13.4 Full-vocab OPD

每个 prefix 上比较整个词表分布:

Lt=vVPS(vct)logPS(vct)PT(vct)=DKL(PSPT)\mathcal{L}_t = \sum_{v \in \mathcal{V}} P_S(v|c_t) \log \frac{P_S(v|c_t)}{P_T(v|c_t)} = D_{\mathrm{KL}}(P_S \| P_T)

13.5 三者对比

粒度Teacher 返回信息量成本稳定性
sampled-token1 个 logprob
top-kK 个 logprob + K 个 id较好
full-vocab全部 logits最好但昂贵

第十四章:k1/k2/k3 估计器

14.1 为什么需要估计器?

只有 sampled token 的 logprob 时,无法直接计算精确 KL。需要单样本估计器。

14.2 Reverse KL 的期望形式

DKL(PSPT)=EyPS[logPS(yct)PT(yct)]D_{\mathrm{KL}}(P_S \| P_T) = \mathbb{E}_{y \sim P_S}\left[\log\frac{P_S(y|c_t)}{P_T(y|c_t)}\right]

14.3 k1(无偏,可负)

k1=logPS(ytct)PT(ytct),ytPSk_1 = \log\frac{P_S(y_t|c_t)}{P_T(y_t|c_t)}, \quad y_t \sim P_S

无偏(E[k1]=DKL(PSPT)\mathbb{E}[k_1] = D_{\mathrm{KL}}(P_S \| P_T)),单样本可正可负,方差高。

14.4 k2(有偏,非负)

k2=12(logPS(ytct)PT(ytct))2k_2 = \frac{1}{2}\left(\log\frac{P_S(y_t|c_t)}{P_T(y_t|c_t)}\right)^2

始终非负,方差低。是 KL 的局部二阶近似(当 PSP_SPTP_T 接近时偏差小)。

14.5 k3(无偏,非负)

k3=PT(ytct)PS(ytct)logPT(ytct)PS(ytct)1k_3 = \frac{P_T(y_t|c_t)}{P_S(y_t|c_t)} - \log\frac{P_T(y_t|c_t)}{P_S(y_t|c_t)} - 1

r=PT/PSr = P_T/P_S,则 k3=rlogr10k_3 = r - \log r - 1 \ge 0

无偏:利用 EyPS[PT(y)/PS(y)]=1\mathbb{E}_{y \sim P_S}[P_T(y)/P_S(y)] = 1(因为 yPS(y)PT(y)/PS(y)=yPT(y)=1\sum_y P_S(y) \cdot P_T(y)/P_S(y) = \sum_y P_T(y) = 1),可证 E[k3]=DKL(PSPT)\mathbb{E}[k_3] = D_{\mathrm{KL}}(P_S \| P_T)

14.6 对比

估计器表达式无偏?非负?方差
k1logPSPT\log\frac{P_S}{P_T}
k212(logPSPT)2\frac{1}{2}(\log\frac{P_S}{P_T})^2
k3PTPSlogPTPS1\frac{P_T}{P_S} - \log\frac{P_T}{P_S} - 1

14.7 k3+(Straight-through)

Forward 数值走 k3,backward 梯度走 k2:

forward_score = k3 # 数值用 k3(无偏、非负)
backward_score = k2 # 梯度用 k2(局部更准确)
return backward_score - backward_score.detach() + forward_score.detach()

Share this post:

Previous Post
eLLM:用算力换显存的自适应 KV Cache
Next Post
OPD 深度解析(一):概率、熵与 KL 散度