本文转载自:微信公众号「菜菜Vibe Coding」
这是本系列的第三篇文章,讲述知识蒸馏的内容,以及 OPD 和强化学习的关系;后面几章的内容有些简略,后续应该还会有几篇来详细的讲解,用来讲述学习 OPD 中遇到的问题。
本文档将分为 3 个部分来介绍:
- 数学基础、信息论基础和 KL 散度(系列第一篇)
- 大模型微调(SFT)、强化学习策略梯度(系列第二篇)
- OPD 的介绍和强化学习的关系(本篇)
Table of contents
Open Table of contents
第七章:知识蒸馏(Off-policy KD)
7.1 什么是知识蒸馏 (Knowledge Distillation, KD)?
知识蒸馏的核心思想是:将大模型(Teacher)学到的”知识”迁移到小模型(Student)中,使小模型在保持较低计算成本的同时,逼近甚至达到大模型的性能。知识蒸馏也是模型压缩的一种方法。
- SFT (Supervised Fine-Tuning) 的局限:SFT 通常只看”标准答案”(Hard Label / One-hot 编码)。它只告诉模型”正确答案是什么”,而忽略了其他错误答案之间的相对关系。
- KD 的优势 (Soft Label):蒸馏不仅看标准答案,还看 Teacher 模型对所有可能答案的评分分布(Soft Label)。这个概率分布中包含了丰富的 “暗知识”(Dark Knowledge)。例如,Teacher 可能预测某输入为”猫”的概率是 0.8,“狗”是 0.15,“汽车”是 0.05。这告诉 Student:“猫和狗有相似性,而和汽车差异很大”,这种类别间的结构化关系是 One-hot 标签无法提供的。
7.2 Off-policy KD Loss
在自回归语言模型中,Off-policy(异策) 意味着 Student 模型是在固定的数据 prefix(通常由 Teacher 生成或来自真实数据集)上进行训练,而不是基于 Student 自己之前生成的 token(On-policy)进行训练。
其 Logits 级别的蒸馏 Loss 通常定义为:
LKD,t=−v∈V∑PT(v∣ctdata)logPS(v∣ctdata)
根据信息论,交叉熵可以分解为:
LKD,t=H(PT)+DKL(PT∥PS)
等价于最小化 Forward KL:因为 Teacher 的分布 PT 是固定的,其熵 H(PT) 对 Student 的参数 θ 来说是常数。因此,最小化上述 Loss 严格等价于最小化 Forward KL 散度 DKL(PT∥PS)。
7.3 温度系数 (Temperature, T) 及其作用
在计算 Soft Label 时,通常会引入一个温度系数 T 来对 Teacher 的 Logits 进行平滑处理:
PT(v∣c)=∑u∈Vexp(zu/T)exp(zv/T)
为什么要加入温度系数 T
- 放大”暗知识”的相对关系:当 T>1 时,Softmax 的输出分布会变得更加平滑(均匀)。原本极小的概率值会被相对放大,使得非目标类别(错误类别)之间的概率差异变得明显,从而让学生模型能更清晰地学习到类别间的细粒度相似性。
- 防止 Student 过拟合:如果 T=1,Teacher 的分布可能已经非常尖锐(接近 One-hot),Student 直接拟合容易退化为普通的 SFT。较大的 T 提供了更柔和的监督信号,具有正则化效果。
- 训练与推理的匹配:在训练时使用 T>1 计算 Loss,但在 Student 推理时通常将 T 设回 1,以恢复模型的判别能力。
7.4 为什么 KD Loss 需要乘以 T2(详细数学推导)
在 Hinton 等人的经典论文中,计算 Soft Label Loss 时会乘以 T2。其核心原因是:为了抵消温度系数 T 变大带来的梯度衰减,保持梯度的量级(Scale)与 T=1 时一致,防止模型学不到东西。
以下是变量严格一致的详细推导过程:
1. 变量定义
- V:词表大小 (Vocabulary size),索引为 i,j,k∈{1,2,…,V}。
- T:温度系数 (Temperature)。
- ui:Teacher 模型在词 i 上的原始 logit。
- zi:Student 模型在词 i 上的原始 logit。
- pi:Teacher 的 soft probability,pi=∑j=1Vexp(uj/T)exp(ui/T)。
- qi:Student 的 soft probability,qi=∑j=1Vexp(zj/T)exp(zi/T)。
- L:蒸馏 Loss (Cross-Entropy),L=−∑i=1Vpilogqi。
2. 第一步:求 Loss 对 Student logit zk 的偏导数
由于 Teacher 的分布 pi 是固定的(视为常数),我们对 L 关于 zk 求导:
∂zk∂L=−i=1∑Vpi∂zk∂logqi
根据链式法则,∂zk∂logqi=qi1∂zk∂qi。对于 Softmax 函数 qi,其对 zk 的导数为标准形式:
∂zk∂qi=T1qi(δik−qk)
(其中 δik 是 Kronecker delta,当 i=k 时为 1,否则为 0)
代入上式:
∂zk∂logqi=qi1⋅T1qi(δik−qk)=T1(δik−qk)
将其代回 Loss 的导数公式:
∂zk∂L=−i=1∑Vpi[T1(δik−qk)]=−T1(i=1∑Vpiδik−i=1∑Vpiqk)
因为 ∑i=1Vpiδik=pk,且概率之和 ∑i=1Vpi=1,所以:
∂zk∂L=−T1(pk−qk)=T1(qk−pk)
(注意:此时梯度中已经出现了一个 T1 的衰减因子)
3. 第二步:对概率差值 (qk−pk) 进行泰勒展开
当 T 较大时,Tui 和 Tzi 都趋近于 0。我们可以利用一阶泰勒展开 ex≈1+x。
展开 Student 的 probability qk:
qk=∑j=1Vexp(zj/T)exp(zk/T)≈∑j=1V(1+zj/T)1+zk/T=V+T1∑j=1Vzj1+zk/T
关键假设:在 Softmax 中,对所有 logits 减去同一个常数不会改变概率分布。因此,我们可以假设 logits 已经过均值中心化 (mean-centered),即 ∑j=1Vzj=0 且 ∑j=1Vuj=0。
在此假设下,分母简化为 V,因此:
qk≈V1+zk/T=V1+V⋅Tzk
同理,Teacher 的 probability pk 也可以展开为:
pk≈V1+V⋅Tuk
计算两者的差值:
qk−pk≈(V1+V⋅Tzk)−(V1+V⋅Tuk)=V⋅Tzk−uk
(注意:这里又出现了一个 T1 的衰减因子)
4. 第三步:代回梯度公式,得出 T2 的必要性
将第二步的结果代入第一步的梯度公式中:
∂zk∂L≈T1(V⋅Tzk−uk)=V⋅T21(zk−uk)
结论分析:从最终公式可以看出,当 T 较大时,原始交叉熵 Loss 产生的梯度与 T21 成正比。如果 T=5,梯度会缩小 25 倍;如果 T=10,梯度会缩小 100 倍。这会导致 Student 模型的参数更新停滞。
补偿方案:如果我们在 Loss 前面乘以 T2,定义缩放后的 Loss 为 Lscaled=T2L,那么新的梯度为:
∂zk∂Lscaled=T2∂zk∂L≈T2⋅V⋅T21(zk−uk)=V1(zk−uk)
乘以 T2 后,梯度恰好近似为 V1(zk−uk)。这等价于最小化 Student logits zk 和 Teacher logits uk 之间的均方误差 (MSE) 的梯度。这保证了无论 T 取何值,梯度的量级都能与 T=1 时保持一致,从而维持稳定的学习率。
5. 工程实践避坑指南 (PyTorch)
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')
loss_kd = (T ** 2) * kl_loss
7.5 基于特征的蒸馏 (Feature-based Distillation)
除了上述基于输出 Logits 的蒸馏(Response-level KD),基于特征的蒸馏是另一种极其重要的 Off-policy 方法。它要求 Student 的中间层表示(Representation)去拟合 Teacher 的中间层表示。
由于 Teacher 的特征是预先计算好并固定的,这同样属于 Off-policy 范式。常见的特征蒸馏包括:
- 隐藏状态对齐 (Hidden States Alignment):强制 Student 的某一层(或每一层)的隐藏状态经过线性投影后,逼近 Teacher 对应层的隐藏状态。
Lhidden=N1i=1∑NWhS(i)−hT(i)22
(其中 W 是用于对齐维度的投影矩阵,N 为序列长度)
- 注意力矩阵对齐 (Attention Maps Alignment):不仅对齐最终的向量表示,还要求 Student 学习 Teacher 的注意力分布,即让 Student 知道”在生成当前词时,应该关注输入序列的哪些位置”。
Lattn=L⋅N1l=1∑Li=1∑NAS(l,i)−AT(l,i)22
(其中 L 为层数,A 为注意力权重矩阵)
为什么需要基于特征的蒸馏?
- 提供更细粒度的监督:Logits 只是模型思考的”最终结果”,而中间特征包含了模型推理的”过程”。拟合特征能让 Student 学到更本质的语言表征和推理逻辑(如 TinyBERT, DistilBERT 的核心思想)。
- 缓解 Logits 蒸馏的局限性:当 Teacher 和 Student 架构差异过大,或任务极其复杂时,仅靠最终的 Soft Label 可能不足以指导 Student,中间层特征提供了更密集的梯度信号。
7.6 总结:Off-policy KD 的核心优势
将 Logits 蒸馏与特征蒸馏结合,构成了现代 LLM 压缩的标准 Off-policy 范式:
- 稳定性高:避免了 On-policy (如 RL) 中高方差的采样问题。
- 信息密度大:Full-vocab 的 Forward KL + 中间层特征对齐,最大化了单个训练样本的信息利用率。
第八章:Exposure Bias 问题
8.1 核心矛盾
训练 prefix = 数据/teacher 的 (x,y<tdata)
推理 prefix = student 自己的 (x,y<tstudent)
一旦 student 前面走偏,后续进入训练未见的区域 → 错误连锁放大。
8.2 生活比喻
学开车:教练总带你走标准路线(SFT训练)。自己开车时第3个路口就偏离了,从第4个路口开始从未练习过 → 开始迷路。
8.3 OPD 解决 Exposure Bias
让训练也发生在 student 自己生成的 prefix 上 → 训练和推理的前缀分布一致。
第九章:OPD 的核心思想
9.1 OPD 的一句话定义
On-Policy Distillation(OPD,同策略蒸馏):
学生模型先用自己的当前策略生成回答,再让教师模型在这些学生自己生成的轨迹上提供监督信号,学生据此更新。
形式上:
- Student 生成轨迹:y∼πθ(⋅∣x)
- 第 t 步的 prefix 是:ct=(x,y<t)(student 自己的前缀)
- Teacher 在这个 prefix 上提供信号
- Student 据此更新参数 θ
9.2 OPD 与 SFT/KD 的根本区别
| SFT / Off-policy KD | OPD |
|---|
| Prefix 来源 | 数据/teacher(固定的) | Student 自己生成(动态的) |
| 训练 prefix | ct=(x,y<tdata) | ct=(x,y<tstudent) |
| 是否有 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)]
写成期望:
Jx(θ)=Ey∼πθ(⋅∣x)[logπθ(y∣x)−logq(y∣x)]
为什么是 reverse KL?因为期望在 student 分布上取——我们只能看到 student 会生成哪些轨迹,在这些轨迹上评估差距。
10.2 梯度推导——逐步详解
这一步是为了介绍 OPD 和 RL 的关系。
第一步:乘积求导法则
J(θ) 是 πθ(y) 和 [logπθ(y)−logq(y)] 的乘积之和。两者都依赖 θ(logq(y) 不依赖 θ),所以用乘积求导法则:
∇θ[A⋅B]=(∇A)⋅B+A⋅(∇B)
应用到求和内部:
∇θ[πθ(y)⋅(logπθ(y)−logq(y))]
=∇πθ(y)⋅(logπθ(y)−logq(y))+πθ(y)⋅∇(logπθ(y))
注意:∇logq(y)=0(teacher 不依赖 θ)
对所有 y 求和:
∇J=y∑∇πθ(y)[logπθ(y)−logq(y)]+y∑πθ(y)∇logπθ(y)
第二步:Log-derivative trick 处理第一项
∇πθ(y)=πθ(y)⋅∇logπθ(y)
推导:πθ(y)=elogπθ(y) → ∇elogπθ(y)=elogπθ(y)⋅∇logπθ(y)=πθ(y)⋅∇logπθ(y)
代入第一项:
y∑∇πθ(y)[logπθ(y)−logq(y)]=y∑πθ(y)⋅∇logπθ(y)⋅[logπθ(y)−logq(y)]
写成期望(∑yπθ(y)⋅(⋯)=Ey∼πθ[⋯]):
=Ey∼πθ[(logπθ(y)−logq(y))⋅∇logπθ(y)]
第三步:第二项为什么等于 0?
Ey∼πθ[∇logπθ(y)]=y∑πθ(y)∇logπθ(y)
用 log-derivative trick 反向:πθ(y)∇logπθ(y)=∇πθ(y)
y∑πθ(y)∇logπθ(y)=y∑∇πθ(y)=∇(y∑πθ(y))=∇1=0
关键:概率分布归一性 ∑yπθ(y)=1 → ∇1=0。每个时间步的 softmax 保证概率和为 1,联合分布也归一化,所以总和对参数的梯度恒为 0。
第四步:合并
∇J=Ey∼πθ[(logπθ(y)−logq(y))⋅∇logπθ(y)]+0
∇J=Ey∼πθ[(logπθ(y)−logq(y))⋅∇logπθ(y)]
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)] ← 写成期望
│ = Σ_y π_θ(y)·∇logπ_θ(y)
│ = Σ_y ∇π_θ(y) ← 再次用 log-derivative trick 反向
│ = ∇(Σ_y π_θ(y)) ← 求和与梯度交换
∇J = E_{y~π_θ}[(logπ_θ(y) - logq(y))·∇logπ_θ(y)] ← 最终 policy-gradient 形式
10.4 结果的含义
∇J=Ey∼πθreward/权重(logπθ(y)−logq(y))⋅score function∇logπθ(y)
- Score function:参数怎么调才能让模型更多生成这条轨迹
- 权重(log-ratio):这条轨迹好不好——student 过度自信 → 减少;teacher 更支持 → 增加
这就是经典的 REINFORCE / policy-gradient 结构:score function × reward。因此就可以将 OPD 和 RL 进行比较。
第十一章:Return-to-Go 与方差分析
11.1 Log-ratio 的定义
rt=logq(yt∣ct)πθ(yt∣ct)=logπθ(yt∣ct)−logq(yt∣ct)
gt=∇θlogπθ(yt∣ct)
log-ratio 就是 student 和 teacher 对同一个 token 的 logprob 之差。
| rt 的值 | 含义 |
|---|
| rt>0 | student 比 teacher 更支持这个 token |
| rt=0 | 完全一致 |
| rt<0 | teacher 比 student 更支持 |
整条轨迹的 log-ratio = 所有 token 的 log-ratio 之和:
logq(y∣x)πθ(y∣x)=t′=1∑Trt′
11.2 直接的 Sequence-level Estimator
在 on-policy 蒸馏(OPD)框架下,学生策略 πθ 与环境交互并收集轨迹,同时以教师策略提供的信号作为奖励。直接的序列级梯度估计器定义为:
g^seq=t=1∑T(t′=1∑Trt′)gt
该估计器用整条轨迹的总回报 R=∑t′=1Trt′ 同时加权所有时间步的梯度,以此提升高回报轨迹的出现概率。
问题:因果性违背与方差增大
从公式可以明显看出,第 t 步的梯度 gt 被所有时间步的奖励 rt′ 加权,其中包含了 t′<t 的奖励。也就是说,动作发生之前已经获得的奖励被用来评价这个动作的好坏。这违背了时间因果性:当前动作不可能影响过去,过去奖励的大小不应成为强化或抑制当前动作的依据。
将总回报拆分为”过去奖励”与”未来奖励”:
R=Rpast(t)+Rfuture(t),Rpast(t)=t′=1∑t−1rt′,Rfuture(t)=t′=t∑Trt′
对第一项求期望。给定 st 时,过去奖励 Rpast(t) 与当前动作 at 无关,因此:
Eat∼πθ[Rpast(t)gt]=Rpast(t)Eat[∇θlogπθ(at∣st)]=0
这说明包含过去奖励的项在期望下为零,不贡献梯度信号,只凭空增加方差。尽管 g^seq 仍是真实策略梯度的无偏估计,但其方差远大于只使用未来奖励的版本,会导致训练收敛缓慢且不稳定。
正确的因果估计器:Reward-to-Go
应将估计器修正为只让每个动作对其之后发生的奖励负责,即使用”奖励到结束”(reward-to-go):
g^causal=t=1∑T(t′=t∑Trt′)gt
这正是经典 REINFORCE 算法中不带折扣的形式。它保留了无偏性,并因移除了与动作无关的过去奖励而显著降低了方差。
进一步的方差缩减
即使采用 reward-to-go,面对长序列时方差仍可能很大。常用改进手段包括:
- 引入基线(baseline):将累积奖励减去一个仅依赖状态的值函数 V(st),得到 g^=∑t(∑t′=tTrt′−b(st))gt,在不改变期望的前提下进一步降低方差。
- 使用折扣因子:对远期奖励乘以 γt′−t,既符合 MDP 的折扣回报定义,也能限制远期随机性的影响。
- 优势函数:当引入 Critic 时,直接用优势 At=Q(st,at)−V(st) 替代 reward-to-go,这是 Actor-Critic 方法的基础。
11.4 方差分析
Token-level:O(T2)
g^tok=t=1∑Trt⋅gt
T 项求和,每项方差 O(1),协方差最坏 O(T2) → 总方差 O(T2)。
Sequence-level:O(T4)
g^seq=t=1∑TRt⋅gt
权重 Rt 本身是 O(T) 个 rt′ 的求和 → ∣Rt∣≤O(T)⋅R → 每项方差 O(T2) → 总方差 O(T4)。
实际影响
| T=50 | T=4000 |
|---|
| token-level O(T2) | ~2,500 | ~16,000,000 |
| sequence-level O(T4) | ~6,250,000 | ~256,000,000,000,000 |
这就是为什么长 reasoning 任务中 sequence-level OPD “容易炸”。
11.5 折扣形式:γ 连接两者
g^γ=t=1∑T(t′=t∑Tγt′−trt′)gt
| γ | 含义 | 方差 |
|---|
| γ=0 | token-level OPD | O(T2),有偏但稳定 |
| γ=1 | sequence-level OPD | O(T4),无偏但高方差 |
| 0<γ<1 | 折中 | 介于两者 |
当 γ<1 时权重 Rtγ 的上界变为常数 1−γR,总方差回到 O(T2)。
第十二章:Token-level OPD
12.1 定义
只保留当前 token 的 rt:
g^tok=t=1∑Trt⋅gt
相对 sequence-level 有偏(丢掉了未来项 ∑t′=t+1Trt′),但方差从 O(T4) 降到 O(T2)。
12.2 Advantage 形式
At=logq(yt∣ct)−logπθ(yt∣ct)=−rt
- At>0:teacher 更支持 → 提高概率
- At<0:student 更自信 → 降低概率
Policy-gradient 更新:
∇J≈At⋅∇logπθ(yt∣ct)
这就是 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-policy | On-policy |
|---|
| Prefix | 固定,预先就有 | Student 实时生成 |
| Teacher forward | 和 student 一起跑 | 需要额外在新 prefix 上跑 |
| Full logits | 天然可用 | 需要额外获取,代价高 |
13.2 Sampled-token OPD
Student rollout y∼πθ,只比较 sampled token:
At=logq(yt∣ct)−logπθ(yt∣ct)
- Teacher 只需返回 1 个 logprob
- 信息量最低,方差高、信号稀疏
13.3 Top-K OPD
Teacher 在 prefix ct 上返回词表中概率最高的 K 个 token:
St=TopKq(ct)⊂V
Top-K 是从词表中选,不是从 prefix 中选。Prefix 是已生成的历史(固定输入),Top-K 选择的是”下一个 token”的候选集。
在支持集内重归一化:
q^(v∣ct)=∑u∈Stq(u∣ct)q(v∣ct),π^(v∣ct)=∑u∈Stπθ(u∣ct)πθ(v∣ct)
为什么重归一化? 只看 K 个 token 时概率总和不到 1,重归一化让它们在子集上重新变成合法分布(总和=1),才能算有意义的 KL。
局部 reverse KL:
LtopK(ct)=v∈St∑π^(v∣ct)logq^(v∣ct)π^(v∣ct)
- 信息量中等,比 sampled-token 稳定
- 有截断偏差:忽略 teacher top-k 外的 token
13.4 Full-vocab OPD
每个 prefix 上比较整个词表分布:
Lt=v∈V∑PS(v∣ct)logPT(v∣ct)PS(v∣ct)=DKL(PS∥PT)
- 信息最完整,精确 KL,无偏无截断
- 显存和计算代价最大
13.5 三者对比
| 粒度 | Teacher 返回 | 信息量 | 成本 | 稳定性 |
|---|
| sampled-token | 1 个 logprob | 低 | 低 | 差 |
| top-k | K 个 logprob + K 个 id | 中 | 中 | 较好 |
| full-vocab | 全部 logits | 高 | 高 | 最好但昂贵 |
第十四章:k1/k2/k3 估计器
14.1 为什么需要估计器?
只有 sampled token 的 logprob 时,无法直接计算精确 KL。需要单样本估计器。
14.2 Reverse KL 的期望形式
DKL(PS∥PT)=Ey∼PS[logPT(y∣ct)PS(y∣ct)]
14.3 k1(无偏,可负)
k1=logPT(yt∣ct)PS(yt∣ct),yt∼PS
无偏(E[k1]=DKL(PS∥PT)),单样本可正可负,方差高。
14.4 k2(有偏,非负)
k2=21(logPT(yt∣ct)PS(yt∣ct))2
始终非负,方差低。是 KL 的局部二阶近似(当 PS 和 PT 接近时偏差小)。
14.5 k3(无偏,非负)
k3=PS(yt∣ct)PT(yt∣ct)−logPS(yt∣ct)PT(yt∣ct)−1
设 r=PT/PS,则 k3=r−logr−1≥0。
无偏:利用 Ey∼PS[PT(y)/PS(y)]=1(因为 ∑yPS(y)⋅PT(y)/PS(y)=∑yPT(y)=1),可证 E[k3]=DKL(PS∥PT)。
14.6 对比
| 估计器 | 表达式 | 无偏? | 非负? | 方差 |
|---|
| k1 | logPTPS | 是 | 否 | 高 |
| k2 | 21(logPTPS)2 | 否 | 是 | 低 |
| k3 | PSPT−logPSPT−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()