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()