Skip to content
alarik.me
Go back

OPD 深度解析(一):概率、熵与 KL 散度

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

本文档面向零基础读者,从信息论基础(熵)出发,逐步推导到 SFT、离线知识蒸馏、强化学习策略梯度,最终到 OPD 的各种变体。读完此文档,你将完整理解 OPD 的数学原理和工程实践。

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

  1. 数学基础、信息论基础和 KL 散度(本篇)
  2. 大模型微调(SFT)、强化学习策略梯度和知识蒸馏
  3. OPD 的实现和工程内容

Table of contents

Open Table of contents

第一章:概率与信息论基础

1.1 什么是概率分布?

在日常生活中,我们常说”明天有 70% 的概率会下雨”。这就是概率——用一个数字描述某件事发生的可能性。

在数学中,概率分布就是把所有可能结果的概率列出来,并且满足:

所有可能结果P(结果)=1\sum_{\text{所有可能结果}} P(\text{结果}) = 1

意思是:所有可能性的概率加起来必须等于 100%。这是概率分布最基本的要求——归一性

例子:掷一个六面骰子,每个面的概率是 16\frac{1}{6}

结果概率
11/6
21/6
31/6
41/6
51/6
61/6

验证归一性: 16×6=1\frac{1}{6} \times 6 = 1

1.2 什么是期望?

期望就是”平均值”,但不是普通的平均,而是按概率加权的平均

EYp[f(Y)]=yp(y)f(y)\mathbb{E}_{Y \sim p}[f(Y)] = \sum_y p(y) \cdot f(y)

解读:每个可能值 yy 的函数值 f(y)f(y),乘以这个值出现的概率 p(y)p(y),然后求和。

Note

为什么写成 E\mathbb{E} 符号? yp(y)f(y)\sum_y p(y) \cdot f(y)Eyp[f(y)]\mathbb{E}_{y \sim p}[f(y)] 是完全等价的东西,后者只是前者的简写——这就是”期望的定义”,没有任何推导,就是定义本身。

1.3 蒙特卡洛采样估计

期望可以用采样近似:

Eyp[f(y)]1Ni=1Nf(y(i)),y(i)p\mathbb{E}_{y \sim p}[f(y)] \approx \frac{1}{N}\sum_{i=1}^{N} f(y^{(i)}), \quad y^{(i)} \sim p

在 OPD 的语境里,yy 是语言模型生成的一条完整回答序列,可能的序列数量是天文数字,不可能真的遍历求和。所以我们总是用采样估计。

第二章:熵与交叉熵

熵和交叉熵是信息论的基础概念,也是机器学习中几乎所有 loss 函数的底层逻辑。理解它们,才能理解 SFT loss 为什么是交叉熵、为什么最小化交叉熵等价于最小化 KL 散度。

2.1 熵——衡量”不确定性”

2.1.1 从生活直觉出发

想象两个场景:

场景B比场景A更”不确定”。熵就是量化这种不确定性的数学工具

2.1.2 定义与公式

H(P)=vP(v)logP(v)H(P) = -\sum_{v} P(v) \log P(v)

逐步拆解每个符号

符号含义
P(v)P(v)事件 vv 发生的概率
logP(v)\log P(v)事件 vv 的”信息量”(概率越低,信息量越大——“罕见事件更有信息价值”)
P(v)logP(v)P(v) \log P(v)事件 vv 的信息量 × 它发生的概率 → 期望信息量
v-\sum_v对所有事件求和,取负号(因为 logP(v)\log P(v) 是负数,取负后变正)
Note

为什么 logP(v)\log P(v) 是”信息量”? 直觉:如果你说”明天太阳会升起”(概率≈1),这句话几乎没有信息价值——大家都知道。但如果你说”明天会地震”(概率≈0.001),这句话信息量巨大——它传达了极其罕见的事件。数学上,信息量定义为 logP(v)-\log P(v):概率越低,信息量越大。

2.1.3 为什么公式是这个形式?推导思路

目标:找一个函数 H(P)H(P) 衡量分布 PP 的”不确定性”,满足以下直觉:

从直觉1出发:如果 P(v)=1P(v)=1,那么 logP(v)=log1=0-\log P(v) = -\log 1 = 0,乘以 P(v)=1P(v)=1 还是 0。所有其他 P(u)=0P(u)=0 的项也是 0。所以 H=0H=0。✓

从直觉2出发:如果所有 P(v)=1NP(v) = \frac{1}{N}NN 个等概率事件),则:

H=v=1N1Nlog1N=1NNlog1N=logNH = -\sum_{v=1}^{N} \frac{1}{N} \log \frac{1}{N} = -\frac{1}{N} \cdot N \cdot \log \frac{1}{N} = \log N

NN 越大(越多等概率事件),HH 越大。✓

从直觉3出发:两个独立事件 XXYY,联合分布 P(x,y)=P(x)P(y)P(x,y) = P(x)P(y)

H(X,Y)=x,yP(x)P(y)logP(x)P(y)=x,yP(x)P(y)[logP(x)+logP(y)]H(X,Y) = -\sum_{x,y} P(x)P(y) \log P(x)P(y) = -\sum_{x,y} P(x)P(y)[\log P(x) + \log P(y)] =xP(x)logP(x)yP(y)+yP(y)logP(y)xP(x)=H(X)+H(Y)= -\sum_x P(x)\log P(x) \cdot \sum_y P(y) + -\sum_y P(y)\log P(y) \cdot \sum_x P(x) = H(X) + H(Y)

✓ 熵满足可加性。

2.1.4 计算示例(逐步)

示例1:完全确定分布 P=(1,0,0)P = (1, 0, 0)

H=1log10log00log0=0H = -1 \cdot \log 1 - 0 \cdot \log 0 - 0 \cdot \log 0 = 0
Note

注意:0log00 \cdot \log 0 在数学上定义为 0(因为 limp0plogp=0\lim_{p \to 0} p \log p = 0

示例2:两个等概率 P=(0.5,0.5)P = (0.5, 0.5)

H=0.5log0.50.5log0.5=log0.5=log20.693H = -0.5 \cdot \log 0.5 - 0.5 \cdot \log 0.5 = -\log 0.5 = \log 2 \approx 0.693

示例3:六面等概率骰子 P=(16,16,16,16,16,16)P = (\frac{1}{6}, \frac{1}{6}, \frac{1}{6}, \frac{1}{6}, \frac{1}{6}, \frac{1}{6})

H=log61.79H = \log 6 \approx 1.79

示例4:偏斜分布 P=(0.9,0.1)P = (0.9, 0.1)

H=0.9log0.90.1log0.10.9×0.105+0.1×2.303=0.325H = -0.9 \cdot \log 0.9 - 0.1 \cdot \log 0.1 \approx 0.9 \times 0.105 + 0.1 \times 2.303 = 0.325

比等概率分布的熵(0.693)低——因为偏斜分布更”确定”。

对比总结

分布不确定性
(1,0,0)(1, 0, 0)0完全确定
(0.9,0.1)(0.9, 0.1)0.325比较确定
(0.5,0.5)(0.5, 0.5)0.693中等不确定
(16,...,16)(\frac{1}{6}, ..., \frac{1}{6})1.79很不确定

2.1.5 熵与语言模型的关系

语言模型在每个时间步输出一个 token 分布 πθ(vct)\pi_\theta(v|c_t)。这个分布的熵告诉我们:

在 OPD 中,teacher 和 student 的 entropy gap(熵的差异)是衡量两者思维模式是否一致的关键指标。

2.2 交叉熵——衡量”用错误的编码描述真实分布”

2.2.1 从生活直觉出发

想象你在给外国朋友描述中国菜:

用不合适的编码去描述,需要更多的信息量。交叉熵就是”用分布 QQ 去编码真实分布 PP 所需的期望信息量”

2.2.2 定义与公式

H(P,Q)=vP(v)logQ(v)H(P, Q) = -\sum_{v} P(v) \log Q(v)

逐步拆解

符号含义
P(v)P(v)真实分布中事件 vv 的概率(权重)
logQ(v)\log Q(v)用分布 QQ 编码时,事件 vv 的信息量
P(v)logQ(v)P(v) \log Q(v)真实事件 vv 的信息量 × 它真实发生的概率
v-\sum_v对所有事件求期望信息量
Note

关键区别:熵 H(P)H(P) 中权重和编码都是 PP(自己编码自己),交叉熵 H(P,Q)H(P,Q) 中权重是 PP 但编码是 QQ(用别人的编码描述自己)。

2.2.3 为什么交叉熵 ≥ 熵?

直觉:用不合适的编码描述,总是比用合适的编码描述需要更多信息量。最优编码就是用自己的分布编码自己——这就是熵。

数学证明

H(P,Q)H(P)=vP(v)logQ(v)+vP(v)logP(v)=vP(v)logP(v)Q(v)=DKL(PQ)0H(P, Q) - H(P) = -\sum_v P(v)\log Q(v) + \sum_v P(v)\log P(v) = \sum_v P(v)\log\frac{P(v)}{Q(v)} = D_{\mathrm{KL}}(P \| Q) \ge 0

所以 H(P,Q)H(P)H(P, Q) \ge H(P),等号当 Q=PQ = P

2.2.4 计算示例(逐步)

真实分布 P=(0.5,0.5)P = (0.5, 0.5)(下雨/不下雨各50%)

情况1:预测完美 Q=(0.5,0.5)Q = (0.5, 0.5)

H(P,Q)=0.5log0.50.5log0.5=log20.693H(P, Q) = -0.5\log 0.5 - 0.5\log 0.5 = \log 2 \approx 0.693

交叉熵 = 熵(预测完美,没有额外浪费)

情况2:预测偏斜 Q=(0.9,0.1)Q = (0.9, 0.1)(你觉得90%下雨)

H(P,Q)=0.5log0.90.5log0.10.5×0.105+0.5×2.303=1.204H(P, Q) = -0.5\log 0.9 - 0.5\log 0.1 \approx 0.5 \times 0.105 + 0.5 \times 2.303 = 1.204

交叉熵 > 熵(1.204 > 0.693),额外浪费 = KL散度 = 1.204 - 0.693 = 0.511

情况3:预测极端错误 Q=(0.99,0.01)Q = (0.99, 0.01)

H(P,Q)=0.5log0.990.5log0.010.5×0.01+0.5×4.605=2.308H(P, Q) = -0.5\log 0.99 - 0.5\log 0.01 \approx 0.5 \times 0.01 + 0.5 \times 4.605 = 2.308

交叉熵更大(2.308),额外浪费更多

对比总结

预测 QQ交叉熵 H(P,Q)H(P,Q)H(P)H(P)KL散度 DKLD_{\mathrm{KL}}额外浪费
(0.5,0.5)(0.5, 0.5)(完美)0.6930.6930
(0.9,0.1)(0.9, 0.1)(偏斜)1.2040.6930.511中等
(0.99,0.01)(0.99, 0.01)(极端)2.3080.6931.615

2.2.5 交叉熵与机器学习的关系

为什么机器学习几乎都用交叉熵作为 loss?

因为在分类/生成任务中:

也就是说:训练模型 = 让模型分布靠近真实分布。交叉熵 loss 就是这个目标的自然度量。

Note

SFT loss 就是交叉熵(第五章会详细推导):真实分布是 one-hot(标准答案),模型分布是 softmax 输出,最小化交叉熵 = 让模型给标准答案高概率。

2.3 交叉熵 = 熵 + KL散度(详细推导)

这是信息论中最重要的等式之一,它把三个概念串联起来。

推导

H(P,Q)=vP(v)logQ(v)H(P, Q) = -\sum_v P(v) \log Q(v)

第一步:拆开 logQ(v)\log Q(v)

logQ(v)=logP(v)Q(v)P(v)=logP(v)+logQ(v)P(v)\log Q(v) = \log P(v) \cdot \frac{Q(v)}{P(v)} = \log P(v) + \log \frac{Q(v)}{P(v)}

这里用了 log(ab)=loga+logb\log(ab) = \log a + \log b

代入:

H(P,Q)=vP(v)[logP(v)+logQ(v)P(v)]H(P, Q) = -\sum_v P(v) [\log P(v) + \log \frac{Q(v)}{P(v)}]

第二步:分成两项

H(P,Q)=vP(v)logP(v)vP(v)logQ(v)P(v)H(P, Q) = -\sum_v P(v) \log P(v) - \sum_v P(v) \log \frac{Q(v)}{P(v)}

第三步:识别每一项

第一项:

vP(v)logP(v)=H(P)-\sum_v P(v) \log P(v) = H(P)

这就是熵的定义。

第二项:

vP(v)logQ(v)P(v)=vP(v)logP(v)Q(v)=DKL(PQ)-\sum_v P(v) \log \frac{Q(v)}{P(v)} = \sum_v P(v) \log \frac{P(v)}{Q(v)} = D_{\mathrm{KL}}(P \| Q)

这就是 KL 散度的定义。

注意符号:logQP=logPQ-\log\frac{Q}{P} = \log\frac{P}{Q},负号和分数翻转抵消了。

结果

H(P,Q)=H(P)+DKL(PQ)H(P, Q) = H(P) + D_{\mathrm{KL}}(P \| Q)

含义图解

交叉熵 H(P,Q) = "用Q编码P所需的总信息量"
├── 熵 H(P) = "用P编码P所需的最少信息量"(不可减少的底线)
└── KL散度 D_KL(P‖Q) = "因为编码不合适而额外浪费的信息量"(可以优化消除)
Important

为什么这个等式重要?

  • 解释了为什么最小化交叉熵等价于最小化 KLH(P)H(P) 对模型参数是常数,优化交叉熵唯一能改变的部分就是 KL 散度
  • 解释了交叉熵的下界H(P,Q)H(P)H(P,Q) \ge H(P),下界就是熵——当模型完美匹配真实分布时达到
  • 连接了信息论和机器学习:训练 loss(交叉熵)= 信息论概念(熵 + KL),让我们可以用信息论的语言理解训练过程

第三章:KL 散度

KL散度是整个 OPD 理论的基石。理解 KL 的两个方向(forward/reverse)和它的折中形式(JSD),才能理解为什么 OPD 选 reverse KL,为什么 DeepSeek V4 不用 sampled-token。

3.1 定义

KL 散度衡量两个分布的”差异大小”:

DKL(PQ)=vP(v)logP(v)Q(v)=vP(v)[logP(v)logQ(v)]D_{\mathrm{KL}}(P \| Q) = \sum_v P(v) \log \frac{P(v)}{Q(v)} = \sum_v P(v)[\log P(v) - \log Q(v)]

为什么这样定义? 直觉是:用分布 QQ 去编码/描述分布 PP 时,比用 PP 自己编码多浪费了多少信息量。logP(v)Q(v)\log\frac{P(v)}{Q(v)} 就是”真实概率 vs 你的预测概率”的比值,P(v)P(v) 加权求和就是平均浪费。

性质

Important

不对称是理解 forward/reverse KL 的关键——两个方向的行为完全不同。

3.2 Forward KL 详细公式与推导

DKL(PTPS)=vVPT(vct)logPT(vct)PS(vct)D_{\mathrm{KL}}(P_T \| P_S) = \sum_{v \in \mathcal{V}} P_T(v|c_t) \log \frac{P_T(v|c_t)}{P_S(v|c_t)}

展开写法(更直观):

DKL(PTPS)=vPT(v)[logPT(v)logPS(v)]D_{\mathrm{KL}}(P_T \| P_S) = \sum_v P_T(v) \cdot \left[\log P_T(v) - \log P_S(v)\right]

逐项解读

为什么是 mode-covering? 逐步推导:

PT(v)>0P_T(v) > 0PS(v)0P_S(v) \to 0 时:

PT(v)logPT(v)PS(v)+P_T(v) \log \frac{P_T(v)}{P_S(v)} \to +\infty

意思是:只要 teacher 认为某 token 有哪怕一点点概率,student 就不能给太低概率,否则 KL 会爆炸。Student 被迫给所有 teacher 支持的 token 都分一些概率——这就是”覆盖所有模式”。

代价:如果 student 容量有限(参数少),它只能在每个模式上分一点概率,结果是每个模式都学不好——“平均但不够好”。

PyTorch 实现

import torch
import torch.nn.functional as F
def forward_kl(teacher_logits, student_logits):
"""计算 D_KL(P_T || P_S),即 forward KL
Args:
teacher_logits: teacher输出的logits,shape [B, L, V]
student_logits: student输出的logits,shape [B, L, V]
Returns:
每个位置的forward KL,shape [B, L]
"""
teacher_log_probs = F.log_softmax(teacher_logits, dim=-1)
student_log_probs = F.log_softmax(student_logits, dim=-1)
teacher_probs = F.softmax(teacher_logits, dim=-1)
kl = teacher_probs * (teacher_log_probs - student_log_probs)
return kl.sum(dim=-1)
def forward_kl_builtin(teacher_logits, student_logits):
teacher_probs = F.softmax(teacher_logits, dim=-1)
student_log_probs = F.log_softmax(student_logits, dim=-1)
cross_entropy = F.cross_entropy(
student_logits.transpose(-1, -2),
teacher_probs.transpose(-1, -2), reduction='none'
)
teacher_entropy = -(teacher_probs * F.log_softmax(teacher_logits, dim=-1)).sum(dim=-1)
return cross_entropy - teacher_entropy

3.3 Reverse KL 详细公式与推导

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

展开写法:

DKL(PSPT)=vPS(v)[logPS(v)logPT(v)]D_{\mathrm{KL}}(P_S \| P_T) = \sum_v P_S(v) \cdot \left[\log P_S(v) - \log P_T(v)\right]

逐项解读

为什么是 mode-seeking? 逐步推导:

PS(v)>0P_S(v) > 0PT(v)0P_T(v) \to 0 时:

PS(v)logPS(v)PT(v)+P_S(v) \log \frac{P_S(v)}{P_T(v)} \to +\infty

意思是:student 把概率放在 teacher 不支持的 token 上,惩罚会爆炸。为了避免这个惩罚,student 倾向于只把概率分配给 teacher 概率最高的那几个 token——这就是”收缩到主要模式”。

代价:student 可能忽略 teacher 的低概率但仍然合理的模式,导致多样性坍塌。

PyTorch 实现

import torch
import torch.nn.functional as F
def reverse_kl(teacher_logits, student_logits):
"""计算 D_KL(P_S || P_T),即 reverse KL"""
teacher_log_probs = F.log_softmax(teacher_logits, dim=-1)
student_log_probs = F.log_softmax(student_logits, dim=-1)
student_probs = F.softmax(student_logits, dim=-1)
kl = student_probs * (student_log_probs - teacher_log_probs)
return kl.sum(dim=-1)

3.4 Forward KL vs Reverse KL 的直觉对比

想象 teacher 的分布有两个峰(两种合理答案),student 容量有限只能是一个峰:

KL 方向Student 学到的比喻数学原因
Forward KL覆盖两个峰 → 宽而低面面俱到但不够好PT(v)>0,PS(v)0P_T(v)>0, P_S(v)\to 0 时KL爆炸
Reverse KL收缩到一个峰 → 精准专注高效但单一PS(v)>0,PT(v)0P_S(v)>0, P_T(v)\to 0 时KL爆炸

用数值例子理解

假设词表只有3个词,teacher分布 PT=(0.45,0.35,0.20)P_T = (0.45, 0.35, 0.20),student三种可能的分布:

Student分布FKLRKL说明
PS=(0.45,0.35,0.20)P_S=(0.45, 0.35, 0.20)(完美匹配)00两个KL都是0
PS=(0.50,0.30,0.20)P_S=(0.50, 0.30, 0.20)(轻微偏移)差距不大
PS=(0.90,0.05,0.05)P_S=(0.90, 0.05, 0.05)(收缩到峰1)大(token2、3被严重低估)(student没把概率放在teacher不支持的地方)RKL偏好这种
PS=(0.33,0.33,0.34)P_S=(0.33, 0.33, 0.34)(均匀覆盖)(每个token都有概率)大(student在teacher低概率的token3上给了太多概率)FKL偏好这种
Important

在 LLM 蒸馏中,Reverse KL 更受欢迎——因为我们希望 student 学到 teacher 最可靠的推理路径(mode-seeking),而不是学到一堆”平均但不够好”的模式。

3.5 JSD(Jensen-Shannon Divergence)——FKL和RKL的折中

JSD 是 forward KL 和 reverse KL 的对称化折中版本

为什么要折中? 纯 forward KL 容易让 student 学得”平均”,纯 reverse KL 容易让 student 多样性坍塌。JSD 试图在两者之间找到平衡。

定义(广义 JSD,参数 β\beta):

DJSD(β)(PT,PS)=βDKL(PTM)+(1β)DKL(PSM)D_{\mathrm{JSD}(\beta)}(P_T, P_S) = \beta \cdot D_{\mathrm{KL}}(P_T \| M) + (1-\beta) \cdot D_{\mathrm{KL}}(P_S \| M)

其中混合分布:

M=βPT+(1β)PSM = \beta \cdot P_T + (1-\beta) \cdot P_S

逐步解读

为什么 β\beta 控制方向?

β\betaJSD 变成直觉
β=0\beta=0DKL(PSPT)D_{\mathrm{KL}}(P_S \| P_T) = Reverse KL只惩罚student在teacher不支持的地方放概率
β=1\beta=1DKL(PTPS)D_{\mathrm{KL}}(P_T \| P_S) = Forward KL只惩罚student低估teacher支持的token
β=0.5\beta=0.5对称 JSD两端等权惩罚

推导验证β=0\beta=0):

DJSD(0)=0DKL(PTPS)+1DKL(PSPT)=DKL(PSPT)D_{\mathrm{JSD}(0)} = 0 \cdot D_{\mathrm{KL}}(P_T \| P_S) + 1 \cdot D_{\mathrm{KL}}(P_S \| P_T) = D_{\mathrm{KL}}(P_S \| P_T)

β=0\beta=0M=PSM = P_S,所以 DKL(PTM)=DKL(PTPS)D_{\mathrm{KL}}(P_T \| M) = D_{\mathrm{KL}}(P_T \| P_S),但 00 系数让它消失了。

推导验证β=0.5\beta=0.5):

M=0.5PT+0.5PSM = 0.5 P_T + 0.5 P_S DJSD(0.5)=0.5DKL(PTM)+0.5DKL(PSM)D_{\mathrm{JSD}(0.5)} = 0.5 D_{\mathrm{KL}}(P_T \| M) + 0.5 D_{\mathrm{KL}}(P_S \| M)

JSD 比 KL 的优势

PyTorch 实现

import torch
import torch.nn.functional as F
def jsd_loss(teacher_logits, student_logits, beta=0.5):
"""计算广义 JSD
beta=0: 等价 reverse KL
beta=1: 等价 forward KL
beta=0.5: 对称JSD
"""
teacher_probs = F.softmax(teacher_logits, dim=-1)
student_probs = F.softmax(student_logits, dim=-1)
teacher_log_probs = F.log_softmax(teacher_logits, dim=-1)
student_log_probs = F.log_softmax(student_logits, dim=-1)
m_probs = beta * teacher_probs + (1 - beta) * student_probs
m_log_probs = torch.log(m_probs.clamp(min=1e-10))
kl_teacher_m = (teacher_probs * (teacher_log_probs - m_log_probs)).sum(dim=-1)
kl_student_m = (student_probs * (student_log_probs - m_log_probs)).sum(dim=-1)
jsd = beta * kl_teacher_m + (1 - beta) * kl_student_m
return jsd

Share this post:

Previous Post
OPD 深度解析(三):知识蒸馏与 OPD 梯度推导
Next Post
OPD 深度解析(二):SFT 与策略梯度推导