使用变分自编码器生成MNIST数字图像
一、AE(AutoEncoder)-自编码器
自编码器是一种无监督学习模型,其核心思想是:将输入样本 x 经过编码器(Encoder)压缩到一个潜在空间(Latent Space),再通过解码器(Decoder)将其还原为输出 x 的预测。
然而,传统自编码器存在一个问题:其潜在空间通常是离散且不连续的。这意味着,如果我们在潜在空间中随机采样一个点并输入给解码器,往往会得到无意义的结果——因为这些点可能并不对应任何真实样本的概率分布区域,也就是说,它们“落在空白区”,无法生成合理的样本。
AE到VAE的过渡——为什么AE不能生成?
要理解VAE为什么能生成,我们首先要弄清楚:为什么普通的AE不行?
想象你养了两只猫,一只橘猫和一只黑猫。你把橘猫的照片喂给AE,编码器把它压缩成一个隐向量 z₁;再把黑猫的照片喂进去,得到隐向量 z₂。现在有个天真的想法:如果在 z₁ 和 z₂ 之间取一个中点 z_mid = (z₁+z₂)/2,让解码器还原——会不会得到一只"介于橘猫和黑猫之间"的花猫?
很遗憾,你得到的极大概率是一团毫无意义的噪点。
为什么? 因为AE的编码器是一个确定性映射——它只学会了把训练集中出现过的每个样本映射到隐空间中的某个点。z₁ 和 z₂ 各自对应一个已知样本,但这两个点之间的"中间地带"是模型从未见过的区域。解码器在这个区域的表现完全不可预测——它没有被训练过去"理解"这些位置应该还原出什么。
换句话说,AE的隐空间是"离散的点映射":有样本的地方是绿洲,样本之间的地方是荒漠。你想在荒漠里随机打井取水——大概率是干枯的。
VAE的解法,是给隐空间"铺上草地"。 VAE不再把每个样本编码成隐空间中的一个点,而是把它编码成一个高斯分布——用均值 μ 标记分布的中心位置,用方差 σ² 描述分布的扩散范围。这样,橘猫对应的就不再是孤零零的一个点 z₁,而是以 μ₁ 为中心、以 σ₁ 为半径的一个"概率云";黑猫对应的也是以 μ₂ 为中心、以 σ₂ 为半径的一个"概率云"。
当这两个"概率云"足够大、彼此靠近时,它们就会产生重叠。重叠区域中的点既"有点像橘猫"也"有点像黑猫"——解码器见到这些点,就会还原出一只"合理过渡"的猫,而不是一团噪声。
这正是VAE能够随机采样并生成新图片的根本原因:整个隐空间被约束成连续、平滑的"宜居带",你在里面随便采一个点,解码器都能给你一个合理的图像。你不再需要依赖训练集里的已知样本——你拥有了真正的"生成"能力。
二、VAE-变分自编码器 为了解决潜在空间不连续的问题,VAE在自编码器的基础上引入了概率建模思想。它不再直接将输入压缩成一个固定向量,而是将每个样本编码为一个高斯分布,用均值向量 μ和方差向量 σ²来表示。每次输入一个样本,都会通过Encoder将其变为 μ 和 σ
在训练过程中,VAE通过优化目标函数,使这些高斯分布在潜在空间中尽可能连续且符合标准正态分布 N(0,1)。这样,在生成阶段,我们就可以直接从 N(0,1)中采样潜在变量 z,再通过解码器生成新的样本。
可以说,VAE通过“给潜在空间加上分布约束”,让模型具备了真正的生成能力。
重参数化技巧:VAE最精巧的设计
读到这里,你可能已经注意到一个关键的问题:VAE需要从高斯分布 N(μ,σ²) 中采样一个 z,然后喂给解码器。但采样操作本身是随机的、不可导的——而深度学习依赖梯度反向传播来更新参数。如果梯度传不到 μ 和 σ,网络就没法学。
这就是VAE论文中最精巧的设计——重参数化技巧(Reparameterization Trick)。
核心思路很简单:把随机性"固定住",不让它参与梯度计算。
具体做法是:我们不直接从 N(μ,σ²) 采样,而是先从标准正态分布 N(0,1) 中采样一个随机数 ε,然后通过一个确定性的变换得到 z:
[ z = \mu + \varepsilon \cdot \sigma,\quad \varepsilon \sim \mathcal{N}(0, 1) ]
在这个公式里,ε 的采样确实发生了,但它在计算图中被视为一个常数输入——就像你把一张图片输入网络一样,你不会对图片本身求导。梯度需要穿过的路径是:z → μ 和 z → σ,这两条路径都是确定性的算术运算(加法和乘法),完全可导。
打个比方: 掷骰子本身是随机的——你无法预测下一次会掷出几点。但你可以决定掷什么样的骰子:是六面骰(方差小,结果在1-6之间)、还是二十面骰(方差大,结果范围更广)。你还可以决定期望掷出什么数字:把骰子的重心偏移一下(通过 μ 调整)。μ 和 σ 就是网络可以学习的"骰子属性",而 ε 就是"命运的随手一扔"——命运不由你控制,但骰子的类型和偏好你可以学。
在本文代码中,你会在 reparameterize 方法里看到这个技巧的精确实现:
std = torch.exp(0.5 * logvar) # σ = exp(0.5 * log(σ²))
eps = torch.randn_like(std) # ε ~ N(0,1),不参与梯度
return mu + eps * std # z = μ + ε·σ,梯度畅通无阻
如果没有这个trick,VAE就无法训练。 因为采样操作的梯度为0,编码器永远收不到"我的 μ 和 σ 应该怎么调整"的信号。重参数化技巧把"不可导的随机采样"变成了"可导的确定性变换 + 固定噪声输入",让梯度可以顺畅地从重构误差一路传回编码器——这是VAE论文在2014年发表时最核心的技术贡献。
三、实例:MNIST手写数字的图像生成
1、实际表现
训练Epoch = 50,进行输出


效果只能说可以看,但是还需要优化。通过VAE生成的图像有个很明显的特点,就是中间清晰,四周模糊。以下是造成周围模糊的三个主要原因:
- 重构目标的均方误差(MSE)导致的模糊性
VAE通常采用均方误差作为重构损失,这相当于假设像素服从独立的高斯分布。模型在优化时会趋向生成“平均意义上的正确像素值”,从而在多个可能结果之间取均值,导致输出图像整体变得平滑、细节模糊。 - 潜在空间采样的不确定性
由于VAE在潜在空间中对每个输入都引入了随机采样(即 z∼N(μ,σ2)),生成时的这种随机性会带来细节的模糊扩散,尤其在解码器的非线性映射不够强时更为明显。 - 潜在空间的全局性编码倾向
编码器通常优先捕捉输入图像的整体结构信息,而对局部高频细节(如边缘、纹理)保留不足。因此生成的图像中心部分(结构清晰的区域)通常重构较好,而边缘部分往往较为模糊。
2、代码
class VAE(nn.Module):
def __init__(self, latent_dim):
super(VAE, self).__init__()
self.encoder = nn.Sequential(
nn.Conv2d(1, 32, kernel_size=3, stride=2, padding=1), # 28x28 -> 14x14
nn.ReLU(),
nn.Conv2d(32, 64, kernel_size=3, stride=2, padding=1), # 14x14 -> 7x7
nn.ReLU(),
nn.Flatten(),
nn.Linear(64 * 7 * 7, 128),
nn.ReLU()
)
self.fc_mu = nn.Linear(128, latent_dim)
self.fc_logvar = nn.Linear(128, latent_dim)
self.decoder = nn.Sequential(
nn.Linear(latent_dim, 64 * 7 * 7),
nn.ReLU(),
nn.Unflatten(1, (64, 7, 7)),
nn.ConvTranspose2d(64, 32, kernel_size=3, stride=2, padding=1, output_padding=1), # 7x7 -> 14x14
nn.ReLU(),
nn.ConvTranspose2d(32, 1, kernel_size=3, stride=2, padding=1, output_padding=1), # 14x14 -> 28x28
nn.Sigmoid()
)
def encode(self, x):
h = self.encoder(x)
mu = self.fc_mu(h)
logvar = self.fc_logvar(h)
return mu, logvar
def reparameterize(self, mu, logvar):
std = torch.exp(0.5 * logvar)
eps = torch.randn_like(std)
return mu + eps * std
def decode(self, z):
return self.decoder(z)
def forward(self, x):
mu, logvar = self.encode(x)
z = self.reparameterize(mu, logvar)
recon_x = self.decode(z)
return recon_x, mu, logvar
def vae_loss(recon_x, x, mu, logvar):
recon_loss = nn.functional.mse_loss(recon_x.view(-1, 784), x.view(-1, 784), reduction='sum')
kl_loss = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
return recon_loss + kl_loss
KL散度的双重作用
在 vae_loss 函数中,你看到了这样一个损失项:
kl_loss = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
这一项被称为KL散度(Kullback-Leibler Divergence),它衡量的是"每个样本的高斯分布 N(μ,σ²) 与标准正态分布 N(0,1) 之间有多大的差异"。它在VAE的训练中扮演着两个至关重要的角色:
作用一:让隐空间连续、有重叠。
如果没有KL项,网络会把每个样本的 μ 学得很大(让不同样本在隐空间中"各占一块地盘"),同时把 σ 学得非常小(退化成普通AE的一个点)。这样的隐空间就是前面说的"荒漠地带"——样本之间没有重叠,随机采样毫无意义。
KL项强制把每个样本的高斯分布向 N(0,1) 靠拢。N(0,1) 是标准正态分布,它的均值为0、方差为1——这意味着所有样本的分布都会被拉到隐空间的中心附近,并且有合理的扩散范围。不同样本的分布因此会产生大量重叠,而重叠的区域恰恰是"合理的插值"——解码器在这些区域生成的结果会自然地从一种数字过渡到另一种数字。
可以这样理解:有KL项时,隐空间像一个平滑的圆盘,盘上任意一点采样都有意义;没有KL项时,隐空间像海面上散布的孤岛,只有站在岛上(训练样本点)才能得到合理结果,跳进海里就是噪声。
作用二:正则化,防止退化和过拟合。
KL损失项的数学形式揭示了它如何实现这种约束:
[ KL = -\frac{1}{2}\sum_{i}\left(1 + \log\sigma_i^2 - \mu_i^2 - \sigma_i^2\right) ]
其中:
- μ² 项:惩罚过大的 |μ|,把每个分布的均值拉向0。如果某个样本的 μ 飞到100,这项损失会非常大。
- σ² 项:惩罚过大的 σ²,防止分布"散得太开"覆盖整个空间——那会让所有样本混成一团,失去区分能力。
- log(σ²) 项:惩罚过小的 σ²。当 σ²→0 时,log(σ²)→-∞,KL损失急剧增大——这阻止了VAE退化成普通AE(σ²=0时,z=μ,没有随机性)。
- +1 项:一个常数基准,保证KL散度始终非负。
三股力量——μ²、σ²、log(σ²)——在训练中相互博弈,最终达到一个优雅的平衡:每个样本的分布既不会太散(保持区分度),也不会太紧(保持连续性),整个隐空间被规整成一个大约以原点为中心、半径为1的"紧凑球体"。在这个球体里,任意一个点都对应一个合理的MNIST数字图像。
VAE vs GAN:生成模型的两条路径
本文生成的MNIST数字图像有一个很明显的特征——数字的形状是对的,但边缘模糊、细节不清晰。这是VAE的"通病"而非本文实现的bug。
VAE生成的典型特点:整体结构正确,细节模糊。 为什么会这样?根源在于VAE使用MSE(均方误差)作为重构损失。MSE本质上是在"取平均"——当模型不确定某个像素应该是黑色还是白色时,它会输出灰色。这种平均化倾向保证了图像在全局结构上的正确性(数字的形状不会错),但代价是所有细节都被"抹平"了。
与之形成鲜明对比的是GAN(生成对抗网络)。GAN不用MSE,而是用一个判别器来评判生成图像是否"逼真"。生成器的目标不是"尽量接近原始图像",而是"骗过判别器"。这种博弈机制让GAN生成的图像细节锐利、以假乱真——但它也有自己的致命弱点:模式坍塌(Mode Collapse)。模式坍塌意味着生成器学会了"作弊"——它只生成少数几种看起来特别逼真的样本,而忽略了数据集中其他类型的样本。比如在MNIST上,一个模式坍塌的GAN可能只会生成"0"和"1",永远不会生成"3"或"8"。
下表总结了VAE和GAN的核心差异:
| 维度 | VAE | GAN |
|---|---|---|
| 生成质量 | 整体正确,细节模糊 | 细节锐利,以假乱真 |
| 模式覆盖 | 覆盖所有类别 | 容易模式坍塌 |
| 训练稳定性 | 稳定,loss平稳下降 | 不稳定,需精心调参 |
| 隐空间结构 | 连续、平滑、可插值 | 无显式概率结构 |
| 理论基础 | 变分推断,数学优雅 | 博弈论,工程驱动 |
有趣的是,VAE和GAN恰好是互补的——一个覆盖全但模糊,一个清晰但可能缺失。这正是后续研究者提出VAE-GAN等混合模型的动机:用VAE的编码器-解码器结构保证模式覆盖,用GAN的判别器来锐化生成结果,取两家之长。
所以,当你看到本文生成的数字图像"中间清晰、四周模糊"时——这不是失败,而恰恰是VAE本质特征的体现。作者在第一节中分析的三个模糊原因(MSE取平均、采样随机性、全局编码倾向),共同构成了VAE的"基因":它天然倾向于生成"安全但模糊"的结果。理解这一点,比得到一个完美的生成结果更有价值。
四、扩展应用
在使用的时候,一般都是以VAE作为编码器使用,通过输入样本,将样本压缩至潜在空间,然后进行后续处理;处理后再通过冻结的预训练VAE解码器进行解码输出。除了这种基础用法,VAE及其变体在实际中还有以下重要应用场景:
1. 数据压缩
VAE的编码器天然就是一个压缩器。以MNIST为例,输入是一张28×28=784维的图像,经过编码器后压缩到 latent_dim(本文中通常设为20或更小)——压缩比接近40:1。更重要的是,解码器能从这20个数字中重建出可辨认的数字图像,说明这20维抓住了数据最本质的结构信息。这比PCA等传统线性降维方法更强大,因为VAE学习的是非线性的、语义上有意义的压缩表示。
2. 异常检测
这是VAE在工业界非常实用的一个应用。思路很简单:在一个"正常"的数据集上训练VAE(比如正常运行的机器传感器数据),让模型学会"什么是正常"。当一个新的样本到来时:
- 如果它是正常的 → 编码再解码的重构误差很小(模型"见过"类似的模式)
- 如果它是异常的 → 重构误差很大(模型"不知道"怎么重建这个没见过的东西)
通过设定一个重构误差的阈值,就可以实现自动异常检测。这被广泛应用于工业设备故障预警、金融欺诈检测、医疗影像异常筛查等领域。
3. 数据增强
训练数据不够?VAE可以从隐空间中采样生成新样本。因为隐空间是连续且平滑的,在标准正态分布 N(0,1) 中随机采样一个 z,解码后就能得到一张"合理的新图片"。这些生成样本虽然不是真实的拍照数据,但它们丰富了训练集的多样性,可以作为数据增强的一种手段,提升下游分类器的泛化能力。
4. 特征解耦与可控生成
标准VAE的隐空间虽然是连续的,但每个维度代表什么语义并不明确——latent_dim=20中的第3维到底控制"数字的粗细"还是"数字的倾斜角度"?答案是不确定的,所有维度混在一起。
β-VAE 是VAE的一个重要变体,它在KL损失项前乘以一个系数β(β>1)。增大β会迫使隐空间的每个维度尽可能独立、每个维度只编码一个独立的"变化因素"。在实践中,β-VAE可以实现令人惊叹的效果:比如在MNIST上,某些维度单独控制"数字是几"(类别),另一些维度单独控制"笔画粗细"(风格),还有一些维度控制"倾斜角度"——实现了解耦表示学习。
这种解耦能力让"可控生成"成为可能:固定类别维度、滑动风格维度,就能生成"不同粗细的同一个数字"——这在创意设计、图像编辑等场景中有巨大的应用潜力。