变分自编码器 VAE

相关:变分下界ELBO笔记(ELBO 推导与直觉) | 去噪扩散概率模型DDPM笔记(多步潜变量模型)

从 ELBO 到 VAE

回顾 ELBO 恒等式(见 KL 散度分解):

VAE 把两个分布都交给神经网络参数化:

  • 编码器(推断网络):输入 ,输出变分后验的参数;
  • 解码器(生成网络):输入 ,输出重构 的参数。

两者取最常用的高斯形式:, 由解码器输出决定。下面逐行对照 main.py 讲。

把 q_θ(z) 换成 q_φ(z | x),等式还成立吗

成立。回看 ELBO 推导(KL 散度分解),恒等式

从头到尾没有用到 的具体形式,对任意变分分布都成立; 不过是”对当前观测 选定一个分布”。逐处替换 即可:

推导中唯一的技巧仍然只是” 不含 、可提出期望”,与 的形式无关。

语义变化:原来一个 服务所有 (换观测要重新优化参数),现在网络 对每个 直接输出专属的变分分布。训练目标是整个数据集上各样本 ELBO 的平均:

采样时只需把重参数化中的 、 换成关于 的函数:。

从联合概率到两项分解:贝叶斯公式的用法

上面的 ELBO 以联合概率 出现;用乘法公式把它拆成”似然 × 先验”:

代入后,两个期望各自归位,联合概率就转化成了两项 KL 的形式:

最后一步用到 KL 的定义 ,因此 。

若把联合概率往另一方向拆(贝叶斯公式):,则得到恒等式的另一半:

两条路径拼起来正是 ELBO 恒等式:

编码器:输出 μ 与 log σ²

网络结构

self.encoder = nn.Sequential(
    nn.Linear(784, 256), nn.ReLU(),
    nn.Linear(256, 64), nn.ReLU(),
    nn.Linear(64, 20),
)
  • 输入 784 维 = 28 × 28 像素展平;
  • 最后一层输出 20 维,前 10 维是 ,后 10 维是 :
mu, log_var = hidden.chunk(2, dim=1)

潜在空间只有 10 维,远小于输入 784 维,这是信息瓶颈: 必须压缩出最关键的特征,才能重构好图像。

为什么存 log σ² 而不是 σ²

方差必须非负,直接输出 需要额外约束;而 可以取任意实数,网络自由输出,使用时再指数还原:

数值上也更稳定:极小的方差在 空间不会下溢成负值。

重参数化:让”采样”变得可导

训练需要关于 、 的梯度,但直接采样不可导:

重参数化把采样拆成”确定性变换 + 标准噪声”:

随机性全部由与参数无关的 承担,梯度可以经 、 正常回传(,)。

std = torch.exp(0.5 * log_var)
eps = torch.randn_like(std)
z = mu + eps * std

解码器:从 z 重构图像

self.decoder = nn.Sequential(
    nn.Linear(10, 64), nn.ReLU(),
    nn.Linear(64, 256), nn.ReLU(),
    nn.Linear(256, 784), nn.Sigmoid(),
)
  • 输入是 10 维隐变量 ;
  • 输出 784 维后接 Sigmoid,把值压到 ,对应像素灰度(数据归一化后也在 )。

解码器输出的是 的参数:把每个像素 视为独立的 Bernoulli 分布,

就是”从 重构第 个像素的概率”,这正是 ELBO 笔记里”重构项”的实现。

重构损失:二元交叉熵(BCE)

把全部像素的对数似然加起来:

训练取负号(最大化似然 ⟺ 最小化负对数似然):

criterion = nn.BCELoss(reduction="sum")
recon_loss = criterion(x_hat, data)

与 逐像素对应:BCE 大 ⟺ 重构图与输入图差异大。这就是 的蒙特卡洛估计(每步只抽一个样本 )。

KL 项:闭式推导

标准先验 与高斯后验 的 KL 有解析式,无需采样:

KL = -0.5 * torch.sum(1 + log_var - mu.pow(2) - log_var.exp())

推导过程

对一维 、,由定义

第一项( 的负熵,用到 ):

第二项(用到 ):

相减得

写成代码的形式( 就是 log_var, 就是 log_var.exp()):

10 个维度各自算完再求和,就是那一行 KL。逐项看它如何起作用:

  • 若编码器输出 、,KL = 0, 恰好等于先验;
  • 或 都会使 KL > 0,产生”偏离先验”的惩罚。

与重构损失的尺度对齐

KL = KL / (x.size(0) * 28 * 28)
recon_loss = recon_loss / (data.size(0) * 28 * 28)
  • BCE 用 reduction="sum" 按 batch 求和;
  • 两者再除以 batch × 784,统一成”每个像素的平均损失”再相加,避免 KL 因维度求和而压过重构项。

训练与生成

训练循环

loss = recon_loss + kl
loss.backward()
optimizer.step()

最小化 (重构误差 + KL 正则),即把 ELBO 笔记里的下界往上抬:重构项让 选能还原 的 ,正则项让 靠近标准正态。

生成新样本

训练完成后不再需要编码器,直接从先验采样、解码:

sample_z = torch.randn(16, 10)
generated = model.decoder(sample_z)

数学上这就是用蒙特卡洛估计边缘分布

KL 项保证了编码器把训练数据映射到标准正态附近,所以从先验随机采 也能落在数据的分布区域,生成的图像才像 MNIST。

与 ELBO 笔记的对应

VAE 组件代码位置对应 ELBO 概念
编码器 encoder + chunk变分分布 (推广为依赖 )
解码器 decoder + Sigmoid
重构损失 BCEnn.BCELoss
KL 正则KL 一行
重参数化z = mu + eps * std让 ELBO 的梯度能回传