ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

DALL-E多模态预训练模型解析:dVAE与Transformer如何协同生成图像

DALL-E多模态预训练模型解析:dVAE与Transformer如何协同生成图像 DALL-E 这个名字第一次出现在我视野里的时候我正陷在一堆 GAN 的训练不稳定问题里出不来。那会儿做图像生成判别器和生成器的博弈就像两个人在拔河稍不留神就一边倒模式崩塌、训练震荡都是家常便饭。所以当我看到 DALL-E 用一套完全不同的思路——把文本和图像塞进同一个 Transformer 里做自回归建模——我第一反应是这条路真的能走通吗后来自己动手复现了简化版本又反复读了它的技术脉络才慢慢理解这套设计背后的取舍。这篇文章就把我对 DALL-E 这套多模态预训练模型的理解完整拆开从它为什么这么设计、dVAE 在里面扮演什么角色、Transformer 怎么同时处理文字和图像到实际复现时容易踩的坑都讲清楚。不管你是刚接触多模态生成的新手还是已经跑过几轮 Transformer 想搞清楚它和视觉任务怎么结合的从业者应该都能从里面拿到点有用的东西。1. 为什么 DALL-E 值得单独拿出来讲1.1 它解决的到底是什么问题文本到图像生成这件事在 DALL-E 之前并不是没人做。早期的方案大多基于 GAN比如 StackGAN、AttnGAN 这类思路是先根据文本生成一个低分辨率草图再逐步放大细化。这些方法能出图但有几个绕不开的毛病一是对复杂文本描述的理解能力有限你让它画一个穿着红色毛衣的柯基犬在滑板上它可能只抓住柯基和滑板颜色和动作就丢了二是生成结果的多样性差同一个 prompt 反复跑出来的图长得都差不多三是训练极不稳定调参成本高得吓人。DALL-E 的核心突破在于换了一套建模范式。它没有用 GAN 那套对抗训练而是把图像生成转化成一个自回归的序列预测问题——本质上和 GPT 生成文本是同一类任务只不过序列里混进了图像 token。这个思路的妙处在于它直接复用了 Transformer 在大规模语言建模上已经被验证过的能力长距离依赖建模、可扩展性、训练稳定性。文本和图像在同一个序列里被统一处理模型自然就学会了跨模态的对应关系。这里有个容易被忽略的点DALL-E 不是文本编码器 图像解码器这种松耦合结构而是把两者真正揉进了一个 Transformer 里。这个设计选择直接决定了它对复杂文本描述的遵循能力。1.2 它和同期方案的本质差异把 DALL-E 和同期其他方案放一起对比差异会更清楚。我整理了一个简单的对照表维度基于 GAN 的方案DALL-E 的自回归方案训练目标对抗博弈判别器与生成器互相拉扯最大似然估计预测下一个 token训练稳定性差容易模式崩塌好损失曲线平滑文本理解通常用独立文本编码器耦合弱文本与图像在同一序列耦合强生成多样性受模式崩塌影响多样性受限采样温度可调多样性可控可扩展性受限于对抗训练的不稳定可随参数量和数据量平滑扩展这张表里最关键的一行是训练目标。自回归建模把生成问题变成了一个监督学习问题每个位置的 token 都有明确的预测目标梯度信号稳定得多。这也是为什么后来一大批多模态生成模型都转向了类似思路。1.3 理解 DALL-E 对今天工作的实际意义可能有人会问DALL-E 是 2021 年初的工作现在模型迭代这么快还有必要回头研究它吗我的看法是非常有必要。原因有三层。第一层它的架构思想是后续很多工作的地基。dVAE 做视觉 token 化、Transformer 做统一序列建模、文本与图像共享注意力这些设计在后来的多模态模型里反复出现只是规模和细节在演进。搞懂了 DALL-E再看后续工作就不是从零开始。第二层它的简化版本非常适合练手。完整复现 DALL-E 需要海量算力但它的核心机制——dVAE 编码、序列拼接、自回归解码——完全可以在小规模数据上跑通。我自己就用一个小型图文数据集复现过虽然出图质量一般但整个流程走下来对多模态建模的理解会深很多。第三层它暴露的问题指明了后续改进方向。比如自回归逐 token 生成速度慢、dVAE 重建有信息损失、文本理解对长描述的覆盖不足这些问题在后来的扩散模型和更大规模模型里被不同程度地解决。知道问题在哪才能理解技术演进的逻辑。2. dVAE把图像变成 Transformer 能吃的 token2.1 为什么图像不能直接喂给 TransformerTransformer 处理的是离散 token 序列。文本天然就是离散的一个词或一个子词就是一个 token。但图像是连续的像素矩阵一张 256×256 的 RGB 图有将近 20 万个数值每个数值还是连续的。如果直接把像素拉平当序列序列长度爆炸不说连续值的建模方式也和 Transformer 的离散 token 预测目标不匹配。所以需要一个中间步骤把图像压缩成一组离散的视觉 token。这个任务就交给了 dVAEdiscrete Variational AutoEncoder离散变分自编码器。它的作用可以类比成图像的 tokenizer——就像文本 tokenizer 把句子切成词元dVAE 把图像切成视觉词元。2.2 dVAE 的编码与解码流程dVAE 的结构分两部分。编码器是一个卷积网络把输入图像下采样并映射到一个 32×32 的网格上每个网格位置输出一个在 8192 个码本向量上的分布。这里 8192 是码本大小32×32 是 token 数量也就是说一张图被压缩成 1024 个视觉 token每个 token 取值在 0 到 8191 之间。解码器反过来接收这 1024 个离散 token通过码本查表还原成连续向量再用卷积网络上采样重建出图像。训练目标是重建损失加上一个对编码分布的约束项让分布尽量接近均匀避免码本坍缩——也就是所有 token 都挤到少数几个码本向量上。# dVAE 编码器的简化结构示意PyTorch 风格 import torch import torch.nn as nn class DVAEEncoder(nn.Module): def __init__(self, vocab_size8192, token_dim256): super().__init__() self.conv nn.Sequential( nn.Conv2d(3, 64, 4, stride2, padding1), # 256 - 128 nn.ReLU(), nn.Conv2d(64, 128, 4, stride2, padding1), # 128 - 64 nn.ReLU(), nn.Conv2d(128, 256, 4, stride2, padding1),# 64 - 32 nn.ReLU(), ) # 输出 32x32 网格每个位置 vocab_size 维 logits self.head nn.Conv2d(256, vocab_size, 1) def forward(self, x): feat self.conv(x) # [B, 256, 32, 32] logits self.head(feat) # [B, 8192, 32, 32] return logits这段代码只是结构示意实际 dVAE 里还有码本、Gumbel-Softmax 采样、温度退火等细节。但核心逻辑就是卷积下采样到固定网格每个网格位置预测一个离散分布。2.3 码本坍缩dVAE 训练里最容易翻车的地方我复现 dVAE 时踩得最狠的坑就是码本坍缩。现象是训练一段时间后编码器输出的分布越来越尖锐几乎所有图像都被映射到同一小撮码本向量上重建出来的图全是模糊的平均色块。根因在于离散采样本身不可导需要用 Gumbel-Softmax 做重参数化而温度参数控制着分布的平滑程度。温度太高采样接近均匀学不到有意义的码本温度太低分布过早尖锐化直接坍缩。我的经验是温度退火要慢从较高的初始温度比如 1.0逐步降到较低值比如 0.1退火周期要覆盖足够多的训练步数。另外编码分布上加一个 KL 约束项鼓励它接近均匀分布也能有效缓解坍缩。实操提示监控码本使用率是个好习惯。每隔若干步统计一下有多少个码本向量被用到了如果使用率持续下降基本就是坍缩的前兆该调温度或加正则了。2.4 dVAE 重建质量对最终生成的影响dVAE 是整条链路的第一环它的重建质量直接决定了最终生成图像的上限。道理很简单Transformer 预测的是视觉 token如果 dVAE 本身重建就丢了很多细节那 Transformer 预测得再准解码出来的图也好不到哪去。我在小数据集上做过对比dVAE 重建损失降不下去的时候后面 Transformer 训练得再久出图都是糊的。反过来先把 dVAE 训到重建清晰、码本使用充分再训 Transformer效果提升立竿见影。所以不要急着上 Transformer先把 dVAE 调好这是顺序问题不是并行问题。3. Transformer 如何同时消化文字和图像3.1 序列是怎么拼起来的DALL-E 的输入序列由三部分组成文本 token、图像 token、以及一个分隔符。文本部分用 BPE 编码最多 256 个 token图像部分就是 dVAE 产出的 1024 个视觉 token中间用一个特殊 token 隔开。整个序列长度大约 1280 左右。模型训练时做的是标准的自回归语言建模给定前面的 token预测下一个 token。文本部分的 token 用文本词表图像部分的 token 用视觉词表两个词表是独立的但共享同一个 Transformer 主体。注意力机制让每个位置都能看到它之前的所有位置所以图像 token 在生成时能看到完整的文本描述文本和图像之间的对应关系就是在这一步学到的。# 序列拼接与自回归预测的简化示意 def build_sequence(text_tokens, image_tokens, sep_id): # text_tokens: [B, L_text] # image_tokens: [B, 1024] seq torch.cat([text_tokens, torch.full((text_tokens.size(0), 1), sep_id), image_tokens], dim1) return seq # 训练时图像部分的每个 token 都作为预测目标 # 损失只在图像 token 位置计算文本部分是条件不计算损失这里有个细节值得说损失通常只在图像 token 位置计算文本部分是作为条件输入的不参与预测。这符合任务目标——我们要的是根据文本生成图像不是根据图像生成文本。3.2 注意力掩码与位置编码的处理自回归建模要求因果掩码也就是每个位置只能看到自己和之前的位置。这个在标准 Transformer 解码器里是标配DALL-E 直接沿用。但多模态序列有个特殊之处文本在前、图像在后所以图像 token 天然能看到全部文本而文本 token 看不到图像。这个顺序不能反反了就成了根据图像生成文本任务就变了。位置编码方面DALL-E 用的是可学习的位置嵌入文本和图像共享同一套位置编码空间。这里有个容易困惑的点文本长度可变图像长度固定 1024位置编码怎么对齐答案是位置编码按序列中的绝对位置分配文本占前若干位分隔符占一位图像占后面的固定区间。因为图像总是从固定位置开始所以图像 token 的位置编码是稳定的。3.3 文本理解能力的来源DALL-E 对文本的理解能力不是来自某个专门的文本编码器而是来自大规模图文配对数据上的联合训练。模型见过海量的描述-图像对在预测图像 token 的过程中它被迫学会把文本里的语义信息映射到视觉特征上。比如红色这个词反复出现在红色物体的描述里模型就逐渐把红色这个 token 和图像中偏红的视觉 token 关联起来。这种学习方式是隐式的、分布式的不像规则系统那样有明确的对应表。好处是泛化能力强能处理训练时没见过的组合坏处是可解释性差你很难说清楚模型到底理解了多少。我在测试时发现对于训练分布内的描述模型遵循度很高对于特别长或特别抽象的描述遵循度会明显下降这跟训练数据的覆盖范围直接相关。3.4 采样策略对生成结果的影响训练完之后生成图像靠的是从模型输出的分布里采样。采样策略直接决定了生成结果的多样性和质量。最基础的是贪心解码每步取概率最大的 token结果稳定但死板同一个 prompt 每次出图都一样。实践中更常用的是温度采样和 top-k 采样。温度参数控制分布的平滑程度。温度趋近 0 就退化成贪心温度升高分布变平采样更随机多样性上来了但可能出乱图。我的经验是温度在 0.7 到 1.0 之间比较平衡。top-k 采样是每步只在概率最高的 k 个 token 里采样过滤掉长尾的低概率 token能有效避免生成明显不合理的视觉 token。k 取 32 到 256 都有人用具体看任务对多样性和质量的偏好。一个实用技巧如果发现生成图整体偏糊先别急着怪模型检查一下是不是采样温度太低导致模式单一。适当升温往往能让细节丰富起来。4. 复现 DALL-E 核心机制时的实操路径4.1 数据准备图文配对数据的清洗要点复现的第一步是搞到干净的图文配对数据。公开数据集里图文对的质量参差不齐很多描述和图像对不上或者描述过于简单比如就一个词。这些脏数据会直接污染训练信号。我的清洗流程是这样的先过滤掉描述长度过短的样本比如少于 5 个词的直接扔再用一个预训练的图文匹配模型给每对打一个相似度分低于阈值的剔除最后人工抽检一批看看剩下的质量。这一步很费时间但省不得。我试过跳过清洗直接训结果模型学出来的东西乱七八糟文本和图像基本对不上白跑了好几天。数据量方面小规模复现用几万到十几万对就能看到初步效果当然要出好图还是得更大规模。关键是质量优先于数量一万对干净数据比十万对脏数据有用得多。4.2 分阶段训练先 dVAE 后 Transformer训练顺序我前面提过这里展开说。第一阶段单独训 dVAE目标是重建损失降到合理水平码本使用率保持健康。这个阶段可以用纯图像数据不需要文本。第二阶段冻结 dVAE训 Transformer让它学会从文本预测视觉 token。为什么要冻结 dVAE因为如果两个一起训dVAE 的码本会随着 Transformer 的梯度一起变导致视觉 token 的语义不稳定Transformer 刚学到的对应关系下一轮就失效了。分开训dVAE 提供一个稳定的 token 空间Transformer 在这个空间里学习映射收敛快得多。# 训练流程的伪命令示意 # 阶段一训练 dVAE python train_dvae.py --data images/ --epochs 50 --lr 1e-4 # 阶段二冻结 dVAE训练 Transformer python train_transformer.py \ --dvae_ckpt dvae_best.pt \ --freeze_dvae \ --data image_text_pairs/ \ --epochs 100 \ --lr 3e-44.3 显存与算力的现实约束完整 DALL-E 有 120 亿参数普通人根本跑不动。但复现核心机制不需要那么大。我的做法是把 Transformer 缩到几千万参数dVAE 也相应缩小图像分辨率降到 64×64 或 128×128token 数量从 1024 降到 256 或 64。这样单卡就能跑起来。显存瓶颈主要在注意力计算上序列长度 1280 的注意力矩阵是 1280×1280还好。真正吃显存的是 batch size 和模型宽度。我的经验是先用小 batch、小模型把流程跑通确认每个环节都对再逐步放大。一上来就追求大模型大概率是卡在显存报错上反复折腾浪费时间。4.4 评估生成效果不能只看图好不好看评估多模态生成模型比评估纯图像模型麻烦因为要同时考虑图像质量和文本遵循度。我一般从三个角度看。第一是重建保真度把 dVAE 单独拎出来看它重建原图的能力这是上限。第二是文本遵循度给定一组描述看生成图是否包含描述里的关键元素比如颜色、物体、数量。第三是多样性同一个描述多次采样看结果是否有变化全一样说明采样太死板。定量指标可以用 FID 衡量图像分布距离但 FID 不反映文本遵循度所以得配合人工评估。我通常会固定一批 prompt生成后自己一张张看记录哪些描述遵循得好、哪些丢了信息这样能定位到是文本理解的问题还是视觉生成的问题。5. 那些文档里不会写的踩坑记录5.1 视觉 token 顺序错乱导致生成图结构崩坏这个问题我卡了挺久。现象是生成的图局部看还行但整体结构完全乱物体位置、比例都不对。排查了半天最后发现是 dVAE 输出的 token 网格在拉平成序列时行列顺序和位置编码的对应关系搞错了。dVAE 输出的是 32×32 的二维网格拉平成一维序列时如果按列优先而不是行优先位置编码就对不上模型学到的空间关系全是错的。修复很简单确认拉平顺序和位置编码分配一致就行。但这个坑隐蔽在于训练损失看起来在正常下降因为模型总能学到某种映射只是这个映射和真实空间结构不符。教训是涉及空间结构的张量变形一定要打印形状和几个具体位置的值来核对别想当然。5.2 文本 token 与图像 token 的词表冲突DALL-E 里文本和图像用独立词表但共享 Transformer 的嵌入层和输出层时如果实现不当两个词表的索引会重叠导致模型分不清某个 token 是文本还是图像。我一开始图省事把两个词表拼成一个大词表结果文本 token 和图像 token 的 id 范围没隔开模型在预测时经常把图像位置预测成文本 token。正确做法是给两个词表分配不重叠的 id 区间或者在嵌入层和输出层做区分。我后来改成文本 id 从 0 开始图像 id 从一个大偏移量开始问题就没了。这个坑的教训是多模态模型里不同模态的 token 空间一定要显式隔离别指望模型自己学会区分。5.3 训练后期损失突然飙升的处理训练到中后期损失突然从平稳状态飙升然后慢慢恢复过一阵又飙升。这种周期性震荡我遇到过好几次。排查下来原因通常是学习率在后期相对当前损失尺度偏大加上自回归任务里长序列的梯度累积效应导致参数更新过猛。解决办法有几个一是加梯度裁剪把梯度范数限制在一个阈值内二是用学习率预热加余弦退火后期学习率降下来三是在损失飙升时自动降低学习率。我一般组合使用梯度裁剪阈值设 1.0学习率退火到初始值的十分之一。这样处理后训练曲线平滑很多。5.4 采样时重复 token 的抑制自回归生成有个通病就是容易陷入重复循环生成一串相同的视觉 token解码出来就是一片纯色区域。文本生成里常用重复惩罚来解决图像 token 生成同样适用。具体做法是在采样时对已经出现过的 token 降低其概率或者用 top-k 加温度的组合来打散分布。我试过纯 top-k 不加温度重复问题还是有加上温度后明显改善。另外重复惩罚的系数不能太大太大会导致生成过于发散图像结构散架。系数在 1.1 到 1.3 之间比较稳妥。6. 从 DALL-E 延伸出去的技术脉络6.1 自回归路线的天花板在哪DALL-E 走的是自回归路线逐 token 生成。这条路线的优势是训练稳定、理论清晰但天花板也很明显生成速度慢。一张图 1024 个 token每个 token 都要跑一次前向生成一张图要上千次前向传播实际用起来延迟很高。而且自回归是串行的没法并行加速。另外自回归建模对全局结构的把控偏弱。它是一步步往外吐 token局部连贯性好但全局布局容易出问题比如画一个人可能头画得挺好身体比例就崩了。这个问题的根源在于生成早期 token 时模型还没想好整体构图后面只能将错就错。6.2 后续方案怎么解决这些问题后来的扩散模型路线比如基于扩散的文本到图像生成换了一套完全不同的思路不是逐 token 预测而是从噪声出发迭代去噪。这个过程可以并行处理所有位置生成速度快很多而且全局一致性更好因为每一步都在看整张图。但扩散模型也有自己的问题比如训练和采样的数学更复杂对噪声调度的设计很敏感。所以技术演进不是简单的替代关系而是不同路线在速度、质量、稳定性之间做不同的权衡。理解 DALL-E 的自回归思路能帮你看清楚这些权衡的来龙去脉。6.3 对今天做多模态应用的启发如果你现在要做多模态相关的应用DALL-E 给我的最大启发是统一序列建模的思路是通用的。把不同模态的数据映射到统一的 token 空间用一个 Transformer 处理这个范式在很多任务上都成立不限于图像生成。比如图文检索、视觉问答、多模态对话都可以借鉴这个思路。另一个启发是分阶段训练的价值。先把各模态的 tokenizer 训好再训跨模态的联合模型这个流程比端到端一起训稳定得多。我在做其他多模态任务时也沿用这个套路效果确实更可控。最后说个我自己的体会。DALL-E 这套东西看论文觉得逻辑挺顺真动手复现才知道细节里全是坑。dVAE 的码本、序列的拼接顺序、词表的隔离、采样的参数每一个环节出问题都会让最终结果差很多。但正是这些坑逼着我把多模态建模的每个环节都想清楚。如果你也在做类似的东西我的建议是别怕从小规模开始把流程跑通比追求大模型重要得多。跑通之后你对这套机制的理解会比读十篇论文都深。
返回列表