Auto-Encoding Variational Bayes
路线图
想让模型分布接近真实数据分布
↓
最大化训练图片在模型下的对数似然
↓
用“先采样潜变量 z,再生成图片 x”的方式建模
↓
计算 pθ(x) 时必须对 z 积分,难以直接优化
↓
直接从先验 p(z) 蒙特卡洛可估计 pθ(x),但有限采样下 log 估计有偏且常与当前 x 错配
↓
引入编码器 qφ(z|x) 近似不可得的真实后验
↓
得到 log pθ(x) 的可优化下界 ELBO
↓
重建期望从 qφ(z|x) 采样估计,并用重参数化回传梯度;KL 项有闭式解
↓
最小化负 ELBO(MSE/BCE 重建项 + KL 项)的 batch 平均
↓
训练后从先验 p(z) 采样,再由解码器生成新图片
0. 记号:一张图、一个 batch、多个 z
- 训练集:D={x(n)}n=1N,其中 x(n) 是第 n 张真实训练图片。
- x:后续推导中,为简洁起见,固定的一张训练图片 x=x(n);它不是模型输出,也不是整个 batch。
- x^:解码器针对 z 生成的重建结果(或输出分布的均值/参数)。
- z:连续潜变量,表示影响图像的隐藏因素,如姿态、光照、物体属性或高层语义。
- p(z):潜变量先验,通常设为 N(0,I)。
- pθ(x∣z):解码器/生成模型;给定 z 后生成图片 x 的条件分布。
- qϕ(z∣x):编码器/近似后验;给定图片 x 后预测其潜变量的分布。
实际训练有两个独立的数量:
| 符号 | 表示什么 | 常见取值 |
|---|
| B | 一个 batch 的图片数:x(1),…,x(B) | 32、64、128 等 |
| K | 对每一张图片采样的潜变量数 | 通常为 1 |
完整的采样记号是:
z(b,k)∼qϕ(z∣x(b)),b=1,…,B,k=1,…,K.
下文先对单张固定图片 x 推导,最后再回到 batch。
1. 为什么目标是最大化数据似然
对于图片空间 我们可以分为两部分 一种是我们希望能从模型中预测出来的图片,他们应该是 合理的/有意义的/特定领域 图片,而不是 随机/非自然/领域外 的像素组合
从训练数据(合理的/有意义的/特定领域的)得到的数据分布 pdata(x);生成模型定义一个可学习的分布 pθ(x)。我们希望模型分布接近数据分布,一种具有明确统计意义的选择是最小化:
DKL(pdata∥pθ)=不随 θ 变化Epdata[logpdata(x)]−Epdata[logpθ(x)].
第一项与参数 θ 无关,因此最小化该 KL 等价于最大化:
θmaxEx∼pdata[logpθ(x)].
我们不知道 pdata,但有从中采样得到的训练集,因此用样本平均近似:
θmaxN1n=1∑Nlogpθ(x(n)).
这就是最大似然的动机:让模型分布接近数据分布。
2. 用潜变量建模图片
直接描述复杂图片分布 pθ(x) 很困难。VAE 采用生成假设:
先从简单分布中抽取隐藏因素 z,再根据 z 生成图片 x。
z∼p(z)⟶x∼pθ(x∣z).
它给出联合分布:
pθ(x,z)=pθ(x∣z)p(z).
其中 p(z) 是我们预先指定、且容易采样的潜变量分布,通常取标准高斯 N(0,I)。它不是由某张图片 x 决定的:生成新图片时,先从这个固定分布抽取 z,再送入解码器。选择简单先验的目的,是让“生成时该从哪里抽取 z”成为一个明确且容易执行的问题。
图片本身的概率必须将潜变量边缘化:
pθ(x)=∫pθ(x,z)dz=∫pθ(x∣z)p(z)dz.
3. 难点:边缘似然与真实后验都难算
对神经网络解码器,上述积分通常没有闭式解,因而无法直接精确优化 logpθ(x)。
从先验 z(k)∼p(z) 做 Monte Carlo 可以估计概率本身:
pθ(x)=Ep(z)[pθ(x∣z)]≈K1k=1∑Kpθ(x∣z(k)).
但训练目标是 logpθ(x)。对样本均值取对数不是 logpθ(x) 的无偏估计;由 Jensen 不等式
≤logE[p^K(x)]=logpθ(x).
它在有限 K 时实际上给出了一个偏低的下界。更关键的是先验 p(z) 不知道当前图片 x,大多数先验样本无法解释它,只有极少数样本有显著权重。
核心澄清:问题不只是“有偏与无偏”
直接先验采样要处理的是:
logEp(z)[pθ(x∣z)].先用样本均值估计内部期望,再取对数,得到的 logp^θ(x) 在有限 K 下是对 logpθ(x) 的向下有偏估计。它并非理论上不能使用:K→∞ 时会收敛;真正的实践困难还包括先验与后验的错配——对固定 x,能解释它的区域 pθ(z∣x) 往往只占 p(z) 的很小部分,因此有限预算下可能几乎抽不到有效样本。
而 VAE 重建项处理的是另一个目标:
Eqϕ(z∣x)[logpθ(x∣z)].对这个期望直接取样本均值是无偏的。但这不表示 VAE 得到了 logpθ(x) 的无偏估计;VAE 改为优化 ELBO:
LELBO(x)≤logpθ(x).ELBO 对真实对数似然而言是下界(偏低),但其组成部分能稳定地计算或估计。qϕ(z∣x) 的作用是使有限采样更有效;KL 项则约束它不能只为重建而任意偏离先验。
我们真正想知道的是:给定图片 x,哪些 z 更可能生成它?这对应真实后验:
pθ(z∣x)=pθ(x)pθ(x,z)=∫pθ(x∣z)p(z)dzpθ(x∣z)p(z).
其分母正是难算的 pθ(x),所以真实后验也不可直接得到。
4. 引入编码器 qϕ(z∣x)
VAE 用一个神经网络 qϕ(z∣x) 近似真实后验 pθ(z∣x)。常用对角高斯:
qϕ(z∣x)=N(μϕ(x),diag(σϕ2(x))).
编码器输入当前图片 x,输出 μϕ(x) 与 logσϕ2(x)(logvar)。和不看 x 的先验相比,qϕ(z∣x) 会把采样集中到更可能解释当前图片的区域。
5. 从后验近似得到 ELBO
考虑近似后验与真实后验的 KL 散度:
DKL(qϕ(z∣x)∥pθ(z∣x))=Eqϕ(z∣x)[logpθ(z∣x)qϕ(z∣x)]=Eqϕ(z∣x)[logpθ(x,z)/pθ(x)qϕ(z∣x)]=logpθ(x)−Eqϕ(z∣x)[logqϕ(z∣x)pθ(x,z)].
最后一步中 logpθ(x) 与积分变量 z 无关,因此能从期望中移出。定义:
LELBO(x)=Eqϕ(z∣x)[logqϕ(z∣x)pθ(x,z)].
便有精确恒等式:
logpθ(x)=证据下界LELBO(x)+≥0DKL(qϕ(z∣x)∥pθ(z∣x)).
因此 LELBO(x)≤logpθ(x)。最大化 ELBO 会提高似然下界;若 qϕ(z∣x) 恰好等于真实后验,二者相等。
6. ELBO 为什么是“重建 − KL”
代入 pθ(x,z)=pθ(x∣z)p(z):
LELBO(x)=Eqϕ(z∣x)[logqϕ(z∣x)pθ(x∣z)p(z)]=Eqϕ(z∣x)[logpθ(x∣z)]+Eqϕ(z∣x)[logqϕ(z∣x)p(z)]=重建项Eqϕ(z∣x)[logpθ(x∣z)]−先验匹配项DKL(qϕ(z∣x)∥p(z)).
| 项 | 约束什么 | 为什么需要它 |
|---|
| 重建项 Eq[logpθ(x∣z)] | 从 z 能否还原当前 x | 让 z 保留图片信息 |
| KL 项 DKL(qϕ(z∣x)∥p(z)) | 编码出的分布是否接近固定先验 | 保证生成时从 p(z) 采样仍有效 |
这两个目标缺一不可:
- 只有重建项:编码器只需为每张图找到任何能让解码器复原它的 z。不同图片的编码区域可以相互零散、没有统一形状;解码器只见过这些区域,而生成时却要从 p(z)=N(0,I) 抽样,因此抽到的 z 可能落在从未训练过的空洞中。这就是普通 AE 虽能重建、却不一定能随机生成的原因。
- 只有 KL 项:最容易的解是对所有图片都令 qϕ(z∣x)=p(z)。此时编码器输出不再依赖 x,z 不携带图片信息;又没有重建项训练解码器保留 x,模型自然无法完成重建。
先验的选择、网络结构、近似分布是人为设计;选定这套模型和下界后,“重建减 KL”是数学结果。“重建”和“先验匹配”则是我们对这个结果的直观解释。
7. 为什么这个目标可训练
ELBO 仍含 z,但它的两部分可以实际计算或估计。
7.1 重建项:从 qϕ(z∣x) 采样
对同一张固定图片 x,独立采样 K 个潜变量:
z(k)∼qϕ(z∣x),k=1,…,K.
则:
Eqϕ(z∣x)[logpθ(x∣z)]≈K1k=1∑Klogpθ(x∣z(k)).
当 z(k) 独立地从 qϕ(z∣x) 采样时,样本均值是该期望的无偏估计。每个 z(k) 都参与求和,但不代表它们的数值贡献相同。相比从先验采样,qϕ(z∣x) 的样本与当前 x 更相关,通常具有更低方差。实践中常取 K=1。
例子:为什么从 p(z) 直接估计似然很困难,而从 qϕ(z∣x) 估计重建项可行?
先用一个极端但直观的离散例子。令潜变量只能取 100 个值:
z∈{1,2,…,100},p(z=j)=1001.固定一张训练图片 x。假设只有 z=17 能生成它:
pθ(x∣z=17)=1,pθ(x∣z=17)=0.那么真实边缘似然是:
pθ(x)=z∑pθ(x∣z)p(z)=1×1001=0.01.若直接从先验 p(z) 抽 10 次,抽到 z=17 的概率只有:
1−(10099)10≈9.6%.也就是说约 90.4% 的时候,10 个样本都没有任何贡献,得到的 p^θ(x)=0,再取 logp^θ(x) 会变成 −∞;即使真实 logpθ(x)=log0.01 是有限的。这说明问题不只是“有一点方差”,而是采样预算会浪费在大量与当前 x 无关的 z 上。
真实后验却是:
pθ(z∣x)={1,0,z=17,z=17.如果编码器已经学到近似后验 qϕ(z∣x)≈pθ(z∣x),它会几乎总是采到 z=17。于是重建项
Eqϕ(z∣x)[logpθ(x∣z)]可由极少样本稳定估计。
7.2 重参数化:让重建误差能够回传
直接从依赖参数的分布采样,不便向编码器参数 ϕ 反向传播。改为采样与参数无关的噪声:
ϵ∼N(0,I),z=μ+σ⊙ϵ=μ+exp(21logσ2)⊙ϵ.
给定 ϵ 后,z 是 μ,σ 的可微函数,梯度即可经由 z 回传到编码器。
def reparameterize(mu, logvar):
std = torch.exp(0.5 * logvar)
eps = torch.randn_like(std)
return mu + eps * std
8. 最终训练损失:重建项与 KL 项
对单张训练图片,最终要最小化的是负 ELBO:
LVAE(x)=重建项−Eqϕ(z∣x)[logpθ(x∣z)]+KL 项DKL(qϕ(z∣x)∥p(z)).
实际训练时,对 batch 中每张 x(b) 的 K 个潜变量样本求平均,再对 B 张图片求平均:
Lbatch=B1b=1∑B[−K1k=1∑Klogpθ(x(b)∣z(b,k))+DKL(qϕ(z∣x(b))∥p(z))].
下面两个子节分别将这两项写成实际可计算的形式。
8.1 KL 项:高斯时可逐项解析计算
重建项需要采样估计,但 KL 项在常用的高斯设定下不需要采样。对当前固定图片 x,令潜变量维度为 d:
qϕ(z∣x)=N(z;μ,diag(σ2)),p(z)=N(z;0,I).
这里 μ,σ 都是编码器针对当前 x 输出的 d 维向量;μj,σj 表示第 j 个维度。由于对角高斯的各维相互独立,KL 可以逐维相加。
从定义开始:
DKL(qϕ(z∣x)∥p(z))=∫qϕ(z∣x)[logqϕ(z∣x)−logp(z)]dz=Eqϕ[logqϕ(z∣x)]−Eqϕ[logp(z)].
先计算先验的对数概率。标准高斯满足:
logp(z)=−2dlog(2π)−21j=1∑dzj2.
而若 zj∼N(μj,σj2),则 Eq[zj2]=μj2+σj2。因此:
Eqϕ[logp(z)]=−2dlog(2π)−21j=1∑dEq[zj2]=−2dlog(2π)−21j=1∑d(μj2+σj2).
展开:标准高斯的对数概率与二阶矩从何而来
因为先验是 d 维标准正态 p(z)=N(z;0,I),各维独立:
p(z)=j=1∏d2π1exp(−2zj2)=(2π)−d/2exp(−21j=1∑dzj2).对上式取对数;乘积变为求和,指数项的对数就是指数:
logp(z)=−2dlog(2π)−21j=1∑dzj2.另一方面,zj∼N(μj,σj2) 的均值和方差定义为:
Eq[zj]=μj,Varq(zj)=Eq[(zj−μj)2]=σj2.将平方展开:
σj2=Eq[zj2−2μjzj+μj2]=Eq[zj2]−2μjEq[zj]+μj2=Eq[zj2]−μj2.移项即可得到:
Eq[zj2]=μj2+σj2.将 logp(z) 代入期望,并利用期望的线性性,就得到正文中的 Eq[logp(z)]。
再计算近似后验自身的对数概率。对角高斯满足:
logqϕ(z∣x)=−2dlog(2π)−21j=1∑d[logσj2+σj2(zj−μj)2].
由于 Eq[(zj−μj)2]=σj2,有:
Eqϕ[logqϕ(z∣x)]=−2dlog(2π)−21j=1∑d(logσj2+1).
两式相减,−2dlog(2π) 抵消:
DKL(qϕ(z∣x)∥p(z))=21j=1∑d(μj2+σj2−1−logσj2).
论文截图中写的是 ELBO 内的 −DKL,故符号相反:
−DKL(qϕ(z∣x)∥p(z))=21j=1∑d(1+logσj2−μj2−σj2).
若编码器输出 logvar = \log\sigma^2,实现为:
kl = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
8.2 重建项:固定方差高斯时对应 MSE
若 pθ(x∣z) 是固定方差高斯,重建项的负对数似然(忽略常数与缩放)对应 MSE;若是 Bernoulli,则常用 BCE。下面推导前一种情形。
设一张图片被展平成 D 维向量。解码器不直接输出一个确定的重建图,而是输出高斯分布的均值:
x^θ(z)=μθ(z),pθ(x∣z)=N(x;x^θ(z),σx2I).
这不是说 x 与 x^θ(z) 相等,而是说:给定 z 后,真实图片 x 被模型视为从一个高斯分布中随机产生。其中:
- N(变量;均值,协方差) 是高斯分布的记号;分号左侧的 x 是该分布要评价/生成的随机变量。
- x^θ(z) 是解码器输入 z 后输出的预测图像,也是该高斯分布的均值。
- σx2I 是协方差矩阵:各像素独立,且都使用同一个方差 σx2。
等价的生成写法是:
x=x^θ(z)+η,η∼N(0,σx2I).
这里观测方差 σx2 是预先固定的,不随 x,z,θ 改变。多维高斯密度为:
pθ(x∣z)=(2πσx2)D/21exp(−2σx21∥x−x^θ(z)∥22).
取负对数:
−logpθ(x∣z)=2Dlog(2πσx2)+2σx21∥x−x^θ(z)∥22.
第一项与解码器参数 θ 无关;第二项前的 2σx21>0 也是固定常数。因此最小化负对数似然,与最小化:
∥x−x^θ(z)∥22=i=1∑D(xi−x^θ,i(z))2
有完全相同的最优解。若再除以 D,就是通常代码里写的 MSE:
MSE(x,x^)=D1i=1∑D(xi−x^i)2.
所以“MSE 重建损失”隐含了一个概率假设:给定 z 后,每个像素等于解码器预测均值加上独立、同方差的高斯噪声。
实现时仍需明确重建项和 KL 项各自的 sum/mean 维度,否则二者的相对权重会随图片尺寸和 batch size 改变。
8.3 最终可实现的 batch loss
令
x^(b,k)=x^θ(z(b,k)),logvarj(b)=log((σj(b))2).
将 batch 负 ELBO 中的重建项替换为 MSE、KL 项替换为对角高斯的闭式解后,忽略固定观测方差带来的常数与缩放,最终最小化的 loss 为:
Lbatch≈B1b=1∑B[K1k=1∑KMSE(x(b),x^(b,k))+21j=1∑d((μj(b))2+exp(logvarj(b))−logvarj(b)−1)].
第一行是对每张图片的 K 次潜变量采样取平均后的重建 MSE;第二行是该图片的正 KL 损失。最常见的设置是 K=1:一个 batch 有 B 张图片,每张图片经重参数化采一个 z,算一次 MSE 和一次 KL,再对 B 个样本取平均。
9. 训练与生成
训练时,输入真实图片 x:
- 编码器输出 μ(x),logσ2(x);
- 用重参数化得到 z;
- 解码器根据 z 输出 x^ 或 pθ(x∣z) 的参数;
- 用重建项与 KL 项更新 θ,ϕ。
生成时,不需要图片,也不使用编码器:
z∼p(z)=N(0,I)⟶x∼pθ(x∣z).
KL 项的作用在这里闭环:训练时让 qϕ(z∣x) 接近 p(z),生成时从先验采到的 z 才会位于解码器见过的潜空间区域。