dev.to #ai短讯
第3a节:f-散度
作者在学习《生成式AI的数学基础》课程,本节笔记聚焦于f-散度家族。通过单个凸函数f来定义和选择特定的散度度量,作为理解生成模型一般原理的补充内容。
我正在研读 Prathosh AP 教授的公开课《生成式 AI 的数学基础》。该播放列表构成了课程的主干。这些笔记是我对每个章节的深度解读:当图像即重点时提供可视化,当先决条件发挥关键作用时进行旁支拓展,并用我自己的话重新表述公式。
此前:第 2 节:生成模型的一般原理。
第 3a 节:f-散度
第 2 节留下了一个关于散度的空白。本笔记用一族散度填补了这一空白。一个凸函数 $f$ 决定了评分标准。尚未撰写的第 3b 节将是凸共轭:这是一种从样本中估计 $D_f$ 的重写形式,也是通向 GANs(生成对抗网络)的桥梁。
我要阐述的核心命题如下:
给定两个具有密度函数 $P_x$ 和 $P_\theta$ 的分布,$$ D_f(P_x \,|\, P_\theta) = \int_{\mathcal{X}} P_\theta(x)\, f!\left(\frac{P_x(x)}{P_\theta(x)}\right) dx $$ 其中 $f: \mathbb{R}+ \to \mathbb{R}$ 是凸函数且下半连续,并满足 $f(1) = 0$。那么 $D_f \ge 0$,且当且仅当 $P_x = P\theta$ 时 $D_f = 0$。
直觉理解
在每一个点 $x$ 处,我比较数据在此处的概率与模型在此处的概率。比率 $P_x / P_\theta$ 在两者一致时为 1,在不一致时为其他值。$f$ 将这种局部不匹配转化为惩罚项。积分则对这些惩罚取平均。不同的 $f$ 惩罚不同类型的失配,因此一个模板便衍生出一族散度。
通过端点进行的负载测试
我将合成流量与真实的访问日志按端点进行对比。
$$ \text{ratio} = \frac{\text{真实流量占比}}{\text{合成流量占比}} $$
比率为 1 意味着生成器在该端点上表现正确。比率为 3 意味着真实用户访问该端点的频率是生成器的三倍。比率为 0.2 意味着生成器过度生成某个真实用户几乎不触碰的端点。
单一评分取决于我们关注哪种失败模式。
- 缺失真实流量模式。当比率较大时,惩罚增加。这是前向 KL 散度。
- 真实用户绝不会发送的虚假流量。当比率接近 0 时,惩罚增加。这是反向 KL 散度。
- 一种不会发散至无穷大的平衡评分。这是 Jensen–Shannon 散度或总变差距离。
$f$ 是惩罚函数。积分是对所有端点的平均。
可视化
$P_x$ 是一个以 0 为中心的标准钟形曲线。$P_\theta$ 是将同一钟形曲线平移至 $\mu$ 后的结果。选定一个 $f$。顶部曲线为惩罚项,它经过点 $(1, 0)$。珊瑚色曲线是被积函数 $P_\theta(x)\, f(P_x(x)/P_\theta(x))$。四个数字则是积分值。
在移动 $\mu$ 时,我希望眼前呈现以下三个要点:
- 当 $\mu = 0$ 时,所有散度均为 0。分布匹配意味着对于所有满足 $f(1) = 0$ 的 $f$,惩罚为零。
- 对于前向 KL 散度,珊瑚色曲线在某些区域低于零轴。单个点可能贡献负值。但总和永远不会为负。下方的证明解释了原因。
- 随着 $\mu$ 增大,Jensen–Shannon 散度和总变差距离趋于平缓。两者均有界。而两种 KL 散度则持续上升。这种饱和现象是后续解释 GAN(旨在优化 Jensen–Shannon 散度)在数据与模型不重叠时可能停止学习的原因。
旁支拓展
密度比率。$u(x) = P_x(x) / P_\theta(u)$ 询问的是:在点 $x$ 处,真实数据出现的概率是模型数据的多少倍。
- $u = 1$:在该点达成一致。
- $u > 1$:模型在此处产出不足。
- $u < 1$:模型在此处产出过剩。
- $u = 0$:模型在数据从未出现的地方分配了概率质量。
如果 $P_x = P_\theta$ 处处成立,那么 $u(x) = 1$ 处处成立。$f$ 仅能观测到这个比率。
凸性。当一个函数呈碗状时,它是凸的:其图像上任意两点之间的直线段位于曲线之上或与曲线重合。对于 $\lambda \in [0, 1]$,
$$ \lambda f(u_1) + (1-\lambda) f(u_2) \;\ge\; f\big(\lambda u_1 + (1-\lambda) u_2\big) $$
输出的平均值至少等于平均值的输出。用 $f(u) = u^2$,$u_1 = 0$,$u_2 = 2$,$\lambda = \tfrac{1}{2}$ 来验证。输出的平均值是 $\tfrac{1}{2}(0 + 4) = 2$。平均值的输出是 $f(1) = 1$。且 $2 \ge 1$。
凸性是迫使 $D_f \ge 0$ 的原因。
期望作为积分。对于连续密度 $p$,$h(x)$ 的期望值是概率加权的平均值:
$$ \mathbb{E}_{x \sim p}[h(x)] = \int_{\mathcal{X}} p(x)\, h(x)\, dx $$
詹森不等式(Jensen's inequality)。对于凸函数 $f$ 和随机变量 $U$,
$$ \mathbb{E}[f(U)] \;\ge\; f(\mathbb{E}[U]) $$
这是两点定义的推广,适用于任何平均值。在碗状曲面上,曲面上点的平均值落在平均值对应值的上方,绝不会低于它。
公式
$$ D_f(P_x \,|\, P_\theta) = \int_{\mathcal{X}} P_\theta(x)\, f^*\left(\frac{P_x(x)}{P_\theta(x)}\right) dx = \mathbb{E}_{x \sim P_\theta}\left[f^*\left(\frac{P_x(x)}{P_\theta(x)}\right)\right] $$
朗读出来:从模型中抽取的点上的平均惩罚,衡量两个密度在每个点上的不匹配程度。
| 符号 | 含义 | 类型/形状 | 作用 |
|---|---|---|---|
| $P_x(x)$ | $x$ 处的真实密度 | 标量 $\ge 0$ | 此处数据的权重 |
| $P_\theta(x)$ | $x$ 处的模型密度 | 标量 $\ge 0$ | 模型的权重,也是平均中的权重 |
| $u = P_x/P_\theta$ | 密度比 | 标量 $\ge 0$ | 局部不匹配程度 |
| $f$ | 散度的生成元 | 凸函数,$f(1) = 0$ | 将不匹配转化为惩罚 |
| $\int \cdot\, dx$ | 对 $\mathcal{X}$ 的积分 | 运算 | 累加惩罚值 |
| $D_f$ | 散度 | 标量 $\ge 0$ | 总分 |
$f$ 的每个条件都有其存在的理由:
- $f(1) = 0$。如果分布匹配,则处处 $u = 1$,因此 $D_f = \int P_\theta \cdot 0\, dx = 0$。
- 凸性。这是下面证明 $D_f \ge 0$ 的关键。
- 下半连续。没有突然的向下跳跃。第 3b 节需要这一性质,以确保共轭函数行为良好。
证明 $D_f \ge 0$
讲义陈述了这一性质。以下是简短的证明。
$$ \begin{aligned} D_f(P_x \,|\, P_\theta) &= \mathbb{E}_{x \sim P_\theta}\left[f^*\left(\tfrac{P_x(x)}{P_\theta(x)}\right)\right] \ &\ge f^*\left(\mathbb{E}_{x \sim P_\theta}\left[\tfrac{P_x(x)}{P_\theta(x)}\right]\right) \ &= f^*\left(\int P_\theta(x)\, \tfrac{P_x(x)}{P_\theta(x)}\, dx\right) \ &= f^*\left(\int P_x(x)\, dx\right) \ &= f(1) = 0 \end{aligned} $$
第一行是期望形式。第二行应用了詹森不等式。第三行将期望写为积分。$P_\theta$ 被消去。密度函数的积分为 1,且 $f(1) = 0$。
用通俗的话说:以模型为权重的平均比值恰好为 1。一个碗状的惩罚函数在 1 附近取平均,其值不可能低于它在 1 处的值,即 0。这就是为什么珊瑚曲线可以在局部为负,而总和保持非负的原因。
一个模板,四种 $f$ 的选择
| 名称 | $f(u)$ | 产生的散度 | 行为特征 | |
|---|---|---|---|---|
| 前向 KL | $u \log u$ | $\displaystyle\int P_x \log\frac{P_x}{P_\theta}\, dx = D_{\mathrm{KL}}(P_x \, | \, P_\theta)$ | 惩罚模型遗漏真实数据的情况。模式覆盖型。这是第 2 节中的模糊团块,对应最大似然估计。 |
| 后向 KL | $-\log u$ | $\displaystyle\int P_\theta \log\frac{P_\theta}{P_x}\, dx = D_{\mathrm{KL}}(P_\theta \, | \, P_x)$ | 惩罚模型在数据不存在的地方生成的情况。模式寻找型。它只集中在一个峰值上。 |
| Jensen–Shannon | $\tfrac{1}{2}\big[u \log u - (u+1)\log\tfrac{u+1}{2}\big]$ | 对称,有界于 $\log 2$ | 原始 GAN 最小化的目标。 | |
| 总变差 | $\tfrac{1}{2}\lvert u - 1 \rvert$ | $\tfrac{1}{2}\int \lvert P_x - P_\theta \rvert\, dx$ | 对称,有界于 1。 |
前向 KL 是使用 $f(u) = u \log u$ 的模板,且 $P_\theta$ 被消去:
$$ \int P_\theta \cdot \frac{P_x}{P_\theta} \log\frac{P_x}{P_\theta}\, dx = \int P_x \log\frac{P_x}{P_\theta}\, dx $$
反向KL对应的是 $f(u) = -\log u$ 的模板。代入后可恢复 $\int P_\theta \log(P_\theta / P_x)\, dx$。上述Jensen–Shannon $f$ 满足 $f(1) = 0$。GAN论文中常使用 $f(u) = u \log u - (u+1)\log(u+1)$,它与该函数仅相差一个常数偏移,这也是为什么该损失被描述为类似于Jensen–Shannon而非完全相同的原因。总变差距离为 $\tfrac{1}{2}\lvert u - 1 \rvert$,这一选择将模板转化为 $\tfrac{1}{2}\int \lvert P_x - P_\theta \rvert$。
一个足够小、可以手工计算的示例
只有两个结果。$P_x = (0.5, 0.5)$,一枚公平硬币。$P_\theta = (0.8, 0.2)$,一个有偏的模型。
比率分别为 $u_1 = 0.5 / 0.8 = 0.625$(过度生成)和 $u_2 = 0.5 / 0.2 = 2.5$(欠生成)。有限集上的模板为 $D_f = \sum_i P_\theta(i)\, f(u_i)$。
| 散度 | 计算过程 | 值 |
|---|---|---|
| 前向KL | $0.8(0.625 \ln 0.625) + 0.2(2.5 \ln 2.5) = 0.8(-0.294) + 0.2(2.291)$ | $\approx 0.223$ |
| 反向KL | $0.8(-\ln 0.625) + 0.2(-\ln 2.5) = 0.8(0.470) + 0.2(-0.916)$ | $\approx 0.193$ |
| Jensen–Shannon | $0.8\, f(0.625) + 0.2\, f(2.5) = 0.8(0.0218) + 0.2(0.1660)$ | $\approx 0.051$ |
| 总变差 | $0.8 \cdot \tfrac{1}{2}(0.375) + 0.2 \cdot \tfrac{1}{2}(1.5) = 0.15 + 0.15$ | $0.300$ |
这四个值均为正数,且它们对差距大小的判断不一致。前向KL不等于反向KL,因此该得分不具有对称性。在前向KL行中,结果1贡献了约 $-0.235$,结果2贡献了约 $+0.458$。存在负项,总和为正。这就是Jensen不等式的体现。如果 $P_\theta = (0.5, 0.5)$,则每个比率都为1,每一行的计算结果均为0。
关键难点,即第3b节的内容
计算 $D_f$ 需要知道 $P_x(x)$ 和 $P_\theta(x)$ 的值。
- $P_x$ 是未知的。第1节只给了我样本。
- $P_\theta$ 是隐式的。第2节允许我采样 $g_\theta(z)$,但不能评估密度值。
因此,我无法构建比率 $u(x)$。根据大数定律,我可以基于样本构建平均值。第3b节必须将 $D_f$ 重写为在 $P_x$ 和 $P_\theta$ 下的期望形式,公式中不包含任何密度值。凸共轭是所用的工具。最终结果即为GAN的目标函数。
本笔记在课程中的定位
- 本笔记回答“哪种散度?”这一问题时,提供的是一个菜单选项,而非单一选择。
- GAN选择了类似Jensen–Shannon的 $f$。
- VAE、扩散模型和自回归模型最小化前向KL,这等同于最大似然估计。
- 当 $P_x$ 和 $P_\theta$ 不重叠时,Jensen–Shannon等受限得分会饱和。Wasserstein距离(它不是f-散度)是后来针对这一问题提出的解决方案。
- 使对齐模型接近参考模型的KL项 $D_{\mathrm{KL}}(\pi_\theta \,|\, \pi_{\mathrm{ref}})$ 也属于这一家族。
我希望能够回答的问题:
- 为什么必须满足 $f(1) = 0$?如果 $f(1) = 5$,会发生什么错误?
- 在上述示例中,结果1对前向KL的贡献为负值。为什么这不会破坏 $D_f \ge 0$ 的性质?
- 为什么我不能直接将数据集代入 $D_f$ 公式并直接计算它?
译文已达到本站中文翻译的字数上限,剩余内容请查看原文。