生成模型 (1.1):变分推断
引言
在这一篇文章中,我们从变分 (Variation) 的视角来看待生成模型。
在讲解诸如VAE、DDPM等经典方法之前,我想先画一些时间讲清楚什么是【变分】。
我一开始学习VAE时,发现VAE中的隐变量是服从高斯分布的,而AE中的隐变量仍然服从某个未知的(往往也是难解的)隐分布。这个时候我可能认为变分就是引入一个隐分布的先验假设。在学习DDPM时,发现优化的ELBO损失是似然函数的“变分下界”,这里又出现了变分,但似乎和高斯分布之类的没有任何关系,然后我就搞不清楚到底变分的意义是什么了。
到后面我了解到数学中有一类称为【变分法】的方法,被用于求解泛函的极值问题。这似乎和优化问题扯上了一些关系,但也是死活没发现泛函体现在哪。
因此,我希望在学习生成模型之前,先把变分(准确来说是【变分推断】)的含义理解清楚。
在本文中,我们将看到数学上的变分法和机器学习中的变分推断实际上都是在处理某个函数空间中的优化问题,其中变分法用于求解泛函的极值所对应的最优函数,而变分推断则用于求解与难解的后验分布最接近且可解的概率分布。
一、什么是推断
在机器学习领域中,推断 (Inference) 是指从观测数据
这里的未知量
推断主要用到的数学工具是贝叶斯定理,因此也被称为贝叶斯推断 (Bayesian Inference):
其中:
称为先验 (prior),是对未知量 的一种无知识的估计 称为后验 (posterior) 称为证据 (evidence) 称为似然 (likelihood)
从这里可以看出机器学习的两种不同流派:
- 概率学派:利用最大似然对未知量
进行点估计,即 ; - 贝叶斯学派:估计未知量
的整个分布,即 ;
进行贝叶斯推断主要有以下三个步骤:
-
根据问题,选定先验
以及似然 所服从的分布; -
计算证据:
(2) -
根据公式
计算后验概率。
可以发现,整个推断的复杂度几乎全部集中在第二步,即求解公式
二、什么是变分推断
2.1. 从计算转为优化
为了解决贝叶斯推断复杂度过高的问题,我们可以使用变分推断 (Variational Inference) 来近似后验概率
变分推断的核心步骤包括以下两步:
- 引入一个带参数
的变分分布 (variational distribution) ,这个分布是可解的; - 通过优化参数
,使得变分分布 尽可能接近真实后验分布 。
可以看到,变分推断实际上是把贝叶斯推断的计算问题变为优化问题,通过优化参数
优化问题收敛后,我们就可以用可解的
公式
KL散度满足以下性质:
2.2. ELBO的引入
上面提到,变分推断的目标是求解以下最优化问题:
我们已经知道,后验分布
首先根据定义展开KL散度:
我们称第二项为经验下界 (Evidence Lower Bound, ELBO),因为它是经验
ELBO项可以进一步做如下变形:
因此,我们的优化问题变为:
即转变为最大化ELBO项。
2.3. 求解变分推断
下面介绍两种求解变分推断问题的实际方法
2.3.1. 平均场变分族
平均场变分族 (Mean-Field Variational Family) 基于的核心假设是:隐变量的不同分量
在这个假设下,我们可以通过坐标上升变分推断 (Coordinate Ascent Variational Inference, CAVI) 方法进行优化。
我们对ELBO的形式进行一定推导:
我们记
考虑ELBO关于
记
注意
我们通过拉格朗日乘子法来求泛函
对
令导数为0得:
因此,最优变分分布满足以下性质:
2.3.2. 黑盒变分推断
黑盒变分推断 (Black Box Variational Inference, BBVI) 是一种通用的、模型无关的变分推断方法,核心思想是将变分推断转化为一个可以通过随机梯度优化的问题,无需为特定模型推导复杂的更新方程。
我们将ELBO看作一个关于参数
其中,
代入得
上面的公式可以写成随机梯度下降的形式来优化: