跳转至

KL 散度数学推导

前言

主要是在看 VAE(Variational Auto-encoder) 的时候,VAE的损失函数涉及到 KL 散度

1784188975915

\[\mathcal{L}_{VAE} = \mathcal{L}_{rec} + \mathcal{L}_{KL}\]

其中 \(\mathcal{L}_{KL} = \frac{1}{2} \sum_{i=1}^{d}(\mu_i^2+\sigma_i^2-\log \sigma_i^2 - 1)\)

作用是让编码器产生的潜向量分布 \(q_{\phi}(z | x)\),尽量接近标准正态分布 \(p(z) = N(0, I)\)

这个损失函数里出现 KL 散度的意味是

VAE的过程:

flowchart LR
A["x"] -- Encoder --> B["q_φ(z|x)"]
B -- 采样 z --> C["p_θ(x|z)"]
  • \(x\):输入图片z:潜变量,也就是压缩后的一组数字
  • \(q_\phi(z|x)\):编码器给出的潜变量分布
  • \(p(z)\):我们希望潜变量服从的先验分布,通常设成标准正态分布
  • \(p_\theta(x|z)\):解码器根据 z 生成图片的概率模型
  • \(\phi\):编码器的参数
  • \(\theta\):解码器的参数

目标是最大化从潜向量中得出训练输入的图片x的概率,即最大化\(p_{\theta}(x)\)

但是

\[p_{\theta}(x) = \int p_{\theta}(x, z) dz\]

要把所有可能的潜变量 z 都积分一遍,很难直接计算

所以 VAE 引入了一个由编码器产生的近似分布: \(q_{\phi}(z|x)\)

\[p_{\theta}(z) = \int q_{\phi}(z | x) \frac{ p_{\theta}(x, z)}{q_{\phi}(z | x)} dz = \mathbb{E}_{q_{\phi}(z|x)} [\frac{ p_{\theta}(x, z)}{q_{\phi}(z | x)}]\]
期望数学定义
\[\mathbb{E}_{z \sim q(z)}[f(z)] = \int q(z) f(z) dz\]

从分布\(q_{z \sim q_{\phi}(z|x)}\) 中采样\(z\),计算\(f(z)\),然后求平均

取对数,然后用琴生不等式

\[\log p_{theta}(x) = \log \mathbb{E}_{q_{\phi}(z|x)} [\frac{ p_{\theta}(x, z)}{q_{\phi}(z | x)}] \geq \mathbb{E}_{q_{\phi}(z|x)} [\log \frac{ p_{\theta}(x, z)}{q_{\phi}(z | x)}]\]

我们称\(LHS\)为ELBO(Evidence Lower Bound) 证据下界/ 对数近似下界

接下来我们把条件概率拆开来

\[ELBO = \mathbb{E}_q [\log \frac{ p_{\theta}(x, z)}{q_{\phi}(z | x)}] = \mathbb{E}_q [\log \frac{ p_{\theta}(x| z) p(z)}{q_{\phi}(z | x)}] = \mathbb{E}_q [\log p_{\theta}(x| z)] + \mathbb{E}_q [\log \frac{ p(z)}{q_{\phi}(z | x)}]\]

KL 散度的数学定义

\[D_{KL}(q||p) = \mathbb{E}_q [\log \frac{q}{p}]\]

加个负号: $\(\mathbb{E}_q[\log \frac{p}{q}] = -D_{KL}(q||p)\)$

代换可得:

\[-ELBO = -\mathbb{E}_q [\log p_{\theta}(x| z)] + D_{KL}(q_{\phi}(z|x)||p(z))\]

即:

\[\mathcal{L}_{VAE} = \mathcal{L}_{rec} + \mathcal{L}_{KL}\]

\(\mathcal{L}_{rec}\)叫图片重建误差

KL 散度

\[D_{KL}(q||p) = \int q(z) \log \frac{q(z)}{p(z)}dz\]

也可以写成:

\[D_{KL}(q||p) = \mathbb{E}_{z \sim q}[\log q(z) - \log p(z)]\]

可以理解为:当真实使用的分布是 q,但你想用 p 来描述它时,会产生多大的差异或额外代价。

继续推导

假设编码器产生的分布是

\[q(z|x) = N(\mu, \sigma^2)\]

我们希望他解决标准正态分布:

\[p(z) = N(0, 1)\]

目标是计算:

\[D_{KL}(q||p) = \mathbb{E}_q [\log q(z) - \log p(z)]\]

两个正态分布的概率密度是:

\[ q(z) = \frac{1}{\sqrt{2\pi\sigma^2}} \exp\left(-\frac{(z - \mu)^2}{2\sigma^2}\right) \]
\[ p(z) = \frac{1}{\sqrt{2\pi}} \exp\left(-\frac{z^2}{2}\right) \]

取对数。

对于 \(q(z)\)

\[ \log q(z) = -\frac{1}{2}\log(2\pi) - \frac{1}{2}\log \sigma^2 - \frac{(z - \mu)^2}{2\sigma^2} \]

对于 \(p(z)\)

\[ \log p(z) = -\frac{1}{2}\log(2\pi) - \frac{z^2}{2} \]

相减,整理一下

\[\log q(z) - \log p(z) = \frac{1}{2} [z^2 - \frac{(z - \mu)^2}{\sigma} - \log \sigma^2]\]
\[D_{KL}(q||p) = \mathbb{E}_q [\log q(z) - \log p(z)] = \frac{1}{2} [\mathbb{E}_q[z^2] - \frac{\mathbb{E}_q [(z - \mu)^2]}{\sigma^2} - \log \sigma^2]\]
  • \(\mathbb{E}_q [(z - \mu)^2] = \sigma^2\)
  • \(Var(z) = \mathbb{E}[z^2] - \mathbb{E}[z]^2 \rightarrow \mathbb{E}_q [z^2] = \sigma^2 + \mu^2\)

代回去:

\[D_{KL}(q||p) = \frac{1}{2} [\sigma^2 + \mu^2 - 1 - \log \sigma^2]\]

上面讲的其实只是一维情况,多维情况要求和

\[D_{\text{KL}}\big(q_\phi(z|x) \,\|\, p(z)\big) = \frac{1}{2} \sum_{i=1}^{d} \Big( \mu_i^2 + \sigma_i^2 - \log \sigma_i^2 - 1 \Big)\]

评论