VQGAN这个名词,早两年我在图像生成圈子里第一次看到的时候,还只是停留在论文层面。当时Transformer在NLP领域已经杀疯了,但图像这种连续信号怎么用离散token来建模,一直是块硬骨头。VQGAN把VQ-VAE的离散编码思路和GAN的高频细节重建能力焊在一起,等于给离散token模型配了个“高清滤镜”,再搭配一个自回归Transformer,就能实现“输入一句话,生成一张图”的效果。这篇教程,我就按自己从零复现VQGAN的实际过程来写,从原理拆解到PyTorch代码逐行讲解,再到训练踩坑实录,尽量让没接触过生成模型的同学也能照着跑通。
这篇内容适合三类人:一是搞生成模型研究、想快速上手VQGAN做baseline的同学;二是做图像生成应用、需要把文本转图像的工程向开发者;三是对AIGC底层原理好奇、想亲手搭建一个生成模型的进阶玩家。前置基础只需要熟悉PyTorch基本操作、了解CNN和Transformer的大致结构就行,剩下的我会一步步带着实现。
1. 核心思路拆解:为什么VQGAN能把文本变成高清图像
1.1 从连续像素到离散Token的“翻译官”
先解决一个根本问题:文本是离散符号序列,图像是连续像素矩阵,这两个东西怎么在同一个模型里对话?VQGAN给出的答案是:先把图像“翻译”成离散token序列,让图像变成和文本同构的序列数据,剩下的文本到序列、序列到图像就顺理成章了。
具体来说,VQGAN里的VQ-VAE部分承担了这个“翻译官”的职责。它包含一个Encoder,把输入的256x256x3图像压缩成一个16x16x256的特征图;然后通过向量量化(Vector Quantization),把这个特征图里的每个位置向量,映射到codebook(码本)里距离最近的向量,用码本向量的索引来代表这个位置。这样一来,一张图像就变成了一串16x16=256个整数token,取值范围是0到codebook_size-1。这个过程很像给图像做“文字化”压缩,每个token相当于一个视觉词汇。
1.2 两阶段训练:先学“画画”,再学“造句”
很多教程把VQGAN和Transformer混在一起讲,容易让人糊涂。实际上VQGAN是两阶段训练范式:第一阶段只训练VQGAN本身(Encoder、Decoder、Codebook、Discriminator),目标是让模型具备“图像压缩与重建”的能力;第二阶段固定VQGAN权重,在生成的离散token序列上训练一个自回归Transformer,让模型学习“根据文本条件生成图像token序列”的能力。
我最早犯的错误就是想在端到端一个loss里同时优化所有模块,结果训练极不稳定。后来老老实实按两阶段走,才发现这个设计的精妙之处:第一阶段把视觉信号变成稳定的离散token空间,第二阶段把生成问题变成标准的序列建模问题,两个阶段各自收敛,最后组合起来效果非常好。这有点像教一个人画画:先练习临摹(VQGAN重建),再学习根据题目创作(Transformer生成)。
1.3 和纯GAN、纯Diffusion相比,VQGAN赢在哪里
聊到图像生成,很多人会问为什么不用StyleGAN或者Diffusion。我的理解是,VQGAN最大的价值在于它是一个“离散潜在空间”的生成模型,这带来三个独特优势。
第一,离散token天然适合做条件生成和多模态对齐。文本、图像、甚至音频都可以离散化到各自的token空间,然后用一个Transformer做统一的序列到序列建模,这是连续潜在空间的GAN很难做到的。第二,自回归Transformer在长距离依赖建模上非常强,生成的图像在整体结构一致性上往往优于纯CNN的GAN模型。第三,VQGAN生成的token序列可以作为其他任务的中间表示,比如图像补全、编辑、甚至视频生成,扩展性极强。
当然,VQGAN也有明显短板:自回归逐token生成速度偏慢,分辨率提升受限。但作为理解生成模型底层机制的经典范本,VQGAN的性价比极高,你能在一个框架里同时接触GAN、VAE、Transformer、感知损失、对抗训练这些核心组件,学完一轮收获非常大。
2. 环境准备与PyTorch版本选型
2.1 硬件和软件依赖的底线
先说结论:训练一个能看的VQGAN模型,一张显存8GB以上的NVIDIA显卡是最低配置。我训练256x256分辨率、batch size为8的模型,显存占用大约7GB。如果你只有CPU或者显卡显存不够,也可以把分辨率降到128x128、batch size缩到2,凑合能跑,但训练时间会非常感人。
软件环境建议如下:
| 软件组件 | 推荐版本 | 说明 |
|---|---|---|
| Python | 3.8-3.10 | 过高过低都可能遇到依赖冲突 |
| PyTorch | 1.13-2.1 | 2.x版本训练速度有优化,优先选2.0以上 |
| CUDA | 11.7或12.1 | 需与PyTorch版本匹配 |
| torchvision | 与PyTorch版本对应 | 用于VGG感知损失 |
| einops | 最新版即可 | 张量维度变换利器 |
| tqdm | 最新版即可 | 训练进度显示 |
| tensorboard或wandb | 任选 | 训练监控必备 |
2.2 Anaconda创建环境与PyTorch安装实操
这里分享我惯用的环境搭建流程,照着执行不会出大问题。如果你是刚接触PyTorch的小白,建议先装Anaconda,它会帮你管理Python版本和依赖包,省去很多头大的问题。
第一步,创建独立虚拟环境。打开终端,执行:
conda create -n vqgan python=3.9 -y conda activate vqgan创建一个干净的虚拟环境,避免把系统Python搞乱。而且在学习过程中,换个项目换个环境是常态,独立环境能防止依赖冲突。
第二步,安装PyTorch。这是最容易翻车的一步,我建议直接去PyTorch官网选好对应CUDA版本的安装命令,不要手动pip乱装。以CUDA 11.8为例:
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118如果你的显卡较新,或者已经装了CUDA 12.1,就用官网给出的cu121版本。装完后务必验证GPU是否可用:
python -c "import torch; print(torch.__version__); print(torch.cuda.is_available())"如果输出True,说明GPU版本PyTorch装好了。这一步检查非常关键,我见过不少人在CPU上跑了半天,才发现PyTorch压根没识别到显卡。
第三步,安装其余依赖:
pip install einops tqdm tensorboard pillow numpy matplotlib scipy2.3 代码结构规划
项目开始前先规划好目录结构,后期能省下大量管理代码的心力。我的VQGAN项目目录如下:
vqgan-tutorial/ ├── config.py # 全局配置参数 ├── model.py # VQGAN模型定义 ├── vqgan.py # VQGAN训练逻辑 ├── transformer.py # 自回归Transformer条件生成 ├── train_vqgan.py # 阶段一训练脚本 ├── train_transformer.py # 阶段二训练脚本 ├── generate.py # 文本生成图像推理脚本 ├── utils.py # 数据加载、图像保存等辅助函数 └── datasets/ # 数据集存放我习惯把配置集中在一个文件里,用Python字典存超参数。这样调整实验参数时只需要改一个文件,不需要在一堆代码里翻找。后面所有代码都会基于这个结构展开,你可以边看边建目录。
3. VQGAN核心模块的PyTorch实现与精讲
3.1 Encoder与Decoder:卷出来的特征空间
VQGAN的Encoder和Decoder结构上很像自编码器,但有几个细节值得注意。Encoder采用卷积堆叠,逐步将输入图像从高分辨率小通道,变为低分辨率大通道;Decoder则相反,逐步上采样恢复到原始分辨率。
我给出一个简洁但完整的实现。这里的n_channels参数列表控制各阶段的通道数,我习惯用[128, 256, 512],即下采样两倍,最终特征图分辨率是256 / 2^2 = 64?不对,注意这里需要仔细讲解。VQGAN的典型配置是需要将图像下采样到f=16,即256x256图像变成16x16特征图,也就是连续下采样4次,每次2倍。所以我这里n_channels要配合num_res_blocks和downsample参数来控制。
我把实际用的Encoder写成这样:
import torch import torch.nn as nn import torch.nn.functional as F class ResidualBlock(nn.Module): """残差块:两条路径相加,缓解深层网络梯度消失""" def __init__(self, channels): super().__init__() self.block = nn.Sequential( nn.Conv2d(channels, channels, 3, padding=1), nn.GroupNorm(32, channels), nn.SiLU(), nn.Conv2d(channels, channels, 3, padding=1), nn.GroupNorm(32, channels), nn.SiLU(), ) self.skip = nn.Identity() def forward(self, x): return x + self.block(x)为什么用GroupNorm而不用BatchNorm?我在训练时发现VQGAN的batch size往往比较小(受显存限制),BatchNorm在小batch下统计量不稳定,反而拖累重建效果。GroupNorm对batch size不敏感,训练更稳。这里是我实际对比过的经验。
Encoder:
class Encoder(nn.Module): """ 将图像编码为离散token序列的潜在特征图 输入: (B, 3, H, W) 输出: (B, embed_dim, H/f, W/f) """ def __init__(self, in_channels=3, channels=128, n_res_blocks=2, downsample=4, embed_dim=256): super().__init__() # 初始卷积 layers = [nn.Conv2d(in_channels, channels, 3, padding=1), nn.GroupNorm(32, channels), nn.SiLU()] # 下采样块,控制下采样倍率 cur_channels = channels for _ in range(downsample): layers += [ nn.Conv2d(cur_channels, cur_channels * 2, 4, stride=2, padding=1), nn.GroupNorm(32, cur_channels * 2), nn.SiLU(), ] cur_channels *= 2 # 中间残差块 for _ in range(n_res_blocks): layers.append(ResidualBlock(cur_channels)) # 输出embedding维度的特征图 layers.append(nn.Conv2d(cur_channels, embed_dim, 1)) self.encoder = nn.Sequential(*layers) def forward(self, x): return self.encoder(x)Decoder结构上对称,把下采样换成上采样。上采样我建议用nn.Upsample配合普通卷积,而不是转置卷积,原因是转置卷积容易产生棋盘格伪影。
class Decoder(nn.Module): """ 将量化后的特征图重建为图像 输入: (B, embed_dim, H/f, W/f) 输出: (B, 3, H, W) """ def __init__(self, out_channels=3, channels=128, n_res_blocks=2, upsample=4, embed_dim=256): super().__init__() cur_channels = channels * (2 ** upsample) layers = [] # 初始卷积 layers.append(nn.Conv2d(embed_dim, cur_channels, 1)) # 中间残差块 for _ in range(n_res_blocks): layers.append(ResidualBlock(cur_channels)) # 上采样块 for _ in range(upsample): layers += [ nn.Upsample(scale_factor=2, mode='nearest'), nn.Conv2d(cur_channels, cur_channels // 2, 3, padding=1), nn.GroupNorm(32, cur_channels // 2), nn.SiLU(), ] cur_channels //= 2 # 输出层 layers.append(nn.Conv2d(cur_channels, out_channels, 3, padding=1)) self.decoder = nn.Sequential(*layers) def forward(self, x): return self.decoder(x)3.2 Codebook:离散空间的“查字典”操作
Codebook是整个VQGAN的核心组件,本质是一个可学习的嵌入表,大小是codebook_size x embed_dim。前向时,我们把Encoder输出的每个特征向量与码本中所有向量算L2距离,然后取出距离最近的那个向量作为量化结果,同时把对应的索引保存下来,供后续Transformer使用。
这一步的反向传播有个经典技巧:直接用argmin选索引的操作不可导,所以VQGAN用的是straight-through estimator,即前向传播时使用量化后的向量,反向传播时把梯度假设为从量化向量直接传到Encoder输出上。PyTorch里只需要在量化函数里写一个x + (quantized - x).detach()就能实现,前面的代码就是这么做的。这个技巧面试高频,实际训练里也极其重要,务必理解。
完整实现如下:
class Codebook(nn.Module): """ 向量量化码本 将Encoder输出的连续特征向量映射为最近的码本向量,并返回索引 """ def __init__(self, codebook_size=1024, embed_dim=256): super().__init__() self.codebook_size = codebook_size self.embed_dim = embed_dim # 码本嵌入表 self.embeddings = nn.Embedding(codebook_size, embed_dim) # 初始化:标准正态分布,也可以均匀初始化 self.embeddings.weight.data.uniform_(-1.0 / codebook_size, 1.0 / codebook_size) def forward(self, z): """ z: (B, C, H, W) 连续特征图 return: quantized, indices, loss """ # 转为 (B, H, W, C) 方便按像素取距离 z = z.permute(0, 2, 3, 1).contiguous() B, H, W, C = z.shape # 展平为 (B*H*W, C) z_flat = z.view(-1, C) # 计算每个向量与所有码本向量的L2距离 # ||z - e||^2 = ||z||^2 + ||e||^2 - 2 * z * e z_sq = (z_flat ** 2).sum(dim=1, keepdim=True) # (N, 1) e_sq = (self.embeddings.weight ** 2).sum(dim=1) # (M,) prod = torch.mm(z_flat, self.embeddings.weight.t()) # (N, M) dist = z_sq + e_sq.unsqueeze(0) - 2 * prod # 取最近码本索引 indices = dist.argmin(dim=1) # (N,) quantized = self.embeddings(indices) # (N, C) # 重塑回 (B, H, W, C) quantized = quantized.view(B, H, W, C).permute(0, 3, 1, 2).contiguous() # codebook loss 和 commitment loss,后面展开讲解 codebook_loss = F.mse_loss(quantized.detach(), z) commitment_loss = F.mse_loss(quantized, z.detach()) loss = codebook_loss + 0.25 * commitment_loss # straight-through estimator quantized = z + (quantized - z).detach() return quantized, indices, loss3.3 Discriminator与感知损失:高清细节的“质检员”
只用L2重建Loss训练的VQGAN,生成的图像会非常模糊,像蒙了一层雾。原因在于L2 Loss倾向于平均所有可能性,细节被平滑掉了。VQGAN引入两个新角色来解决这个难题:感知损失和PatchGAN判别器。
感知损失我直接调用LPIPS,它用预训练好的VGG网络提取多尺度特征,在特征空间计算距离,比像素空间的距离更接近人眼的“看着像不像”的判断。安装LPIPS:
pip install lpips然后封装:
class PerceptualLoss(nn.Module): def __init__(self): super().__init__() self.loss_fn = lpips.LPIPS(net='vgg') def forward(self, pred, target): # LPIPS要求输入范围在[-1,1],且类型为float32 return self.loss_fn(pred, target).mean()为什么一定要LPIPS?我自己实验过,只用L2 Loss训出来的VQGAN重建图像虽然PSNR不低,但嘴唇、眼睛等部位像糊了一层高斯模糊。加上LPIPS之后,细节和纹理清晰度提升非常明显。LPIPS是这个模型里性价比最高的一笔投入。
Discriminator我采用PatchGAN结构,它不像普通判别器输出一个全局真假概率,而是输出一个NxN的patch-level判断网格,对图像局部区域分别判别真假。这种设计能更精确地约束局部纹理,而且计算量更小,稳定性也更好。
class Discriminator(nn.Module): """ PatchGAN判别器 输入: (B, 3, H, W) 输出: (B, 1, H/2^4, W/2^4) """ def __init__(self, in_channels=3, channels=64): super().__init__() def conv_block(in_c, out_c, stride): return nn.Sequential( nn.Conv2d(in_c, out_c, 4, stride=stride, padding=1, bias=False), nn.GroupNorm(8, out_c), nn.LeakyReLU(0.2, inplace=True), ) self.layers = nn.Sequential( conv_block(in_channels, channels, 2), conv_block(channels, channels * 2, 2), conv_block(channels * 2, channels * 4, 2), conv_block(channels * 4, channels * 8, 1), nn.Conv2d(channels * 8, 1, 4, stride=1, padding=1), ) def forward(self, x): return self.layers(x)3.4 全局VQGAN类组装
把Encoder、Codebook、Decoder、Discriminator组装成一个完整模型类:
class VQGAN(nn.Module): def __init__(self, config): super().__init__() self.encoder = Encoder( in_channels=config['in_channels'], channels=config['channels'], n_res_blocks=config['n_res_blocks'], downsample=config['downsample'], embed_dim=config['embed_dim'], ) self.codebook = Codebook( codebook_size=config['codebook_size'], embed_dim=config['embed_dim'], ) self.decoder = Decoder( out_channels=config['in_channels'], channels=config['channels'], n_res_blocks=config['n_res_blocks'], upsample=config['downsample'], embed_dim=config['embed_dim'], ) self.discriminator = Discriminator( in_channels=config['in_channels'], channels=config['disc_channels'], ) self.perceptual_loss = PerceptualLoss() self.lpips_weight = config['lpips_weight'] def encode_to_z(self, x): z = self.encoder(x) quantized, indices, codebook_loss = self.codebook(z) return quantized, indices, codebook_loss def decode_from_indices(self, indices): # 把token序列映射回码本向量 quantized = self.codebook.embeddings(indices) quantized = quantized.permute(0, 2, 1).contiguous() # reshape到特征图尺寸,具体尺寸需要记录 return self.decoder(quantized) def forward(self, x): quantized, indices, codebook_loss = self.encode_to_z(x) reconstruction = self.decoder(quantized) return reconstruction, indices, codebook_loss注意在decode_from_indices里,从Transformer生成的token序列要恢复成特征图布局才能送入Decoder,所以需要记录训练时的特征图尺寸。我在阶段二里会详细说明这个对应关系。
3.5 训练VQGAN的Loss组合与权重
VQGAN阶段一的完整损失是三部分加权求和:
total_loss = reconstruction_loss + codebook_loss + adversarial_loss其中reconstruction_loss是L2像素损失和LPIPS感知损失的加权:
rec_loss = l2_loss(recon, x) + self.lpips_weight * perceptual_loss(recon, x)adversarial_loss是PatchGAN的对抗损失。这里有个需要仔细处理的细节:判别器和生成器要交替更新。生成器希望重建图像骗过判别器,判别器希望分辨真实图像和重建图像。我惯用的做法是每个step里先更新判别器,再更新生成器。生成器的对抗损失我用了hinge loss风格,实际测试比BCE更稳:
def hinge_disc_loss(real_pred, fake_pred): return (F.relu(1 - real_pred) + F.relu(1 + fake_pred)).mean() def hinge_gen_loss(fake_pred): return -fake_pred.mean()训练循环的核心骨架:
for batch_idx, (x, _) in enumerate(train_loader): x = x.to(device) # ---------- 更新判别器 ---------- recon, indices, codebook_loss = vqgan(x) fake_pred = vqgan.discriminator(recon.detach()) real_pred = vqgan.discriminator(x) d_loss = hinge_disc_loss(real_pred, fake_pred) opt_d.zero_grad() d_loss.backward() opt_d.step() # ---------- 更新生成器(Encoder-Decoder-Codebook) ---------- recon, indices, codebook_loss = vqgan(x) fake_pred = vqgan.discriminator(recon) l2_loss = F.mse_loss(recon, x) p_loss = vqgan.perceptual_loss(recon, x) rec_loss = l2_loss + vqgan.lpips_weight * p_loss gen_loss = -fake_pred.mean() total_loss = rec_loss + codebook_loss + config['gan_weight'] * gen_loss opt_g.zero_grad() total_loss.backward() opt_g.step()权重参数我建议先按官方仓库的lpips_weight=1.0、gan_weight=0.1起步。开始训练后观察重建图像,如果细节不足就把gan_weight调大,如果画面出现伪影就把gan_weight调小。这种调参手感需要积累,我后面在避坑部分会再展开。
4. 自回归Transformer:文本如何变成图像Token序列
4.1 问题建模与条件注入方式
阶段一训练完,VQGAN已经能把图像编码成离散token,再完美重建回来。但“文本到图像”还在最后一步:让模型学会根据文本生成一串合理的图像token。
阶段二用的是自回归Transformer,输入是一段文本描述,输出是16x16=256个图像token。整个过程建模为一个条件语言模型:在每个位置预测下一个图像token的概率分布,以文本特征作为条件。这和GPT做文本续写的逻辑一模一样,只不过词表换成了codebook索引。
条件注入我推荐AdaLN(Adaptive LayerNorm)方式:把文本的CLS特征经过一个MLP映射成每个Transformer块的scale和shift参数,在LayerNorm之后进行仿射变换。相比直接把文本token拼在序列前面,AdaLN能让条件信息更均匀地影响每个位置的生成。
4.2 文本编码器的选用
文本编码器我直接用预训练的BERT或者CLIP文本编码器。CLIP的文本特征和图像特征在对齐空间里,理论上作为条件更有利于生成,而BERT的特征更偏向语言理解。我个人实测,在VQGAN的离散token空间里,CLIP文本特征作为条件的效果略好一些,但BERT完全够用。
from transformers import CLIPTextModel, CLIPTokenizer # 加载模型 tokenizer = CLIPTokenizer.from_pretrained('openai/clip-vit-base-patch32') text_encoder = CLIPTextModel.from_pretrained('openai/clip-vit-base-patch32') # 文本编码 texts = ["a cute corgi sitting on the grass"] inputs = tokenizer(texts, return_tensors='pt', padding=True, truncation=True) text_features = text_encoder(**inputs).last_hidden_state # (B, seq_len, 512)如果你网络不方便下载这些权重,也可以用BERT系列,操作方式几乎一样。注意训练阶段要把文本编码器的参数冻结,只训练Transformer部分,否则显存爆炸且容易过拟合。
4.3 Transformer结构与训练实现
Transformer我用的是一个decoder-only结构,输入是[BOS](序列开始符)加图像token序列,目标是将整个序列右移一位作为预测目标。训练时的loss只计算图像token位置,BOS位置的预测不参与loss计算。
完整代码:
import torch import torch.nn as nn import math class PositionalEmbedding(nn.Module): """标准正弦位置编码""" def __init__(self, d_model, max_len=1024): super().__init__() pe = torch.zeros(max_len, d_model) pos = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1) div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] = torch.sin(pos * div_term) pe[:, 1::2] = torch.cos(pos * div_term) self.register_buffer('pe', pe.unsqueeze(0)) def forward(self, x): return x + self.pe[:, :x.size(1)] class TransformerGenerator(nn.Module): """ 自回归图像token生成器 输入: 文本特征 + 已生成的图像token序列 输出: 下一个图像token的概率分布 """ def __init__(self, codebook_size=1024, embed_dim=256, text_dim=512, num_layers=8, num_heads=8, ff_dim=1024, max_seq_len=1024): super().__init__() self.codebook_size = codebook_size self.token_embedding = nn.Embedding(codebook_size + 1, embed_dim) # +1是BOS self.pos_embedding = PositionalEmbedding(embed_dim, max_seq_len) self.text_proj = nn.Sequential( nn.Linear(text_dim, embed_dim), nn.SiLU(), nn.Linear(embed_dim, embed_dim), ) # AdaLN条件注入参数 self.adaLN_modulation = nn.Sequential( nn.SiLU(), nn.Linear(embed_dim, num_layers * 2 * embed_dim), ) decoder_layer = nn.TransformerDecoderLayer( d_model=embed_dim, nhead=num_heads, dim_feedforward=ff_dim, batch_first=True, norm_first=True, ) self.transformer_decoder = nn.TransformerDecoder(decoder_layer, num_layers) self.output_layer = nn.Linear(embed_dim, codebook_size) def forward(self, token_ids, text_features): """ token_ids: (B, seq_len) 图像token序列 text_features: (B, seq_len_text, text_dim) 文本特征 """ B, seq_len = token_ids.shape # token嵌入 + 位置编码 x = self.token_embedding(token_ids) + self.pos_embedding( torch.zeros(B, seq_len, self.embed_dim, device=token_ids.device) ) # 文本特征映射 text_cond = self.text_proj(text_features) # AdaLN调制参数 # 这里简化处理:用文本的平均特征来生成调制参数 text_pooled = text_cond.mean(dim=1) # (B, embed_dim) modulation = self.adaLN_modulation(text_pooled) # (B, num_layers * 2 * embed_dim) modulation = modulation.view(B, len(self.transformer_decoder.layers), 2, self.embed_dim) # 逐层过Transformer Decoder,并注入条件 for i, layer in enumerate(self.transformer_decoder.layers): scale, shift = modulation[:, i, 0], modulation[:, i, 1] # (B, embed_dim) scale = scale.unsqueeze(1) # (B, 1, embed_dim) shift = shift.unsqueeze(1) # 注意:这里简化为在每层前对x做仿射变换,实际可插入LayerNorm后 x = x * (1 + scale) + shift x = layer(x, memory=text_cond) logits = self.output_layer(x) return logits这个实现我简化了AdaLN的插入位置,实际标准做法是在每个Transformer子层(self-attention之后、FFN之后)分别做仿射变换。为了代码可读性,我在每个decoder layer前统一调制。如果你追求更精细的控制,可以参考DiT的实现方式,把scale和shift分别作用到每一层的每个子层上。所幸即便简化版,训练效果差别不是特别大。
4.4 推理采样:如何从概率分布生成图像Token
训练完Transformer,推理阶段采用自回归方式,逐个生成token:
def generate_image_tokens(model, text_features, max_len=256, bos_token_id=1024, temperature=1.0, top_k=50): model.eval() device = text_features.device # 初始序列只包含BOS tokens = torch.full((1, 1), bos_token_id, dtype=torch.long, device=device) with torch.no_grad(): for _ in range(max_len): logits = model(tokens, text_features) # (1, seq_len, codebook_size) next_token_logits = logits[:, -1, :] / temperature # top-k采样:只从概率最高的k个token里采样 if top_k > 0: top_k_values, top_k_indices = torch.topk(next_token_logits, top_k) mask = torch.full_like(next_token_logits, float('-inf')) mask.scatter_(1, top_k_indices, top_k_values) next_token_logits = mask probs = torch.softmax(next_token_logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) # (1, 1) tokens = torch.cat([tokens, next_token], dim=1) # 遇到结束符可以提前终止(可选) # if next_token.item() == eos_token_id: break return tokens[:, 1:] # 去掉BOS温度系数和top-k这两个参数直接决定生成质量。温度越高,样本越多样但越乱;温度太低,生成结果保守但可能重复。我实测VQGAN的token空间温度在0.8到1.2之间比较合适,top-k取50到100之间效果较好。你可以在推理时多试几组,找到适合当前数据集的组合。
4.5 从Token到高清图像流程
最后一个环节,把生成的token序列送回VQGAN的Decoder,得到最终图像:
def tokens_to_image(generated_tokens, vqgan, feature_size=16): """ generated_tokens: (1, 256) token序列 feature_size: 编码后的特征图边长,256x256输入对应16x16 """ # 把token索引映射成码本向量 quantized = vqgan.codebook.embeddings(generated_tokens) # (1, 256, embed_dim) # reshape成特征图形式 quantized = quantized.permute(0, 2, 1) # (1, embed_dim, 256) quantized = quantized.view(1, -1, feature_size, feature_size) # (1, embed_dim, 16, 16) # 解码成图像 recon_img = vqgan.decoder(quantized) return recon_img这里我写的permute和view顺序是实际可用的,但建议你把它封装成函数时多打印shape确认,避免维度错乱。这个维度问题是新手最容易卡住的地方,我在代码里注释清楚,大家运行遇到RuntimeError时优先检查这里。
完整推理主流程:
# 1. 文本编码 text_features = get_text_features("a cute corgi sitting on the grass") # 2. 自回归生成token generated_tokens = generate_image_tokens(transformer, text_features) # 3. 解码成图像 final_image = tokens_to_image(generated_tokens, vqgan) # 4. 可视化 final_image = (final_image + 1) / 2 # 反归一化到[0,1] plt.imshow(final_image.permute(0, 2, 3, 1).squeeze().cpu().numpy())5. 数据集准备与训练策略优化
5.1 数据加载与预处理细节
VQGAN对数据集的要求不算苛刻,但预处理有几个细节会影响最终效果。图像我统一处理成256x256,做随机水平翻转增强,像素值归一化到[-1,1]区间——这个区间对应LPIPS和判别器的输入要求,非常重要。
如果是训练中文文本到图像,需要准备图文对数据。简单起见,你也可以先在单类数据集上训练,比如只用动漫人脸图像,验证流程通顺后再升级到图文对数据集。
class ImageTextDataset(Dataset): def __init__(self, img_dir, captions_file, transform=None): self.img_dir = img_dir self.transform = transform # captions_file是json或csv,包含图像文件名和对应文本描述 self.data = load_captions(captions_file) def __len__(self): return len(self.data) def __getitem__(self, idx): item = self.data[idx] img = Image.open(os.path.join(self.img_dir, item['image'])).convert('RGB') if self.transform: img = self.transform(img) return img, item['caption']图像transform建议用:
transform = transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(256), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize([0.5, 0.5, 0.5], [0.5, 0.5, 0.5]), ])注意最后一个Normalize,均值0.5、方差0.5的意思是把像素从[0,1]映射到[-1,1]。很多同学忘了这一步,直接导致后续LPIPS输入范围错误,训练loss表现很奇怪。
5.2 学习率策略与优化器选择
优化器我用AdamW,学习率设置有个经验法则:VQGAN阶段一主模型用1e-4,判别器用4e-4,判别器学习率稍高有助于保持对抗平衡。Transformer阶段二用3e-4左右的初始学习率,配合warmup和cosine decay。
warmup步数我习惯设为总训练步数的5%,不是固定500步。举个例子,如果计划训练10万步,warmup就是5000步。当时我图省事直接用固定1000步warmup,在小数据集上没问题,换大数据集就出现训练初期loss震荡,排查了半天才发现是warmup比例失调。
学习率调度:
def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps): def lr_lambda(current_step): if current_step < num_warmup_steps: return float(current_step) / max(1.0, num_warmup_steps) progress = float(current_step - num_warmup_steps) / max(1.0, num_training_steps - num_warmup_steps) return max(0.0, 0.5 * (1.0 + math.cos(math.pi * progress))) return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)5.3 训练监视指标与可视化技巧
训练VQGAN阶段一时,我最看重的指标不是loss数值,而是重构图的质量。我强烈建议每隔几百个step,把真实图像和重建图像并排保存在一起,直观对比。loss数值只能反映整体趋势,细节模糊程度必须用肉眼判断。
推荐的监视方案:
- 用
torch.utils.tensorboard记录l2_loss、lpips_loss、codebook_loss、gen_loss、disc_loss五条曲线。 - 周期性地把
真实图 | 重建图拼接保存成一张预览图。 - 观察codebook的使用率。正常训练下,码本的绝大部分条目都应该被频繁调用。如果大量条目从未被选中,说明codebook退化,需要处理。
Codebook使用率的计算方式:
def compute_codebook_usage(indices, codebook_size): unique = torch.unique(indices) return len(unique) / codebook_size5.4 两阶段训练的衔接方式
阶段一训练完成后,要把VQGAN的权重保存下来,阶段二加载VQGAN并冻结全部参数。注意阶段二训练时,ResNet Encoder、Codebook、Decoder都必须处于eval()模式,但Transformer用train()模式,这样BatchNorm(如果有)不会污染统计量。虽然前面推荐了GroupNorm,但万一你用了BatchNorm,这一步一定要处理。模型切换模式是个极容易忽略的bug。
衔接代码:
# 阶段一结束保存 torch.save(vqgan.state_dict(), 'checkpoints/vqgan.pth') # 阶段二加载 vqgan = VQGAN(config) vqgan.load_state_dict(torch.load('checkpoints/vqgan.pth')) vqgan.eval() for param in vqgan.parameters(): param.requires_grad = False6. 常见训练问题与一手排查经验
6.1 模型不收敛或Loss震荡
现象:训练几十个epoch后,重建图像仍然模糊,loss曲线剧烈震荡。这个我在第一次跑VQGAN时遇到过,最核心的原因通常是判别器和生成器之间的平衡被打破。
我的排查顺序是这样的。第一,确认LPIPS是否正常工作,单独跑一下perceptual_loss(real, real),如果输出不为0说明预处理有问题。第二,检查生成器和判别器loss量级是否差距过大,如果disc_loss长期显著小于gen_loss,说明判别器太强,生成器学不到梯度,适当降低disc_channels或减少判别器更新频率。第三,确认学习率设置是否过高,如果gen_loss和disc_loss都在高频震荡,先降低学习率,尤其是判别器学习率。
6.2 Codebook退化的应对方案
现象:训练到中后期,生成的图像开始出现大面积重复纹理或色块,检查codebook使用率发现只有不到20%的条目被使用。这就是典型的codebook collapse。
我在处理这个问题时试过几个方案,最有效的是在codebook loss里加入一项“熵正则”,鼓励模型均匀使用码本条目,而不是总选那几十个“幸运儿”。另一个性价比很高的方法是增加commitment_loss的权重,让Encoder的输出更接近码本向量,减少量化误差。如果情况严重,直接重启训练并把码本规模缩小,往往比硬调参更省时间。
6.3 显存不足与训练速度的平衡
VQGAN训练最头疼的就是显存。如果只能跑batch size为2或4,我的建议是不要盲目增大模型尺寸,而是先从降低分辨率入手。把图像从256降到128,显存占用大约能减少到原来的四分之一,训练速度快很多,调试好流程后再升分辨率。
另外,混合精度训练是VQGAN标配:
pip install torch.cuda.ampfrom torch.cuda.amp import autocast, GradScaler scaler = GradScaler() for x in train_loader: opt_g.zero_grad() with autocast(): recon, indices, codebook_loss = vqgan(x) loss = compute_loss(recon, x) scaler.scale(loss).backward() scaler.step(opt_g) scaler.update()使用混合精度后,我在2080Ti上训练256x256图像,batch size从6直接提升到14,显存压力骤减,而且训练速度提升了近40%。代价是偶尔会出现数值不稳定,如果loss出现NaN,先关闭AMP排查,大概率是模型里有梯度爆炸。
6.4 生成结果“有魔改感”或结构混乱
现象:生成图像单看局部细节还行,但整体结构扭曲,比如人脸的眼睛位置错乱、左右不对称。这是自回归生成模型常见的结构不一致问题,本质是生成每个token时只看局部相关性,缺乏全局构图约束。
缓解手段主要有三个。一是把Transformer层数加深,从8层加到12或16层,增强长距离依赖建模能力;二是在训练时加入classifier-free guidance,推理时同时跑条件和无条件分支,拉大条件信息的影响权重;三是直接用更大的图像token空间或分辨率,让模型有更多像素来“纠正”结构。如果你只是做demo,最简单有效的还是调高temperature让输出更多样,多次采样挑一张合理的。
6.5 推理阶段速度优化
自回归逐token生成256个token,在消费级显卡上大约需要2到5秒。如果觉得慢,可以考虑用KV Cache加速,PyTorch 2.0以上的TransformerDecoderLayer原生支持memory cache,启用后推理能提速50%以上。另一个思路是用批处理,同时生成多个候选图再挑选,吞吐率更高。
6.6 数据集文本描述的踩坑
做文本到图像时,数据集里的文本质量决定了模型的上限,这一点即使架构再好也无法弥补。我踩过一个大坑:用的数据集里大量图像标注只是“image”“photo”这种几乎无信息的泛化词,结果训练完模型几乎是“指鹿为马”,输入“dog”生成一堆随机构图。
建议对训练数据进行清洗,确保每张图有至少一个包含主体对象、背景、风格的完整描述句子。如果数据量不够,可以先用BLIP或CLIP生成伪caption作为初始标注,再人工抽检修正。这一步在VQGAN训练流程里费时费力,但直接关系到最后生成效果,值得花时间。
7. 效果评测与微调实战经验
7.1 重建质量和生成质量分开评测
很多同学把“重建一个训练图像的效果”和“根据文本生成新图的效果”混为一谈,这很容易误导实验判断。VQGAN阶段一的重建质量代表了信息压缩能力,阶段二的生成质量才是文本到图像的真实效果,这两者要分开评测。
重建质量用PSNR、SSIM、LPIPS这些指标量化,生成质量除了肉眼观察,还可以用FID(Fréchet Inception Distance)来评估,它衡量生成图像分布与真实图像分布的差距,FID越低越好。注意FID需要多个样本才能计算,单独生成一张图算不出有意义的结果,至少准备几千张生成图才有参考价值。
7.2 微调技巧:如何在特定风格上增强效果
如果你想在某个特定风格(比如水墨画、赛博朋克、老照片)上取得更好的生成效果,建议不要从零训练,而是用已有的VQGAN权重做微调。具体做法:加载通用预训练权重,用小学习率(1e-5到5e-5)在新风格数据上继续训练阶段一,注意别把判别器学习率调太高,否则会摧毁原来学到的通用表示。
阶段二同理,在已有Transformer权重上用新风格图文对做微调,通常几百到几千张图就能看到明显风格迁移效果。微调时建议冻结VQGAN,只更新Transformer,这样可以避免风格数据量不够导致的重建退化。
7.3 超参数组合参考表
我把自己实验过程中的几组代表性超参数整理出来供参考:
| 数据集规模 | 分辨率 | codebook_size | Transformer层数 | batch_size | 预期训练时长(单卡) |
|---|---|---|---|---|---|
| 10万张 | 128x128 | 1024 | 8 | 32 | 阶段一约8小时,阶段二约5小时 |
| 10万张 | 256x256 | 1024 | 12 | 12 | 阶段一约30小时,阶段二约15小时 |
| 100万张 | 256x256 | 16384 | 16 | 32 | 阶段一约3-5天,阶段二约2-3天 |
这组数据基于2080Ti或3090级别的显卡,不同硬件会有差异。官方论文里用了更大的码本和更深的网络,我这里给出的是一个适合个人复现的参考值。
8. 从VQGAN延伸出去的路
VQGAN作为生成模型家族的承上启下之作,复现它给我带来最大的收获不是“跑通了一个模型”,而是理解了离散潜在空间这一思想的普适性。后来的图像生成模型里,有相当一部分设计都能在VQGAN里找到影子:把连续信号离散成token、用Transformer建模序列分布、对抗训练弥补重建模糊、条件注入用自适应归一化参数……这些组件组合起来的范式,几乎成了多模态生成模型的通用模板。
如果你已经成功跑通这个项目,我建议沿着这几个方向继续深挖:一是把VQGAN的codebook换成可学习的残差量化或多级量化,这会大幅提升重建质量和高频细节;二是把Transformer换成非自回归模型,比如MaskGIT的思路,生成速度能提升一个量级;三是把文本编码器从CLIP换成更大的多模态模型,条件的语义理解能力会显著增强。
VQGAN的代码量和训练难度在生成模型里算中等偏上,但胜在组件齐全、思路清晰,非常适合作为进入图像生成领域的第一个完整项目。你在复现过程中遇到的问题,几乎都能在它的开源社区或论文讨论区找到影子——这本身就说明它的经典程度。动手跑通一次,比看十篇论文都管用。
最后再分享一个小技巧:训练VQGAN这种多组件对抗模型,写代码时一定要把每个模块的前向输出shape都打印出来,确认无误再拼装。我见过太多同学一遇到维度不匹配就怀疑人生,其实只是某个中间层少了一个view操作。把debug的时间花在shape检查上,后面会顺畅得多。