VQGAN原理与PyTorch实战:文本生成高清图像全解析

📅 发布时间:2026/9/16 21:47:07
VQGAN原理与PyTorch实战:文本生成高清图像全解析
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码本里距离最近的向量用码本向量的索引来代表这个位置。这样一来一张图像就变成了一串16x16256个整数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凑合能跑但训练时间会非常感人。软件环境建议如下软件组件推荐版本说明Python3.8-3.10过高过低都可能遇到依赖冲突PyTorch1.13-2.12.x版本训练速度有优化优先选2.0以上CUDA11.7或12.1需与PyTorch版本匹配torchvision与PyTorch版本对应用于VGG感知损失einops最新版即可张量维度变换利器tqdm最新版即可训练进度显示tensorboard或wandb任选训练监控必备2.2 Anaconda创建环境与PyTorch安装实操这里分享我惯用的环境搭建流程照着执行不会出大问题。如果你是刚接触PyTorch的小白建议先装Anaconda它会帮你管理Python版本和依赖包省去很多头大的问题。第一步创建独立虚拟环境。打开终端执行conda create -n vqgan python3.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的典型配置是需要将图像下采样到f16即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, padding1), nn.GroupNorm(32, channels), nn.SiLU(), nn.Conv2d(channels, channels, 3, padding1), 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不敏感训练更稳。这里是我实际对比过的经验。Encoderclass Encoder(nn.Module): 将图像编码为离散token序列的潜在特征图 输入: (B, 3, H, W) 输出: (B, embed_dim, H/f, W/f) def __init__(self, in_channels3, channels128, n_res_blocks2, downsample4, embed_dim256): super().__init__() # 初始卷积 layers [nn.Conv2d(in_channels, channels, 3, padding1), nn.GroupNorm(32, channels), nn.SiLU()] # 下采样块控制下采样倍率 cur_channels channels for _ in range(downsample): layers [ nn.Conv2d(cur_channels, cur_channels * 2, 4, stride2, padding1), 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_channels3, channels128, n_res_blocks2, upsample4, embed_dim256): 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_factor2, modenearest), nn.Conv2d(cur_channels, cur_channels // 2, 3, padding1), nn.GroupNorm(32, cur_channels // 2), nn.SiLU(), ] cur_channels // 2 # 输出层 layers.append(nn.Conv2d(cur_channels, out_channels, 3, padding1)) 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_size1024, embed_dim256): 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(dim1, keepdimTrue) # (N, 1) e_sq (self.embeddings.weight ** 2).sum(dim1) # (M,) prod torch.mm(z_flat, self.embeddings.weight.t()) # (N, M) dist z_sq e_sq.unsqueeze(0) - 2 * prod # 取最近码本索引 indices dist.argmin(dim1) # (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网络提取多尺度特征在特征空间计算距离比像素空间的距离更接近人眼的“看着像不像”的判断。安装LPIPSpip install lpips然后封装class PerceptualLoss(nn.Module): def __init__(self): super().__init__() self.loss_fn lpips.LPIPS(netvgg) 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_channels3, channels64): super().__init__() def conv_block(in_c, out_c, stride): return nn.Sequential( nn.Conv2d(in_c, out_c, 4, stridestride, padding1, biasFalse), nn.GroupNorm(8, out_c), nn.LeakyReLU(0.2, inplaceTrue), ) 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, stride1, padding1), ) 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_channelsconfig[in_channels], channelsconfig[channels], n_res_blocksconfig[n_res_blocks], downsampleconfig[downsample], embed_dimconfig[embed_dim], ) self.codebook Codebook( codebook_sizeconfig[codebook_size], embed_dimconfig[embed_dim], ) self.decoder Decoder( out_channelsconfig[in_channels], channelsconfig[channels], n_res_blocksconfig[n_res_blocks], upsampleconfig[downsample], embed_dimconfig[embed_dim], ) self.discriminator Discriminator( in_channelsconfig[in_channels], channelsconfig[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_weight1.0、gan_weight0.1起步。开始训练后观察重建图像如果细节不足就把gan_weight调大如果画面出现伪影就把gan_weight调小。这种调参手感需要积累我后面在避坑部分会再展开。4. 自回归Transformer文本如何变成图像Token序列4.1 问题建模与条件注入方式阶段一训练完VQGAN已经能把图像编码成离散token再完美重建回来。但“文本到图像”还在最后一步让模型学会根据文本生成一串合理的图像token。阶段二用的是自回归Transformer输入是一段文本描述输出是16x16256个图像token。整个过程建模为一个条件语言模型在每个位置预测下一个图像token的概率分布以文本特征作为条件。这和GPT做文本续写的逻辑一模一样只不过词表换成了codebook索引。条件注入我推荐AdaLNAdaptive 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_tensorspt, paddingTrue, truncationTrue) 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_len1024): super().__init__() pe torch.zeros(max_len, d_model) pos torch.arange(0, max_len, dtypetorch.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_size1024, embed_dim256, text_dim512, num_layers8, num_heads8, ff_dim1024, max_seq_len1024): 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_modelembed_dim, nheadnum_heads, dim_feedforwardff_dim, batch_firstTrue, norm_firstTrue, ) 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, devicetoken_ids.device) ) # 文本特征映射 text_cond self.text_proj(text_features) # AdaLN调制参数 # 这里简化处理用文本的平均特征来生成调制参数 text_pooled text_cond.mean(dim1) # (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, memorytext_cond) logits self.output_layer(x) return logits这个实现我简化了AdaLN的插入位置实际标准做法是在每个Transformer子层self-attention之后、FFN之后分别做仿射变换。为了代码可读性我在每个decoder layer前统一调制。如果你追求更精细的控制可以参考DiT的实现方式把scale和shift分别作用到每一层的每个子层上。所幸即便简化版训练效果差别不是特别大。4.4 推理采样如何从概率分布生成图像Token训练完Transformer推理阶段采用自回归方式逐个生成tokendef generate_image_tokens(model, text_features, max_len256, bos_token_id1024, temperature1.0, top_k50): model.eval() device text_features.device # 初始序列只包含BOS tokens torch.full((1, 1), bos_token_id, dtypetorch.long, devicedevice) 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_samples1) # (1, 1) tokens torch.cat([tokens, next_token], dim1) # 遇到结束符可以提前终止可选 # 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_size16): 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, transformNone): 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这些指标量化生成质量除了肉眼观察还可以用FIDFréchet Inception Distance来评估它衡量生成图像分布与真实图像分布的差距FID越低越好。注意FID需要多个样本才能计算单独生成一张图算不出有意义的结果至少准备几千张生成图才有参考价值。7.2 微调技巧如何在特定风格上增强效果如果你想在某个特定风格比如水墨画、赛博朋克、老照片上取得更好的生成效果建议不要从零训练而是用已有的VQGAN权重做微调。具体做法加载通用预训练权重用小学习率1e-5到5e-5在新风格数据上继续训练阶段一注意别把判别器学习率调太高否则会摧毁原来学到的通用表示。阶段二同理在已有Transformer权重上用新风格图文对做微调通常几百到几千张图就能看到明显风格迁移效果。微调时建议冻结VQGAN只更新Transformer这样可以避免风格数据量不够导致的重建退化。7.3 超参数组合参考表我把自己实验过程中的几组代表性超参数整理出来供参考数据集规模分辨率codebook_sizeTransformer层数batch_size预期训练时长(单卡)10万张128x1281024832阶段一约8小时阶段二约5小时10万张256x25610241212阶段一约30小时阶段二约15小时100万张256x256163841632阶段一约3-5天阶段二约2-3天这组数据基于2080Ti或3090级别的显卡不同硬件会有差异。官方论文里用了更大的码本和更深的网络我这里给出的是一个适合个人复现的参考值。8. 从VQGAN延伸出去的路VQGAN作为生成模型家族的承上启下之作复现它给我带来最大的收获不是“跑通了一个模型”而是理解了离散潜在空间这一思想的普适性。后来的图像生成模型里有相当一部分设计都能在VQGAN里找到影子把连续信号离散成token、用Transformer建模序列分布、对抗训练弥补重建模糊、条件注入用自适应归一化参数……这些组件组合起来的范式几乎成了多模态生成模型的通用模板。如果你已经成功跑通这个项目我建议沿着这几个方向继续深挖一是把VQGAN的codebook换成可学习的残差量化或多级量化这会大幅提升重建质量和高频细节二是把Transformer换成非自回归模型比如MaskGIT的思路生成速度能提升一个量级三是把文本编码器从CLIP换成更大的多模态模型条件的语义理解能力会显著增强。VQGAN的代码量和训练难度在生成模型里算中等偏上但胜在组件齐全、思路清晰非常适合作为进入图像生成领域的第一个完整项目。你在复现过程中遇到的问题几乎都能在它的开源社区或论文讨论区找到影子——这本身就说明它的经典程度。动手跑通一次比看十篇论文都管用。最后再分享一个小技巧训练VQGAN这种多组件对抗模型写代码时一定要把每个模块的前向输出shape都打印出来确认无误再拼装。我见过太多同学一遇到维度不匹配就怀疑人生其实只是某个中间层少了一个view操作。把debug的时间花在shape检查上后面会顺畅得多。