3. 变分自编码器
Reference:[2022 Auto-Encoding Variational Bayes]
标准的变分自编码器(Variational autoencoder)主要用于建模连续的隐变量,对离散隐变量的建模因为采样过程无法传递梯度信息存在问题,但可利用Gumbel-Softmax、STE、梯度旋转(见RQ-VAE github)等技术扩展到离散隐变量,此外 VAE 支持对离散和连续的观测变量进行建模。
考虑 N 个独立同分布(i.i.d.)样本的数据集 {x(i)}i=1N,样本 x(i) 是连续或离散的变量,假设数据由一个随机过程生成,该过程涉及一个未被观测到的连续随机变量 z,具体两步为:(1)从某个先验分布 pθ∗(z) 中生成一个 z(i);(2)在给定条件分布 pθ∗(x∣z) 下生成观测样本 x(i)。假设先验分布 pθ∗(z) 和条件分布 pθ∗(x∣z) 来自参数化的分布族 pθ(z) 和 pθ(x∣z),并且他们的概率密度函数(PDFs)或概率质量函数(PMFs)几乎所有位置关于 θ 和 z 都是可微的。
VAE想解决两个挑战:
1、Intractability:当 z 是高维连续变量,且 pθ(x∣z) 是复杂的函数时,边际似然 pθ(x)=∫pθ(z)pθ(x∣z)dz 是难解的,导致无法直接使用以计算损失函数和更新 θ;推断隐变量的真实后验也存在困难,因为 pθ(z∣x)=pθ(x∣z)pθ(z)/pθ(x),导致 EM 算法无法使用,因为 E 步需要真实的后验分布(由式(9)可知);在不具备共轭性质的复杂模型中,即使是使用平均场假设,变分推断过程中的积分式 Eq[lnp(x,z)] (由式(65)可知)也无法得到解析解。
2、A large dataset: 数据集过大,以至于批量(更准确地说是全量)优化成本过高;希望使用小批量甚至单个数据点来更新参数。蒙特卡洛 EM 等采样方法通常会很慢,因为它涉及为每个数据点执行昂贵的采样循环。
3.1 变分下界
由式(24)或式(63),可以得到每个数据点的边际似然:
logpθ(x(i))=∫qϕ(z∣x(i))logqϕ(z∣x(i))pθ(x(i),z)dz−∫qϕ(z∣x(i))logqϕ(z∣x(i))pθ(z∣x(i))dz=L(θ,ϕ;x(i))+DKL(qϕ(z∣x(i))∣∣pθ(z∣x(i)))(146)
因为 KL 散度的非负性,所以
\log p_{\theta}(\mathbf{x}^{(i)}) \geq \mathcal{L}(\theta, \phi; \mathbf{x}^{(i)}) = \mathbb{E}_{q_{\phi}(\mathbf{z} \mid \mathbf{x}^{(i)})}[-\log q_{\phi}(\mathbf{z}\mid \mathbf{x}^{(i)}) + \log p_{\theta}(\mathbf{x}^{(i)}, \mathbf z)] \tag{147}
变分下界还可以写为
L(θ,ϕ;x(i))=Eqϕ(z∣x(i))[−logpθ(z)qϕ(z∣x(i))+logpθ(z)pθ(x(i),z)]=−DKL(qϕ(z∣x(i))∣∣pθ(z))+Eqϕ(z∣x(i))[logpθ(x(i)∣z)](148)
欲针对变分参数 ϕ 和生成参数 θ 求导和优化变分下界 L(θ,ϕ;x(i)),然而下界关于变分参数 ϕ 的梯度计算存在一些问题。通常使用一般的蒙特卡洛梯度估计(求导和积分的互换见[2020 Monte Carlo Gradient Estimation in Machine Learning]:
\triangledown_{\phi} \mathbb{E}_{q_{\phi}(\mathbf{z})}[f(\mathbf{z})]=\mathbb{E}_{q_{\phi}(\mathbf{z})}[f(\mathbf{z})\triangledown_{\phi}\log q_{\phi}(\mathbf{z})]\simeq \frac{1}{L} \sum_{l=1}^L f(\mathbf{z}) \triangledown_{\phi}\log q_{\phi}(\mathbf{z}^{(l)}) \tag{149}
此处 z(l)∼qϕ(z∣x(i)),然而这样基于得分函数的蒙特卡洛估计方差非常大(见[2012 Variational Bayesian inference with Stochastic Search]),无法适用于此文研究(但在强化学习中常见,如REINFORCE)。
3.2 Stochastic Gradient Variational Bayes (SGVB) estimator
注意假设的近似后验具有 qϕ(z∣x) 的形式,但是同样可以假设 qϕ(z),主要用于全局模型参数。需注意式(65)中的 q(Z) 如果是近似局部参数严格表述应为 q(zi;λi)。
在相对温和的条件限制下,可以用一个可微分的转换 gϕ(ϵ,z) 和一个辅助(噪声)变量 ϵ 为近似后验 qϕ(z∣x) 重参数化变量 z~∼qϕ(z∣x):
\tilde{\mathbf z} = g_{\phi}(\epsilon, \mathbf x) \quad \text{with } \epsilon \sim p(\epsilon) \tag{150}
然后我们可以获得 qϕ(z∣x) 关于函数 f(z) 的期望的蒙特卡洛估计:
\mathbb{E}_{q_{\phi}(\mathbf z\mid \mathbf x^{(i)})}[f(\mathbf z)] = \mathbb{E}_{p(\epsilon)}\left[f(g_{\phi}(\epsilon, \mathbf x^{(i)})) \right] \simeq \frac{1}{L} \sum_{l=1}^L f(g_{\phi}(\epsilon^{(l)}, \mathbf x^{(i)})) \quad \text{where } \epsilon^{(l)} \sim p(\epsilon) \tag{151}
将该技术作用于式(147),得到第一种 SGVB 估计器 L~A(θ,ϕ;x(i))≃L(θ,ϕ;x(i)) (注意其为 ELBO 需最大化):
\tilde{\mathcal{L}}^A (\theta, \phi; \mathbf x^{(i)}) = \frac{1}{L} \sum_{l=1}^L \left[\log p_{\theta}(\mathbf x^{(i)}, \mathbf z^{(i,l)})-\log q_{\phi}(\mathbf z^{(i,l)}\mid \mathbf x^{(i)})\right] \quad \text{where } \mathbf z^{(i,l)}=g_{\phi}(\epsilon^{(i,l)}, \mathbf x^{(i)}), \epsilon^{(l)} \sim p(\epsilon) \tag{152}
通常式(148)中的 DKL(qϕ(z∣x(i))∣∣pθ(z)) 可以解析的计算,因此只有期望的重构误差 Eqϕ(z∣x(i))[logpθ(x(i)∣z)] 需要使用采样进行估计。这个 KL 散度项可以被理解为规范变分参数 ϕ,促使估计后验尽量和先验 pθ(z) 接近。因此针对式(148)可以得到第二种 SGVB 估计器 L~B(θ,ϕ;x(i))≃L(θ,ϕ;x(i)),相较于通用估计器具有更小的方差:
L~B(θ,ϕ;x(i))=−DKL(qϕ(z∣x(i))∣∣pθ(z))+L1l=1∑L[logpθ(x(i)∣z(i,l))]where z(i,l)=gϕ(ϵ(i,l),x(i)),ϵ(l)∼p(ϵ)(153)
给定具有 N 个数据点的数据集 X 中的多个数据点,我们可以基于小批量构造全数据集的边际似然的下界(注意这是为了得到整个数据集的无偏估计,实践中使用 batch 内平均损失更稳健;使用平均损失时 Full VB 要注意 weight decay 的设置):
\mathcal{L}(\theta, \phi; X) \simeq \tilde{\mathcal{L}}^M (\theta, \phi; X^M) = \frac{N}{M} \sum_{i=1}^M \tilde{\mathcal{L}} (\theta, \phi; \mathbf x^{(i)}) \tag{154}
作者在实验中发现只要小批量 M 足够大,例如 M=100,每个数据点采样数 L 可以设置为1(因为批次内各个点的独立随机噪声会相互抵消)。
算法5:
初始化原始参数 θ 和变分参数 ϕ
重复:
从完整数据集 X 中选取小批量数据 XM
从噪声分布 p(ϵ) 中抽取样本 ϵ
计算梯度 g←▽θ,ϕL~M(θ,ϕ;XM,ϵ)
利用梯度 g 更新 参数 (θ,ϕ)
直到参数 (θ,ϕ) 收敛
3.3 完全变分贝叶斯
目标是将 VAE 的重参数化技巧和随机梯度下降方法,从仅仅推断局部隐变量 z (latent variables) 扩展到同时推断模型的全局参数 θ (global parameters) 。注意在前文中参数 θ 是被当做固定的未知量,通过最大似然(或 MAP)来寻找其点估计 。但在本节中将把 θ 也看作是一个随机变量,给它设定一个先验分布,并使用变分推断来求它的后验分布。
首先,假设参数 θ 服从一个超先验分布 pα(θ),其中 α 是超参数。
\log p_\alpha(X) = D_{KL}(q_\phi(\theta) || p_\alpha(\theta|X)) + \mathcal{L}(\phi; X) \tag{155}
要最大化整个数据集 X 的边际似然 logpα(X),只需要最大化后面的变分下界(由式(146)可知):
\mathcal{L}(\phi; X) = \int q_\phi(\theta) \left( \log p_\theta(X) + \log p_\alpha(\theta) - \log q_\phi(\theta) \right) \mathrm{d}\theta \tag{156}
因为数据集 X 由 N 个独立同分布的样本组成,所以 logpθ(X)=∑i=1Nlogpθ(x(i)) ,对于每一个单一样本再次使用变分推断来引入关于局部隐变量 z 的近似后验 qϕ(z∣x(i)) :
\log p_\theta(\mathbf{x}^{(i)}) = D_{KL}(q_\phi(\mathbf{z}\mid \mathbf{x}^{(i)}) \mid \mid p_\theta(\mathbf{z}|\mathbf{x}^{(i)})) + \mathcal{L}(\theta, \phi; \mathbf{x}^{(i)}) \tag{157}
单一样本的变分下界为
\mathcal{L}(\theta, \phi; \mathbf x^{(i)}) = \int q_\phi(\mathbf z\mid \mathbf x^{(i)}) \left( \log p_\theta(\mathbf x^{(i)}\mid \mathbf z) + \log p_\theta(\mathbf z) - \log q_\phi(\mathbf z \mid \mathbf x^{(i)}) \right) \mathrm{d}\mathbf z \tag{158}
分别对隐变量 z 和全局参数 θ 进行重参数化,可得
\tilde{\mathbf{z}} = g_\phi(\boldsymbol{\epsilon}, \mathbf{x}^{(i)}) \quad \text{其中 } \boldsymbol{\epsilon} \sim p(\boldsymbol{\epsilon}), \quad \tilde{\boldsymbol{\theta}} = h_\phi(\boldsymbol{\zeta}) \quad \text{其中 } \boldsymbol{\zeta} \sim p(\boldsymbol{\zeta}) \tag{159}
设简化记号(使用单样本估计,为了保证无偏性,必须乘回数据集大小 N)
f_\phi(\mathbf x, \mathbf z, \theta) = N \cdot \left( \log p_\theta(\mathbf x\mid \mathbf z) + \log p_\theta(\mathbf z) - \log q_\phi(\mathbf z \mid \mathbf x) \right) + \log p_\alpha(\boldsymbol \theta) - \log q_\phi(\boldsymbol \theta) \tag{160}
最终得到蒙特卡洛估计器(注意到在最终的估计器中省略了数据集的似然中的KL散度)
\mathcal{L}(\phi; X) \simeq \frac{1}{L} \sum_{l=1}^L f_\phi \left(\mathbf x^{(l)}, g_\phi(\boldsymbol \epsilon^{(l)}, \mathbf x^{(l)}), h_\phi(\boldsymbol \zeta^{(l)}) \right) \tag{161}
如果所有的先验和近似后验都是高斯分布,可以把式(160)中的一些项解析计算出来,从而进一步降低方差 。隐变量先验 pθ(z)=N(0,I),近似后验 qϕ(z∣x)=N(μz,σz2I) 。参数先验 pα(θ)=N(0,I),近似后验 qϕ(θ)=N(μθ,σθ2I) 。
此时,θ 和 z 都可以写成“均值 + 标准差 × 标准正态噪声”的形式 。最终的低方差估计器为
L(ϕ;X)≃L1l=1∑LN⋅(21j=1∑J(1+log((σz,j(l))2)−(μz,j(l))2−(σz,j(l))2)+logpθ(x(i)∣z(i)))+21j=1∑J(1+log((σθ,j(l))2)−(μθ,j(l))2−(σθ,j(l))2)(162)
3.4 重参数化技巧
令 z 为一个连续的随机变量(注意重参数技巧无法直接应用于离散分布),并且 z∼qϕ(z∣x) 作为条件分布,通常可以将其表示为一个确定性变量 z=gϕ(ϵ,x),其中 ϵ 是一个辅助变量,具有边际分布 p(ϵ),gϕ(⋅) 是以 ϕ 为参数的向量函数。重参数化技巧可以用来改写关于 qϕ(z∣x) 的期望,使得该期望的蒙特卡洛估计值关于 ϕ 可微。
首先根据变量代换定理,只要 z 和 ϵ 间为确定性映射,则
q_\phi(\mathbf{z}|\mathbf{x}) \prod_i dz_i = p(\boldsymbol{\epsilon}) \prod_i d\epsilon_i \tag{163}
改写期望
\int q_\phi(\mathbf{z}|\mathbf{x})f(\mathbf{z}) d\mathbf{z} = \int p(\boldsymbol{\epsilon})f(\mathbf{z}) d\boldsymbol{\epsilon} = \int p(\boldsymbol{\epsilon})f(g_\phi(\boldsymbol{\epsilon}, \mathbf{x})) d\boldsymbol{\epsilon} \tag{164}
这样积分的变量里就不含参数 ϕ,采样不阻断梯度的传导
\int q_\phi(\mathbf{z}|\mathbf{x})f(\mathbf{z}) d\mathbf{z} \simeq \frac{1}{L} \sum_{l=1}^L f(g_\phi(\boldsymbol{\epsilon}^{(l)}, \mathbf{x})) \quad \text{where } \boldsymbol{\epsilon}^{(l)} \sim p(\boldsymbol{\epsilon}) \tag{165}
以单变量高斯分布举例:使 z∼p(z∣x)=N(μ,σ2),一个有效的重参数化是 z=μ+σϵ,其中 ϵ 是辅助的噪声变量 ϵ∼N(0,1),因此
\mathbb{E}_{\mathcal{N}(z; \mu, \sigma^2)}[f(z)] = \mathbb{E}_{\mathcal{N}(\epsilon; 0, 1)}[f(\mu + \sigma \epsilon)] \simeq \frac{1}{L} \sum_{l=1}^L f(\mu + \sigma \epsilon^{(l)}) \quad \text{where } \epsilon^{(l)} \sim \mathcal{N}(0,1) \tag{166}
对于哪些 qϕ(z∣x) 可以选择这样的可微变换 gϕ(⋅) 和辅助变量 ϵ∼p(ϵ)?有三种基本方法:
1、易处理的逆 CDF。在这种情况下,令 ϵ∼U(0,1),并令 gϕ(ϵ,x) 是 qϕ(z∣x) 的逆 CDF。例如:指数分布、柯西分布、Logistic 分布、瑞利分布、帕累托分布、威布尔分布、倒数分布、Gompertz 分布、Gumbel 分布和 Erlang 分布。
2、类似于高斯分布的例子,对于任何“位置-尺度”(location-scale)分布族,可以选择标准分布(location=0, scale=1)作为辅助变量 ϵ,并令 g(⋅)=location+scale⋅ϵ。例如拉普拉斯分布、椭圆分布、学生 t 分布、Logistic 分布、均匀分布、三角分布和高斯分布。
3、组合,通常可以将随机变量表达为辅助变量的不同变换。例子:对数正态分布(正态分布变量的指数化)、伽马分布(指数分布变量的总和)、狄利克雷分布(伽马变量的加权总和)、贝塔分布、卡方分布和 F 分布。
当所有三种方法都失败时,存在对逆累积分布函数的良好近似,其计算的时间复杂度与概率密度函数相当(见[1986 Sample-based non-uniform random variate generation])。
3.5 示例:变分自编码器
假设隐变量的先验分布为 pθ(z)=N(z;0,I),这是简化的做法,通常出于
1、正态分布的良好计算性质,先验与近似后验的 KL 散度可以解析表示。
2、正则化,使其合理的聚集
3、协方差矩阵设为 I 保证各维度间相互独立,引导模型学习独立的特征
令 pθ(x∣z) 为多元高斯,并且分布参数为 z 经过一个 MLP 得到的值,注意此时后验分布 pθ(z∣x) 是没有解析解的。假设 qϕ(z∣x) 是一个具有对角协方差矩阵的多元高斯分布
logqϕ(z∣x(i))=logN(z;μ(i),σ2(i)I)(167)
其中 μ(i),σ2(i) 是数据点 x(i) 经过神经网络参数为变分参数 ϕ 的 MLP 得到的。因此可以采样得到 z(i,l)=gϕ(x(i),ϵ(i))=μ(i)+σ(i)⊙ϵ(l) where ϵ(l)∼N(0,I),并根据式(153)可以得到最终的损失函数
\mathcal{L}(\theta, \phi; \mathbf{x}^{(i)}) \simeq \frac{1}{2} \sum_{j=1}^J \left( 1 + \log((\sigma_j^{(i)})^2) - (\mu_j^{(i)})^2 - (\sigma_j^{(i)})^2 \right) + \frac{1}{L} \sum_{l=1}^L \log p_\theta(\mathbf{x}^{(i)}|\mathbf{z}^{(i,l)})\tag{168}
前半部分为两个正态分布的 KL 散度解析表达式,后半部分为重构损失。
3.6 复现实验
Loss A(1000 epoch,z_dim=128):
Loss B(1000 epoch,z_dim=128):