生成模型 (1.1):变分推断

生成模型 #生成模型 #变分推断 约 21 分钟 · 7221 字

引言

在这一篇文章中,我们从变分 (Variation) 的视角来看待生成模型。

在讲解诸如VAE、DDPM等经典方法之前,我想先画一些时间讲清楚什么是【变分】。

我一开始学习VAE时,发现VAE中的隐变量是服从高斯分布的,而AE中的隐变量仍然服从某个未知的(往往也是难解的)隐分布。这个时候我可能认为变分就是引入一个隐分布的先验假设。在学习DDPM时,发现优化的ELBO损失是似然函数的“变分下界”,这里又出现了变分,但似乎和高斯分布之类的没有任何关系,然后我就搞不清楚到底变分的意义是什么了。

到后面我了解到数学中有一类称为【变分法】的方法,被用于求解泛函的极值问题。这似乎和优化问题扯上了一些关系,但也是死活没发现泛函体现在哪。

因此,我希望在学习生成模型之前,先把变分(准确来说是【变分推断】)的含义理解清楚。

在本文中,我们将看到数学上的变分法和机器学习中的变分推断实际上都是在处理某个函数空间中的优化问题,其中变分法用于求解泛函的极值所对应的最优函数,而变分推断则用于求解与难解的后验分布最接近且可解的概率分布。

以下的内容主要来自知乎中的高赞回答 和博客。

一、什么是推断

在机器学习领域中,推断 (Inference) 是指从观测数据 计算未知量 的分布的过程,即计算后验分布 的过程。

这里的未知量 可以包括多种,比如说模型参数、隐变量,甚至是未来的观测数据。

推断主要用到的数学工具是贝叶斯定理,因此也被称为贝叶斯推断 (Bayesian Inference):

(1)

其中:

  • 称为先验 (prior),是对未知量 的一种无知识的估计
  • 称为后验 (posterior)
  • 称为证据 (evidence)
  • 称为似然 (likelihood)

从这里可以看出机器学习的两种不同流派:

  1. 概率学派:利用最大似然对未知量 进行点估计,即 ;
  2. 贝叶斯学派:估计未知量 的整个分布,即 ;

进行贝叶斯推断主要有以下三个步骤:

  1. 根据问题,选定先验 以及似然 所服从的分布;

  2. 计算证据:

    (2)
  3. 根据公式 计算后验概率。

可以发现,整个推断的复杂度几乎全部集中在第二步,即求解公式 的积分上。公式 实际上是一个多维积分,其复杂度取决于隐变量 的维度。在实际的机器学习问题中,隐变量(如模型参数)的维度通常都非常高,可能是成千上万维的一个高维向量。因此,在解决实际问题时贝叶斯推断几乎不可用。这也是为什么我们称 为难解的 (intractable)。

二、什么是变分推断

2.1. 从计算转为优化

为了解决贝叶斯推断复杂度过高的问题,我们可以使用变分推断 (Variational Inference) 来近似后验概率 。

变分推断的核心步骤包括以下两步:

  1. 引入一个带参数 的变分分布 (variational distribution) ,这个分布是可解的;
  2. 通过优化参数 ,使得变分分布 尽可能接近真实后验分布 。

可以看到,变分推断实际上是把贝叶斯推断的计算问题变为优化问题,通过优化参数 来逼近后验分布:

(3)

优化问题收敛后,我们就可以用可解的 分布来替代难解的 分布。

公式 中的 是度量两个分布之间距离的度量。一般来说,我们采用KL散度作为这个度量。在上一篇文章中,我们已经介绍了KL散度:

(4)

KL散度满足以下性质:,当且仅当 时等号成立。

2.2. ELBO的引入

上面提到,变分推断的目标是求解以下最优化问题:

(5)

我们已经知道,后验分布 是难解的,因此我们需要做适当的变形。

首先根据定义展开KL散度:

(6)

我们称第二项为经验下界 (Evidence Lower Bound, ELBO),因为它是经验 的一个下界估计:

(7)

ELBO项可以进一步做如下变形:

(8)

因此,我们的优化问题变为:

(9)

即转变为最大化ELBO项。

2.3. 求解变分推断

下面介绍两种求解变分推断问题的实际方法

2.3.1. 平均场变分族

平均场变分族 (Mean-Field Variational Family) 基于的核心假设是:隐变量的不同分量 之间是相互独立的。因此,变分分布 可以分解为:

(10)

在这个假设下,我们可以通过坐标上升变分推断 (Coordinate Ascent Variational Inference, CAVI) 方法进行优化。

我们对ELBO的形式进行一定推导:

(11)

我们记 表示除 之外的隐变量。我们在固定 的情况下,希望优化 。

考虑ELBO关于 的泛函:

(12)

记 ,则:

(13)

注意 满足约束条件

(14)

我们通过拉格朗日乘子法来求泛函 的极值。构造如下的拉格朗日函数:

(15)

对 求导得:

(16)

令导数为0得:

(17)

因此,最优变分分布满足以下性质:

(18)

2.3.2. 黑盒变分推断

黑盒变分推断 (Black Box Variational Inference, BBVI) 是一种通用的、模型无关的变分推断方法,核心思想是将变分推断转化为一个可以通过随机梯度优化的问题,无需为特定模型推导复杂的更新方程。

我们将ELBO看作一个关于参数 的函数,则求偏导得:

(19)

其中,

(20)

代入得

(21)

上面的公式可以写成随机梯度下降的形式来优化:

(22)