Variational Auto-Encoder的原理整理
Variational Auto-Encoder的原理整理
本文致谢:
论文地址:
Auto-Encoding Variational Bayes https://arxiv.org/abs/1312.6114
B站视频:
【大白话02】一文理清 VAE 变分自编码器 | 原理图解+公式推导 https://www.bilibili.com/video/BV1ix4y1x7MR
【VAE变分自编码器原理解析】https://www.bilibili.com/video/BV1op421S7Ep/
【一个视频看懂VAE的原理以及关于latent diffusion的思考】https://www.bilibili.com/video/BV1wx421k74m/
【生成模型VAE】十分钟带你了解变分自编码器及搭建VQ-VAE模型(Pytorch代码)!简单易懂!—GAN/机器学习/监督学习 https://www.bilibili.com/video/BV1Uj411Y7Zq/
【深度学习-自编码器(Auto-Encoders)基本原理及项目实 战[基于PyTorch实现]】 https://www.bilibili.com/video/BV18v41147bT/
博客地址:
苏剑林. (Mar. 18, 2018). 《变分自编码器(一):原来是这么一回事 》[Blog post]. Retrieved from https://kexue.fm/archives/5253
苏剑林. (Mar. 28, 2018). 《变分自编码器(二):从贝叶斯观点出发 》[Blog post]. Retrieved from https://spaces.ac.cn/archives/5343
【VAE学习笔记】全面通透地理解VAE(Variational Auto Encoder) https://blog.csdn.net/a312863063/article/details/87953517
【Variational Autoencoders 】Amaires@May 2024 https://amaires.github.io/VAE/
前言
首次进入生成模型的世界,第一个接触的模型是Diffusion,但是Diffsuion中有VAE的思想,所以正好结合B站视频和博客整体学一学,写一些自己的学习笔记,结合了自己的基础,所以有些地方不会细说。
Auto-Encoder的基本原理和代码实战
VAE中文为变分自编码器,谈到自编码器,离不开谈论Auto-Encoder,也就是AE, 其就是一个神经网络,只不过是分为两个部分,如下图所示。

第一部分为编码器,就是一个降维的过程,由原图像大小到256到64再到20, 这个隐藏层都可以自己设定;第二部分为解码器,就是一个升维的过程,一般都是20到64到256再到图像大小,一般都是对称的,总的来说编码器和解码器就是普通的神经网络层,只不过是中间有个状态表示隐空间,也就是编码器结束的时候,形成的向量就是隐向量,损失函数就是一般回归的L2损失函数或者说MSE。代码如下,一看就懂了。
- model.py
from torch import nn
class AutoEncoder(nn.Module):
def __init__(self, input_size=784, output_size=784, **kwargs):
super(AutoEncoder, self).__init__()
self.input_size = input_size
self.output_size = output_size
self.is_image = kwargs["is_image"]
if self.is_image:
self.image_size = kwargs['image_size']
# [b, 784]
self.encoder = nn.Sequential(
nn.Linear(input_size, 256),
nn.ReLU(),
nn.Linear(256, 64),
nn.ReLU(),
nn.Linear(64, 20) # [b, 20]
)
# [b, 20] => [b, 784]
self.decoder = nn.Sequential(
nn.Linear(20, 64),
nn.ReLU(),
nn.Linear(64, 256),
nn.ReLU(),
nn.Linear(256, output_size),
nn.Sigmoid() # [b, 784]
)
def forward(self, x):
"""
:param x:[b, 1, 28, 28]
:return:
"""
batch_size = x.size(0)
# flatten
x = x.view(batch_size, self.input_size)
# encoder
x = self.encoder(x)
# decoder
x = self.decoder(x)
# reshape
if self.is_image:
x = x.view(batch_size, 1, self.image_size, self.image_size)
return x
- main.py
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import transforms, datasets
from AE import AutoEncoder
import visdom
def main():
mnist_train = datasets.MNIST(root='./mnist', train=True, transform=transforms.ToTensor(), download=True)
mnist_test = datasets.MNIST(root='./mnist', train=False, transform=transforms.ToTensor(), download=True)
mnist_train = DataLoader(dataset=mnist_train, batch_size=32, shuffle=True)
mnist_test = DataLoader(dataset=mnist_test, batch_size=32, shuffle=True)
x, _ = iter(mnist_train).__next__()
print("x:", x.shape) # torch.Size([32, 1, 28, 28])
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = AutoEncoder(
input_size = x.shape[2] * x.shape[2],
output_size = x.shape[2] * x.shape[2],
image_size = x.shape[2],
is_image = True
).to(device)
print(model)
criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
vis = visdom.Visdom()
# 可视化启动服务, python -m visdom.server
for epoch in range(1000):
for batch_idx, (x, y) in enumerate(mnist_train):
# [b,1, 28, 28] => [b, 1*28*28]
x = x.to(device)
x_hat = model(x)
loss = criterion(x_hat, x)
optimizer.zero_grad()
loss.backward()
optimizer.step()
print(epoch, "loss:", loss.item())
x, _ = next(iter(mnist_test))
x = x.to(device)
with torch.no_grad():
x_hat = model(x)
vis.images(x.cpu().numpy(), nrow=8, win='x', opts=dict(title='x'))
vis.images(x_hat.cpu().numpy(), nrow=8, win='x_hat', opts=dict(title='x_hat'))
if __name__ == "__main__":
main()
谈到这里,其实我们也发现了,自编码器好像只能用于降维和特征提取,不能用于生成,因为L2损失好像只能保证数据和原来越来越接近,所以神经网络只能学习到已有的图片的特征。所以更严谨的说法是,AE更适合于重构问题,而不适合生成问题,而这一原因是中间的隐向量的分布没有被约束,这里的分布可能有点不懂了,本人也因为基础知识不牢固,被分布唬住了,这里的分布就是普通的正态分布之类的,但是一定不要混淆的是概率密度曲线(PDF)和概率质量函数(PMF), PMF是描述离散随机变量的, y轴就是x轴样本对应的概率;PDF是描述连续值的,与PMF的y轴表示不同,样本对应的曲线y值就是概率,从负无穷到该点累加的才是x轴样本对应的概率, 也就是P(x)不一定是概率,在连续情形下是概率密度。
VAE的基本原理
VAE的整理训练过程
如果说AE是重构,那么一般的生成模型的本质是构建一个从隐变量\(z\)生成目标数据\(x\)的模型,首先假设\(Z\)服从某种常见的分布(比如正态分布),目的是训练一个模型\(X=g(Z)\),这个模型能够将原来的概率分布映射到训练集的概率分布,最终目的是进行分布之间的变换。
接着由一般生成模型回归到VAE模型,由于AE模型的隐向量的分布没有被约束,为啥不直接约束一下,比如让Z就是服从于正态分布,如下图所示。

上面做法的问题是我们不知道\(Z_i\)和\(X_i\)的对应关系,就无法求损失,那么其实只要一个\(X_i\)求出一个专属的\(\mu_i和\sigma_i\)对应\(Z_i\)就迎刃而解了,如下图所示。

然后到这里就可以说一下VAE的具体架构,如下图所示。


不同于AE, VAE的encoder网络学习的是一个\(p(z\mid x)\)的分布,但是由上图右下角的贝叶斯公式可以知道,这个后验分布很难求出数值解,或者说分母隐空间的\(z\)很难全都找出来,致使\(p(z\mid x)\)很难表示,所以VAE这里利用变分推断的思想,也就是通过一个简单的、可处理的分布(称为变分分布)来近似一个复杂的、难以直接计算的后验分布,就是用\(q_{\phi}(z\mid x)\)近似\(p(z\mid x)\),所以\(q_{\phi}(z\mid x)\)才是encoder真正学出来的分布,这里一般是高斯分布,所以学习的分布更通俗来说,其实编码器学习的是分布的均值和方差。

有了均值和方差,但是没有组合关系,如何构成分布和达到可以支持神经网络反向传播进行求导呢,这里有一个很巧妙的做法,参数重整化或者说重参数技巧,就是\(z=\mu + \sigma \times \epsilon, \epsilon \sim N(0, 1)\),进行参数重整化处理后,该式线性可导,支持神经网络反向传播求导。

有了\(z\sim p(z \mid x)\)的条件,便可以从中采样,这个采样就是通过\(N(0, 1)\)完成的,获得的采样值\(z=\mu + \sigma \times \epsilon\)接着经过解码器生成新的\(X\), 那这里解码器到底是什么,就是对应框架图中的\(p(x \mid z)\), 怎么理解呢,这还是一个分布,不过这个分布的\(\mu\)就是\(f(z_i)\),神经网络得到的值,方差一般是一个固定的常数,这样的话最后神经网络反向传播的损失依然可以是MSE, 也就是L2损失\(\frac{1}{2}(x-\hat{f}(x))^2\)。这就是整体VAE的训练流程。
谈到这里,我猜你们和我一样还是不理解分布到底是啥,为什么约束了隐向量的分布就可以实现生成任务,我结合其他教程的理解是:编码器求均值和方差,均值其实就类似之前AE模型的隐向量,很好理解吧,方差是啥呢,就是一个噪声强度,\(\mu + \sigma \times \varepsilon\),这里由于噪声的存在,可以动态调节重构误差MSE,进而影响到decoder的生成能力,但是就会有一个问题,这里的损失利用了MSE,也就是重构的损失,训练过程中为了重构更好,肯定会想办法让方差为0,这样的话,就没有随机性了,就会退化成普通的AE,噪声就失去了存在的意义,那是如何解决的呢?
VAE的公式推导
为了揭开VAE的神秘面纱,下面将从公式层面解决上面的问题和深入了解VAE的原理,或者说讲解一下总损失目标函数。
这里的推导是基于贝叶斯联合分布进行推理,还有一种方法是基于最大似然估计进行推导,本人觉得第一种方式简单易理解,因此之后的推理都是基于贝叶斯联合分布进行推理。
之前我们提到变分推断的用\(q_{\phi}(z\mid x)\)近似\(p(z\mid x)\), 那如何去衡量分布的相似呢,VAE就采用了KL散度,其定义为:
因此我们的目标是\(\min~ KL\Big(q_{\phi}(z\mid x) \Big\Vert p(z\mid x)\Big)\)
经过如上推导,一般通常把$\mathbb{E}{z\sim q(z \mid x)} \log \left[ \frac{p(z, x)}{q_\phi(z \mid x)} \right] \(称为ELBO(Evidence Lower Bound,证据下界),上式的\)\log p(x)\(其实给定\)x$后,就是一个固定的值,为了最小化KL散度,因此只需要最大化ELBO值。

继续推导,化简ELBO:
由此,ELBO分为了两部分,第一部分是重构项(L2 loss),也就是我们平常说的MSE(\(\dfrac{1}{2}\mid \mid x - \hat{x}\mid \mid ^2\)), 第二项是先验匹配项(KL loss),也就是让编码器和采样的正态分布分布形式更相似。
到了这里,就可以给出保证方差不为0的方法,就是让\(p(z \mid x) \rightarrow \mathcal{N}(0, I)\),又因为以下推导\(p(z)=\mathcal{N}(0, I)\),所以VAE这里做的一个处理是让所有的\(p(z \mid x)\)向标准正态分布看齐,也就是中间的\(P(z)\)向标准正态分布看齐,从而防止零噪声的发生。
我们再来看看仔细的训练过程和采样过程流程图,采样的过程就是随机从标准正态分布中采样,然后输入编码器就可以实现生成。如下图,这里还有个仔细需要注意,我们求\(\sigma\)是借助\(\log \sigma\)辅助求解的,因为方差是恒大于等于0的,但是神经网络的值是实数,所以对方差取一个对数,就可以把值域返回扩展到实数域。

到这里,我们可能还不理解重构项是如何变成L2 loss的,以及先验匹配项的具体形式是咋样的,因为只有知道具体的公式才能构建损失函数进行神经网络训练。
为了将重构项化为L2 loss,需要先引入蒙特卡洛近似的思想,为了近似获得积分的结果,我们可以从\(q(z\mid x)\)中采样具有代表性的n个点,然后代入\(\log P(x\mid z)\),求出均值。
在论文中,其假设 $ P(x \mid z) \sim \mathcal{N}(\mu, \sigma^2 I) $,其中 $ \mu $ 和 $ \sigma^2 $ 都是需要使用神经网络去逼近。
但是,一般地,我们假设 $ P(x \mid z) \sim (f(z), cI) $,也就是其均值用神经网络去逼近,对于其协方差矩阵,我们设定为常数 $ c $ 和 $ I $ 相乘,所以依然是各个维度之间相互独立。我们来看看它的极大似然估计得什么(假设采样 $ n $ 个样本), 具体推导如下。
为了得到先验匹配项也就是KL loss的具体形式,需要引入两个高斯分布之间的 KL 散度公式
假设 \(p(x) = \mathcal{N}(x; \mu_1, \sigma_1^2)\) 和 \(q(x) = \mathcal{N}(x; \mu_2, \sigma_2^2)\),则:
两个高斯分布之间的 KL 散度公式推导
步骤 1:写出高斯分布的概率密度函数\[p(x) = \frac{1}{\sqrt{2\pi}\sigma_1} \exp\left(-\frac{(x-\mu_1)^2}{2\sigma_1^2}\right), \quad q(x) = \frac{1}{\sqrt{2\pi}\sigma_2} \exp\left(-\frac{(x-\mu_2)^2}{2\sigma_2^2}\right) \]步骤 2:代入 KL 散度定义
\[D_{\text{KL}}(p \| q) = \int_{-\infty}^{\infty} p(x) \log \frac{p(x)}{q(x)} dx \]展开对数项:
\[\log \frac{p(x)}{q(x)} = \log \frac{\sigma_2}{\sigma_1} - \frac{(x-\mu_1)^2}{2\sigma_1^2} + \frac{(x-\mu_2)^2}{2\sigma_2^2} \]步骤 3:分离积分项
\[D_{\text{KL}}(p \| q) = \log \frac{\sigma_2}{\sigma_1} \underbrace{\int p(x) dx}_{=1} + \mathbb{E}_p\left[-\frac{(x-\mu_1)^2}{2\sigma_1^2}\right] + \mathbb{E}_p\left[\frac{(x-\mu_2)^2}{2\sigma_2^2}\right] \]步骤 4:计算期望项
- 第一项期望:
\[ \mathbb{E}_p\left[-\frac{(x-\mu_1)^2}{2\sigma_1^2}\right] = -\frac{\mathbb{E}_p[(x-\mu_1)^2]}{2\sigma_1^2} = -\frac{\sigma_1^2}{2\sigma_1^2} = -\frac{1}{2} \]
- 第二项期望:
\[\mathbb{E}_p\left[\frac{(x-\mu_2)^2}{2\sigma_2^2}\right] = \frac{\mathbb{E}_p[(x-\mu_2)^2]}{2\sigma_2^2} \]展开平方项:\[\mathbb{E}_p[(x-\mu_2)^2] = \mathbb{E}_p[(x-\mu_1 + \mu_1 - \mu_2)^2] = \sigma_1^2 + (\mu_1 - \mu_2)^2 \]步骤 5:合并所有项
\[D_{\text{KL}}(p \| q) = \log \frac{\sigma_2}{\sigma_1} - \frac{1}{2} + \frac{\sigma_1^2 + (\mu_1 - \mu_2)^2}{2\sigma_2^2} \]最终公式
\[D_{\text{KL}}(p \| q) = \log \frac{\sigma_2}{\sigma_1} + \frac{\sigma_1^2 + (\mu_1 - \mu_2)^2}{2\sigma_2^2} - \frac{1}{2} \]
根据如上公式,\(KL(q_{\phi}(z \mid x) \| P(z)) = \log{\dfrac{1}{\sigma}} + \dfrac{\sigma^2 + \mu^2}{2} - \dfrac{1}{2}=\dfrac{1}{2}(\mu^2 + \sigma^2 - \log \sigma^2 - 1)\) 。
对于这两个损失,可以分析一下,L2 loss就是重构损失,作用是有区域, KL loss就是先验匹配项,作用是有规则分布,两者同时会发挥有区域有分布的作用,如下图。

最后补充一点,B站up主或者博客博主都认为编码器采样的\(\mu和\sigma\)形成的高斯分布可以叠加成任何分布,如下图所示。

但是我的理解,虽然这里和高斯混合模型的不同高斯分布叠加可以形成任何分布的思想很想,但是这里的实现和叠加还是有一些区别,根据万能近似定理,具有足够容量的神经网络(如解码器)可以逼近任何连续函数。因此,VAE 可以通过潜在空间 z 和解码器的非线性变换来逼近复杂的多模态分布,而不需要显式地叠加多个高斯分布。
VAE的局限性
最后谈一谈VAE的局限,一是针对采样过程,采样是在标准正态分布中随机采样,然后输入到编码器中就可以实现生成,但是VAE优化的目标之一就是最小化\(KL \Big(q_{\phi}(z\mid x) \Big \vert \Big \vert p(z) \Big)\),也就是说真实的\(q_{\phi}(z \mid x)\)和真正的\(\mathcal{N}(0, 1)\)还有一定的差距,因此会造成不匹配问题,后验概率和先验概率并不完全相同。( Latent Diffusion Model(潜在扩散模型), 逐步添加噪声到标准正态分布,然后再逐步去噪声)
二是均方误差(MSE)可能会导致模型优化出一个模糊的图像,比如一张有猫的图,猫耳朵不在正确的位置,但是是存在的,这个图片是不合理的,但是这也是MSE优化的方向,真实性会大打折扣。(Generative Adversarial Network(生成对抗网络),使用判别器代替MSE)
VAE代码
- model.py
import torch
from torch import nn
import numpy as np
class VAE(nn.Module):
def __init__(self, input_size=784, output_size=784, **kwargs):
super(VAE, self).__init__()
self.input_size = input_size
self.output_size = output_size
self.is_image = kwargs["is_image"]
if self.is_image:
self.image_size = kwargs['image_size']
# [b, 784]
self.encoder = nn.Sequential(
nn.Linear(input_size, 256),
nn.ReLU(),
nn.Linear(256, 64),
nn.ReLU(),
nn.Linear(64, 20) # [b, 20]
)
# [b, 10] => [b, 784]
self.decoder = nn.Sequential(
nn.Linear(10, 64),
nn.ReLU(),
nn.Linear(64, 256),
nn.ReLU(),
nn.Linear(256, output_size),
nn.Sigmoid() # [b, 784]
)
def forward(self, x):
"""
:param x:[b, 1, 28, 28]
:return:
"""
batch_size = x.size(0)
# flatten
x = x.view(batch_size, self.input_size)
# encoder
# [b, 20], including mean and sigma
h_ = self.encoder(x)
# [b, 20] => [b, 10] and [b, 10]
mu, sigma = h_.chunk(2, dim=1)
# reparameterization trick, epsilon ~ N(0, 1)
h = mu + sigma * torch.randn_like(sigma)
# \dfrac{1}{2}(\mu^2 + \sigma^2 - \log \sigma^2 - 1)
kld = 0.5 * torch.sum(
torch.pow(mu, 2) +
torch.pow(sigma, 2) -
torch.log(1e-8 + torch.pow(sigma, 2)) - 1
) / (batch_size * self.image_size * self.image_size)
# decoder
x_hat = self.decoder(h)
# reshape
if self.is_image:
x_hat = x_hat.view(batch_size, 1, self.image_size, self.image_size)
return x_hat, kld
- main.py
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from torchvision import transforms, datasets
from VAE import VAE
import visdom
def main_VAE():
mnist_train = datasets.MNIST(root='./mnist', train=True, transform=transforms.ToTensor(), download=True)
mnist_test = datasets.MNIST(root='./mnist', train=False, transform=transforms.ToTensor(), download=True)
mnist_train = DataLoader(dataset=mnist_train, batch_size=32, shuffle=True)
mnist_test = DataLoader(dataset=mnist_test, batch_size=32, shuffle=True)
x, _ = iter(mnist_train).__next__()
print("x:", x.shape) # torch.Size([32, 1, 28, 28])
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = VAE(
input_size = x.shape[2] * x.shape[2],
output_size = x.shape[2] * x.shape[2],
image_size = x.shape[2],
is_image = True
).to(device)
print(model)
criterion = nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
vis = visdom.Visdom()
# 可视化启动服务, python -m visdom.server
for epoch in range(1000):
for batch_idx, (x, y) in enumerate(mnist_train):
# [b,1, 28, 28] => [b, 1*28*28]
x = x.to(device)
x_hat, kld = model(x)
loss = criterion(x_hat, x) # 这里是负对数似然,也就是-mse
if kld is not None:
ELBO = - loss - 1.0 * kld # ELBO = mse - kld
loss = - ELBO # 最终也变成最小化问题,越小越好
optimizer.zero_grad()
loss.backward()
optimizer.step()
print(epoch, "loss:", loss.item(),"kld:", kld.item())
x, _ = next(iter(mnist_test))
x = x.to(device)
with torch.no_grad():
x_hat, kld = model(x)
vis.images(x.cpu().numpy(), nrow=8, win='x', opts=dict(title='x'))
vis.images(x_hat.cpu().numpy(), nrow=8, win='x_hat', opts=dict(title='x_hat'))
if __name__ == "__main__":
main_VAE()

浙公网安备 33010602011771号