跳至内容
返回

KL散度:信息论视角下概率分布对齐工具,原理、辨析与论文实战

发布于:

0 前言

KL 散度(Kullback-Leibler Divergence,相对熵)是深度学习、知识蒸馏、生成模型、时序建模中高频出现的基础工具。很多同学会把它和余弦相似度混淆:二者都用来衡量“两个向量的相似度”,但适用场景、数学含义完全不同。

本文从基础定义、核心特性出发,对比 KL 与余弦相似度的差异,结合两篇论文(CDSD、TimeDistill)讲解工程中如何落地使用,最后补充实践要点与参考资料。

1 什么是 KL 散度

1.1 数学定义

离散概率分布:

DKL(PQ)=xP(x)logP(x)Q(x)D_{KL}(P \parallel Q)=\sum_{x} P(x)\log\frac{P(x)}{Q(x)}

连续概率分布:

DKL(PQ)=p(x)logp(x)q(x)dxD_{KL}(P \parallel Q)=\int p(x)\log\frac{p(x)}{q(x)}dx

  • PP目标真实分布QQ模型近似分布
  • 物理含义:用近似分布 QQ 去描述真实分布 PP,带来的平均信息损失(额外编码代价)

关键特性

  1. 非负性DKL(PQ)0D_{KL}(P\parallel Q)\ge0;当且仅当 P=QP=Q 时等于 0。
  2. 不对称性DKL(PQ)DKL(QP)D_{KL}(P\parallel Q)\neq D_{KL}(Q\parallel P)它不是数学严格意义的距离,只能叫“散度/差异度量”。方向非常关键,不能随意调换 P、Q。
  3. 行为差异
    • DKL(PQ)D_{KL}(P\parallel Q)(正向 KL,zero-avoiding):强迫模型 Q 覆盖 P 所有非零区域,不能漏真实存在的可能性;
    • DKL(QP)D_{KL}(Q\parallel P)(反向 KL,mean-seeking):模型倾向于拟合分布主峰,忽略尾部小概率区域。

通俗比喻:P 是真实世界,Q 是手绘地图;KL 散度就是拿着这份不完美地图导航,付出的平均绕路成本。地图越接近真实世界,KL 值越小,绕路越少。

1.2 和交叉熵的关系

DKL(PQ)=H(P,Q)H(P)D_{KL}(P\parallel Q)=H(P,Q)-H(P)

H(P)H(P) 是真实分布 P 的信息熵,是常数。最小化 KL 等价于最小化交叉熵。这也是分类任务交叉熵损失的信息论来源。

直觉拆解

总不确定度(交叉熵 H(P,Q))
  = P 本身的不确定度(信息熵 H(P),数据自带的)
  + 模型 Q 额外引入的不确定度(KL,模型的错)

1.3 相对熵、交叉熵与余弦相似度的三角关系

三个概念正好构成“信息论家族 vs 几何家族”两大阵营,但又在对比学习/蒸馏中被 softmax 桥接

概念公式家族在问什么
相对熵(KL)DKL(PQ)=P(x)logP(x)Q(x)D_{KL}(P\parallel Q)=\sum P(x)\log\frac{P(x)}{Q(x)}信息论Q 描述 P 的额外信息损失
交叉熵H(P,Q)=P(x)logQ(x)H(P,Q)=-\sum P(x)\log Q(x)信息论用 Q 预测时的总不确定度
余弦相似度cos(A,B)=ABAB\cos(A,B)=\frac{A\cdot B}{\|A\|\|B\|}几何两个向量方向有多像

关系一:KL = 交叉熵 − 信息熵(差一个常数,训练时等价)

训练中 P(标签分布)固定 → H(P) 是常数 → 最小化 KL ⇔ 最小化交叉熵。所以分类任务用交叉熵损失,本质就是在最小化 KL——只是换了个名字

关系二:余弦 vs KL/交叉熵——完全不同的家族

余弦相似度KL / 交叉熵
家族几何(线性代数)信息论
输入任意向量(不用归一化)概率分布(softmax 后,和=1)
关注方向夹角概率质量分配
对称性对称非对称
本质相似度函数(越大越好)差异/损失度量(越小越好)

关系三:三者在同一个流程里怎么碰面(softmax 桥接)

对比学习(SimCLR/InfoNCE)流程:
特征向量 z_i, z_j
   ↓ 余弦相似度(几何:算相似度)
logit = cos(z_i, z_j) / τ
   ↓ softmax(把相似度变成概率分布)
P_ij
   ↓ 交叉熵/KL(信息论:算损失)
L = -log P_ij

三者各司其职:余弦负责“测量像不像”,softmax 负责“把像不像翻译成概率”,交叉熵/KL 负责“优化这个概率”——几何工具与信息论工具在 softmax 处握手。

蒸馏场景同理:概率对齐用 KL(软标签/频谱分布),特征对齐用余弦(向量特征)——这正是第 2 节“选择口诀”的底层原因。

2 KL 散度 vs 余弦相似度:什么时候选谁?

很多场景中两者都输入向量,但二者关注的对象完全不同:

指标核心关注点适用对象特点
KL 散度分布形状、概率质量分配概率分布向量(经过 softmax,总和=1)非对称;只对概率分布有物理意义;衡量信息损失
余弦相似度向量方向夹角任意特征向量对称;衡量向量指向是否相近,不要求归一化为概率;不关心概率质量分配

一句话选择口诀

  • 如果你手里是概率分布(分类输出、频谱经过 softmax、蒸馏软标签),关心“概率质量分配是否对齐”,优先用 KL 散度;
  • 如果你手里是隐层特征向量,只关心特征的指向语义,不需要概率归一化,优先用余弦相似度。

注意:KL 散度输入必须是合法概率分布(非负,求和为 1);直接对原始 logits 喂 KL 会得到无意义结果,需要先做 softmax。PyTorch 的 KLDivLoss 要求输入log(Q),目标直接传概率 P,接口容易踩坑。

3 论文实战:两篇真实使用 KL 的论文

注:经论文原文核实,本组 8 篇论文中实际使用 KL 散度的是 CDSD(分类概率对齐)与 TimeDistill(频域周期分布对齐)两篇。

3.1 CDSD:分类概率对齐(自蒸馏)

CDSD 任务场景中,KL 散度用来对齐自蒸馏中老师与学生的分类概率分布(老师和学生是同一个模型的不同分支):

  • PP:基于 DIR 分支(老师)计算出的分类软概率 y;
  • QQ:骨干中间层分类器(学生)输出的概率 y₁/y₂/y₃;
  • 目标:最小化 Lcc=KL(y,y1)+KL(y,y2)+KL(y,y3)L_{cc} = KL(y, y_1) + KL(y, y_2) + KL(y, y_3),让中间层学习 DIR 分支输出的类别概率权重,不仅仅学习 one-hot 硬标签;
  • 作用:把域不变分支学到的类别间“暗知识”蒸馏到骨干中间层,属于典型知识蒸馏用法。

关键点:输入是分类头 softmax 之后的概率向量,维度等于类别数。

3.2 TimeDistill:频域上对齐周期分布

TimeDistill 是面向长时序预测的跨架构蒸馏框架,把 Transformer 教师模型的周期知识蒸馏给轻量 MLP 学生模型。

  • 流程:
    1. 分别对教师、学生的预测序列做 FFT,提取频谱幅度;
    2. 对频谱幅度做带温度系数的 softmax,把频谱转换为“周期概率分布”QYtQ_{Yt}(教师)、QYsQ_{Ys}(学生);
    3. 使用 KL 散度作为周期蒸馏损失:LYperiod=KL(QYtQYs)\mathcal{L}_{Yperiod}=KL(Q_{Yt}\parallel Q_{Ys})
  • 本质:不再直接对齐原始时间序列,而是在频域对齐两个模型捕捉到的周期模式分布,强迫学生复现教师识别出来的高频/低频周期分量,提升长时序预测能力。

这是 KL 一个很巧妙的非分类场景:把 FFT 频谱归一化成概率分布,用 KL 衡量频谱的“概率质量”是否对齐。

3.3 两种真实用法对比:CDSD vs TimeDistill

对比一下能看出 KL 的两种典型打开方式:

维度CDSD(CVPR 2022)TimeDistill(KDD 2026)
对齐对象分类概率分布周期分布(频谱)
输入是什么分类头 softmax 输出(维度=类别数)FFT 频谱幅度 → 冷温度 softmax(维度=频率数)
公式Lcc=KL(y,y1)+KL(y,y2)+KL(y,y3)L_{cc} = KL(y, y_1) + KL(y, y_2) + KL(y, y_3)LYperiod=KL(QYtQYs)\mathcal{L}_{Yperiod} = KL(Q_{Yt}\parallel Q_{Ys})
谁和谁对齐自蒸馏:DIR 分支分类概率 y(老师)vs 中间层分类概率 y₁/y₂/y₃(学生)教师预测频谱 vs 学生预测频谱
学什么类别间的“暗知识”(软概率权重)教师捕捉到的周期模式(哪些频率占主导)
应用领域目标检测时序预测

两种用法的共同模式:先把“想对齐的东西”变成合法概率分布(softmax),再用 KL 衡量分布差异——KL 的输入必须是概率分布,这是它和余弦相似度最大的区别(余弦直接吃任意向量)。

启发:KL 的适用范围比“分类蒸馏”更广——只要你能把任意信号(频谱、注意力、关系矩阵)归一化成概率分布,KL 就能用来对齐它。这解释了为什么 TimeDistill 能把 FFT 频谱“当概率用”:先 softmax 归一化,再 KL 对齐,本质上是在对齐“周期的概率质量分配”。

4 工程实践避坑

  1. 方向不能搞反DKL(PQ)D_{KL}(P\parallel Q)DKL(QP)D_{KL}(Q\parallel P) 行为完全不同,需要明确谁是目标分布 P,谁是模型输出 Q。
  2. 输入必须是概率:原始 logits 不能直接输入 KL,需要经过 softmax;PyTorch 接口 F.kl_div(input=log_q, target=p),第一个参数要传 log 之后的模型概率。
  3. 零概率数值稳定性:当 P(x)>0,Q(x) 趋近 0,KL 会趋向无穷大。工程上通常加极小 epsilon 做数值截断。
  4. 什么时候不要用 KL:原始特征向量没有概率含义,不要强行用 KL;此时优先余弦相似度、MSE。

5 延伸:现实计算的痛点(大模型视角)

理论 KL 需要遍历整个词表/全部离散状态求和。当分布维度极高(LLM 词表 50k+),完整计算 KL 开销爆炸,显存难以承受。实际训练中往往用蒙特卡洛采样、重要性采样、k3 控制变量估计器做 KL 的近似估计,而不是直接算理论公式,PPO、GRPO 中都大量使用这类估计技巧。

参考资料

视频

文章


在以下平台分享此文章:

上一篇
论文速览手册
下一篇
AdaIN:从“调整特征统计量”到任意风格迁移