VQGAN 矢量量化生成模型
FreeGuideOnline
19阅读
2026-07-13
python
编码器前向
z_e = encoder(real_img) z_q, indices, commit_loss = vector_quantize(z_e, codebook) rec_img = decoder(z_q)
重建损失(包含感知损失和L1)
rec_loss = l1_loss(rec_img, real_img) + perceptual_loss(rec_img, real_img)
对抗损失
fake_pred = discriminator(rec_img) real_pred = discriminator(real_img) adv_loss = hinge_generator_loss(fake_pred)
总生成器损失
g_loss = rec_loss + commit_loss * beta + adv_loss * lambda_adv g_loss.backward() optimizer_G.step()
判别器损失
d_loss = hinge_discriminator_loss(real_pred, fake_pred) d_loss.backward() optimizer_D.step()
这里 `vector_quantize` 实现了最近邻搜索和直通梯度(straight-through estimator),以便梯度回传。实际实现会用到 `torch.cdist` 和自定义梯度。
### 阶段二:训练自回归 Transformer
冻结 VQGAN 编码器(得到索引序列),然后训练一个因果 GPT 模型:
```python
# 获取图像token序列
with torch.no_grad():
z_e = vqgan_encoder(img) # (B, h, w, d)
_, indices, _, _ = vector_quantize(z_e, frozen_codebook)
indices = indices.reshape(B, h*w) # 展平为序列长度 T
# 加上起始标记(如 <sos>)
input_seq = torch.cat([sos_token * torch.ones(B,1).long(), indices], dim=1)
# Transformer 输入
logits = transformer(input_seq[:, :-1]) # 预测下一个词
loss = cross_entropy(logits.reshape(-1, K), indices.reshape(-1))
loss.backward()
optimizer.step()