VAE/VQ-VAE
VAE 与 VQ-VAE 详解
第一部分:VAE(变分自编码器, Variational AutoEncoder)
1. 核心问题:AE(普通自编码器)的不足
- AE 结构:编码器 z=f(x)z = f(x)z=f(x),解码器 x^=g(z)\hat{x} = g(z)x^=g(z)。目标是让 x^≈x\hat{x} \approx xx^≈x。
- AE 的问题:它学到的隐空间(latent space,即 zzz 所在的空间)不连续、不规则。编码器只是把每个输入映射到一个确定的点 zzz。在两个真实样本的编码点之间(空白区域),解码器可能会生成毫无意义的图像。因此 AE 不能用来随机生成新样本。

2. VAE 的核心思想:概率化 + 强制分布
VAE 不再把输入 xxx 映射成一个点 zzz,而是映射成一个概率分布(通常是高斯分布)。然后从这个分布中采样得到 zzz,再送入解码器。
关键约束:强制所有编码出的分布都向标准正态分布 N(0,I)\mathcal{N}(0, I)N(0,I) 看齐。这样整个隐空间就被“标准化”成连续、完整的标准正态分布空间。随机从这个空间采样一个点,解码器就能生成有意义的、多样的新样本。
3. VAE 的结构
- 编码器:输入 xxx,输出两个向量 μ\muμ 和 log(σ2)\log(\sigma^2)log(σ2)(或直接用 σ\sigmaσ)。它们定义了一个高斯分布 N(μ,σ2I)\mathcal{N}(\mu, \sigma^2 I)N(μ,σ2I)。
- 采样层:z=μ+σ⊙ϵz = \mu + \sigma \odot \epsilonz=μ+σ⊙ϵ,其中 ϵ∼N(0,I)\epsilon \sim \mathcal{N}(0, I)ϵ∼N(0,I)。(这是重参数化技巧,使梯度可以回传)
VAE 中有一个问题:
z∼N(μ,σ^2)
采样操作本身不可导。如果直接从这个分布中采样,梯度无法从 Decoder 回传到 Encoder。
所以 VAE 使用 Reparameterization Trick:
z=μ+σ⊙ϵ
其中:
ϵ∼N(0,I)
这样随机性被转移到了 ϵ 上,而 μ 和 σ 仍然是可导的。
- 解码器:输入 zzz,输出重构 x^\hat{x}x^。
4. 损失函数(重要)
VAE 的损失由两部分组成:
- 重构损失:衡量 x^\hat{x}x^ 和 xxx 的差异。通常是 MSE(对于连续数据)或交叉熵(对于二值图像)。目标:让解码器能根据 zzz 重建出输入。
- KL 散度损失:衡量编码器输出的分布 N(μ,σ2)\mathcal{N}(\mu, \sigma^2)N(μ,σ2) 与标准正态分布 N(0,I)\mathcal{N}(0, I)N(0,I) 之间的距离。目标:迫使隐空间规整、连续。
KL=−12∑i=1d(1+log(σi2)−μi2−σi2) KL = -\frac{1}{2} \sum_{i=1}^{d} \left(1 + \log(\sigma_i^2) - \mu_i^2 - \sigma_i^2\right) KL=−21i=1∑d(1+log(σi2)−μi2−σi2)
其中 ddd 是隐变量 zzz 的维度。
总损失:LVAE=Lrecon+KL\mathcal{L}_{VAE} = \mathcal{L}_{recon} + KLLVAE=Lrecon+KL。
5. 优点与缺点
- 优点:
- 隐空间连续且完整,可以随机生成新样本。
- 可以在隐空间进行插值(如将“笑脸”图片的 zAz_AzA 和“眼镜”图片的 zBz_BzB 平滑过渡,生成“笑脸→眼镜”的中间结果)。
- 理论基础扎实,基于变分推断。
- 缺点:
- 模糊性:这是 VAE 最常被诟病的问题。生成的图像往往偏模糊,不如 GAN 清晰。原因有两个:
- 重构损失(如 MSE)天然倾向于取平均值,导致边缘模糊。
- KL 损失强制所有后验分布向同一个先验靠拢,可能“牺牲”了部分细节。
- 常需要权衡 β\betaβ 系数(LVAE=Lrecon+β⋅KL\mathcal{L}_{VAE} = \mathcal{L}_{recon} + \beta \cdot KLLVAE=Lrecon+β⋅KL,即 β\betaβ-VAE)。
- 模糊性:这是 VAE 最常被诟病的问题。生成的图像往往偏模糊,不如 GAN 清晰。原因有两个:
6. 代码
- VAE 的伪代码
import torch
import torch.nn as nn
import torch.nn.functional as F
class VAE(nn.Module):
def __init__(self, input_dim=784, hidden_dim=400, latent_dim=20):
super().__init__()
self.encoder = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
)
self.fc_mu = nn.Linear(hidden_dim, latent_dim)
self.fc_logvar = nn.Linear(hidden_dim, latent_dim)
self.decoder = nn.Sequential(
nn.Linear(latent_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, input_dim),
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)
z = mu + eps * std
return z
def decode(self, z):
return self.decoder(z)
def forward(self, x):
mu, logvar = self.encode(x)
z = self.reparameterize(mu, logvar)
x_hat = self.decode(z)
return x_hat, mu, logvar
def vae_loss(x, x_hat, mu, logvar):
recon_loss = F.mse_loss(x_hat, x, reduction="sum")
kl_loss = -0.5 * torch.sum(
1 + logvar - mu.pow(2) - logvar.exp()
)
loss = recon_loss + kl_loss
return loss, recon_loss, kl_loss
第二部分:VQ-VAE(矢量量化变分自编码器, Vector Quantized Variational AutoEncoder)
1. 为什么需要 VQ-VAE?
VAE 使用连续的隐变量,这是导致生成结果模糊的部分原因。VQ-VAE 提出:如果使用离散的隐表示,可能能更好地捕捉数据中的类别性信息(比如语音中的音素、图像中的物体部件),同时结合强大的自回归模型(如 PixelCNN)来生成清晰的样本。
2. 核心思想:离散化 + 嵌入空间(Codebook)

VQ-VAE 不再直接输出连续向量 zzz,而是维护一个可学习的嵌入空间(codebook):e1,e2,…,eK∈RDe_1, e_2, \dots, e_K \in \mathbb{R}^De1,e2,…,eK∈RD,其中 KKK 是离散类别数,DDD 是每个向量的维度。
过程:
- 编码器:输入 xxx,输出连续特征图 ze(x)z_e(x)ze(x),形状如 [h,w,D][h, w, D][h,w,D]。
- 量化(Vector Quantization, VQ):对于 ze(x)z_e(x)ze(x) 中的每一个 DDD 维向量,在 codebook 中找到最接近的嵌入向量 eke_kek(如用欧氏距离)。将该位置的原始向量替换为 eke_kek,得到量化后的 zq(x)z_q(x)zq(x)。
zq(x)=ek,其中 k=argminj∥ze(x)−ej∥2 z_q(x) = e_k, \quad \text{其中 } k = \arg\min_j \| z_e(x) - e_j \|_2 zq(x)=ek,其中 k=argjmin∥ze(x)−ej∥2 - 解码器:输入量化后的 zq(x)z_q(x)zq(x),输出重构 x^\hat{x}x^。
- 先验(Prior):训练完编码器/解码器/codebook 后,在离散隐变量上训练一个自回归模型(如 PixelCNN),来学习隐变量的分布 p(z)p(z)p(z)。生成时,先由自回归模型采样得到离散编码序列,再通过解码器生成数据。
3. 损失函数(也是关键)
VQ-VAE 的损失有三部分,巧妙解决了梯度无法通过 argmin 操作传递的问题(使用 straight-through estimator,即直通估计器,前向用 argmin,反向传播时直接将梯度从 zqz_qzq 拷贝给 zez_eze):
- 重构损失:与 VAE 相同,优化编码器和解码器。它让 Decoder 能够根据量化后的 latent code 重建图像。
- Codebook 损失:让 codebook 中的向量 eee 靠近编码器输出的 zez_eze(Encoder 输出了一堆连续特征,codebook 需要学习去覆盖这些特征分布。)。
- 承诺损失 Commitment Loss:让编码器的输出 zez_eze 靠近它被映射到的 codebook 向量 eee(Encoder 要“承诺”使用某些 codebook vector,而不是输出离 codebook 很远的连续向量。)。
LVQ−VAE=logp(x∣zq(x))+∥sg[ze(x)]−e∥22+β∥ze(x)−sg[e]∥22 \mathcal{L}_{VQ-VAE} = \log p(x|z_q(x)) + \| \text{sg}[z_e(x)] - e \|_2^2 + \beta \| z_e(x) - \text{sg}[e] \|_2^2 LVQ−VAE=logp(x∣zq(x))+∥sg[ze(x)]−e∥22+β∥ze(x)−sg[e]∥22
- sg[⋅]\text{sg}[\cdot]sg[⋅] 表示 stop-gradient(停止梯度),即该项不参与梯度更新。
- 实践中常简化:第一项是重构损失,后两项合并优化:更新 eee 时使用 ∥ze−e∥22\|z_e - e\|_2^2∥ze−e∥22,更新 zez_eze 时只使用第三项的梯度。
4. VQ-VAE 的优势
- 清晰生成:通过离散编码 + 强大的自回归先验(如 PixelCNN),生成结果比普通 VAE 清晰很多,避免了“模糊平均”问题。
- 无后验塌陷:离散隐空间自然避免了 VAE 中常见的后验坍缩到先验的问题。
- 层级化表示:可以构建多层 VQ-VAE,在不同抽象层级捕捉信息(比如底层编码纹理,高层编码形状)。
- 可用于多种模态:图像、语音、视频,尤其适用于需要离散 token 的下游任务(如文本到图像生成)。
5. 缺点
- 训练分两阶段:先训练 VQ-VAE(编码器、解码器、codebook),再单独训练自回归先验。比较繁琐且耗时。
- codebook 可能利用率不足(部分嵌入向量从未被使用),需要额外的技巧(如重启机制)。
- 相比普通 VAE,超参数(codebook 大小 KKK、维度 DDD、承诺系数 β\betaβ)比较敏感,需要较多调参。
6. 代码
- Vector Quantizer
import torch
import torch.nn as nn
import torch.nn.functional as F
class VAE(nn.Module):
def __init__(self, input_dim=784, hidden_dim=400, latent_dim=20):
super().__init__()
self.encoder = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
)
self.fc_mu = nn.Linear(hidden_dim, latent_dim)
self.fc_logvar = nn.Linear(hidden_dim, latent_dim)
self.decoder = nn.Sequential(
nn.Linear(latent_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, input_dim),
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)
z = mu + eps * std
return z
def decode(self, z):
return self.decoder(z)
def forward(self, x):
mu, logvar = self.encode(x)
z = self.reparameterize(mu, logvar)
x_hat = self.decode(z)
return x_hat, mu, logvar
def vae_loss(x, x_hat, mu, logvar):
recon_loss = F.mse_loss(x_hat, x, reduction="sum")
kl_loss = -0.5 * torch.sum(
1 + logvar - mu.pow(2) - logvar.exp()
)
loss = recon_loss + kl_loss
return loss, recon_loss, kl_loss
- VQ-VAE
class VQVAE(nn.Module):
def __init__(self, encoder, decoder, num_embeddings=1024, embedding_dim=256):
super().__init__()
self.encoder = encoder
self.quantizer = VectorQuantizer(
num_embeddings=num_embeddings,
embedding_dim=embedding_dim,
beta=0.25,
)
self.decoder = decoder
def forward(self, x):
z_e = self.encoder(x)
z_q, vq_loss, indices = self.quantizer(z_e)
x_hat = self.decoder(z_q)
return x_hat, vq_loss, indices
第三部分:VAE vs VQ-VAE 直观对比
| 特性 | VAE | VQ-VAE |
|---|---|---|
| 隐空间类型 | 连续高斯分布 | 离散码本(codebook) |
| 生成方式 | 直接从连续空间中采样 → 解码 | 自回归模型(如 PixelCNN)采样离散索引 → 查表获得向量 → 解码 |
| 编码器输出 | μ,logσ^2 | 连续特征 z_e |
| 采样方式 | 重参数化技巧 | 最近码本查询 |
| 损失函数 | 重建损失+KL 散度 | 重建损失+Codebook 损失+ 承诺随时 |
| 生成质量 | 较模糊 | 清晰(尤其配合强大自回归先验) |
| 隐空间结构 | 连续、可插值 | 离散、不连续(不同 token 间无直接过渡) |
| 训练复杂度 | 单阶段,容易 | 两阶段(自编码器 + 自回归先验),更复杂 |
| 代表应用 | 文本生成、异常检测、平滑插值 | 语音合成(WaveNet)、图像生成(VQ-VAE-2, DALL-E)、无监督学习 |
常见QA
- 问题 1:VAE 和普通 AutoEncoder 有什么区别?
可以回答:
普通 AutoEncoder 的 Encoder 输出一个确定性的 latent vector,而 VAE 的 Encoder 输出 latent distribution 的参数,例如均值和方差。VAE 通过从该分布采样 latent,并使用 KL divergence 约束 posterior 接近标准正态分布,使 latent space 更连续、更规整、更适合采样生成。
- 问题 2:VAE 为什么需要 KL loss?
可以回答:
KL loss 用来约束 Encoder 输出的近似后验 q(z∣x) 接近先验分布 p(z)=N(0,I)。如果没有 KL loss,VAE 就退化成普通 AutoEncoder,latent space 可能是不连续、不规整的,无法直接从标准正态分布采样生成合理样本。
- 问题 3:VAE 为什么需要重参数化技巧?
可以回答:
因为从 N(μ,σ^2) 中直接采样是不可导的,梯度无法从 Decoder 传回 Encoder。重参数化技巧把采样写成 z=μ+σϵ,其中 ϵ∼N(0,I)。这样随机性被转移到独立噪声 ϵ 上,而 μ 和 σ 仍然可以通过反向传播更新。
- 问题 4:VQ-VAE 的 codebook 是什么?
可以回答:
Codebook 是一个可学习的 embedding table,其中包含 K 个 embedding vector。Encoder 输出连续 latent 后,VQ-VAE 会为每个 latent vector 找到最近的 codebook embedding,并用该 embedding 替换原连续向量,从而实现连续特征到离散 token 的量化。
- 问题 5:VQ-VAE 为什么需要 commitment loss?
可以回答:
Commitment loss 用来约束 Encoder 输出靠近它选择的 codebook vector。如果没有这个约束,Encoder 输出可能在连续空间中不断变化,离 codebook embedding 很远,导致量化不稳定。Commitment loss 让 Encoder “承诺”使用某些 codebook entries,从而稳定训练。
- 问题 6:VQ-VAE 中 argmin 不可导怎么办?
可以回答:
VQ-VAE 使用 straight-through estimator。前向传播时使用 nearest codebook vector,反向传播时近似地把 quantized latent 的梯度直接传给 Encoder 输出。代码中常用 z_q = z_e + (z_q - z_e).detach() 实现。
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)