摘要
我们提出 Generative Infinite-Vocabulary Transformers(GIVT):它生成由实值元素构成的向量序列,而不是有限词表中的离散 token。为此,我们对 decoder-only Transformer 做两项非常简单的修改:输入端用输入向量的线性投影取代有限词表查找表;输出端用多元 Gaussian mixture model 的参数取代通常映射到 categorical distribution 的 logits。受 VQ-GAN 与 MaskGIT 图像生成范式启发——其中 Transformer 对 VQ-VAE 的离散潜在序列建模——我们让 GIVT 对 β-VAE 未量化的实值潜在序列建模。在 class-conditional 图像生成中,GIVT 超过 VQ-GAN 及其改进版本和 MaskGIT,表现可与近期 latent diffusion model 竞争。将 GIVT 用于 UViM 框架的 VAE 变体时,它在 panoptic segmentation 与 depth estimation 上也取得了很强结果。
引言
在引入后不久,变换器成为自然语言处理中的主导架构,[72]最近在计算机视觉中也变得非常流行。[18, 63, 40]Dosovitskiy等人..[18]表明,通过将 图像分割成补丁序列,线性嵌入这些补丁,并 随后将所得特征序列馈送到Transformer编码器 ,可以在大型模型和数据规模下产生优于基于CNN的 架构的强大图像分类器。这一策略现在 是许多判别性视觉任务的标准,包括 分类[18]、检测[40],以及 分割[63]。如何将 生成式Transformer解码器应用于图像则不那么明显。生成,因为 它们被设计为消费和预测来自某个 固定、有限词汇表的离散标记。这种结构自然适合自然 语言,对于语言来说,仅解码器模型能够实现强大的序列 生成建模和高效训练[72, 52].
为了将这些能力用于图像,最近的工作[54, 20, 7, 6, 39, 46]采用了两个阶段的方法 首先训练一个向量量化变分自编码器 (VQ-VAE)[49]将图像映射为离散标记序列,然后训练一个 变换器解码器来建模潜在离散标记分布。这种基于VQ-VAE的图像标记化的一个优势是,它能够 实现交错的多模态生成模型,只需将不同模态(包括文本和图像)的词汇表拼接起来即可[1, 29, 2]。然而,这种方法也存在几个问题。 首先,VQ的非连续性要求可微分的 近似来实现基于随机梯度的优化[49]. 其次,词汇量较小的VQ-VAE虽然使潜在建模变得容易,但也使得潜在编码信息量不足,这阻碍了对图像生成中低层细节的控制,并在使用这些标记进行密集预测[33, 42]或低层判别任务[1, 29]时影响质量。另一方面,词汇量过大可能导致词汇利用率低[46],因此高保真VQ-VAE设置通常依赖一系列先进技术,如熵损失[7]或码本分割[33]。此外,大词汇量会导致相应的嵌入矩阵变大,从而增加内存消耗,这在多模态环境中尤其是个问题。
在这项工作中,我们首次(据我们所知)展示了如何完全移除量化从用于视觉数据的生成式Transformer中。事实上,从业者似乎普遍认为这几乎不可能,因为Transformer解码器在许多方面与离散表示紧密相关。令人惊讶的是,我们不仅展示了简单的修改能使Transformer解码器直接生成未量化向量的序列,而且这种方法在图像生成质量和表示学习能力上优于基于VQ的方法。我们称这种 变换器为生成式无限词汇 变换器(GIVT)。1具体来说,与标准变换器解码器架构相比,我们做了两个改变[72, 52], 参见2:1)在输入方面,GIVT不是使用离散标记序列来查找有限词汇的嵌入,而是线性嵌入实数向量序列;2)在输出方面,GIVT不是预测有限词汇上的类别分布,而是预测元高斯混合模型(GMM)的参数。 我们以与标准变换器解码器相同的方式训练GIVT:使用因果注意力掩码和教师强制[72],并且还探索了如MaskGIT中的快速渐进掩码双向建模[13, 7, 6].
与使用VQ-VAE的两阶段方法类似,并且类似于潜在扩散模型的两阶段方法[55, 51],我们首先使用高斯先验-VAE[30, 24]学习一个较低维度的潜空间,然后使用GIVT对其进行建模。我们强调,训练-VAE和GIVT仅依赖于深度学习工具箱中的标准技术,而不是VQ-VAE文献中的高级训练技术,如辅助损失[49, 7]在潜在表示、码本重新初始化[37]或专用优化算法上[33, 27].
我们的主要贡献可总结如下:
我们展示了GIVT在类条件图像生成中优于VQGAN[55](及其后续变体)和MaskGIT[7],通常以较大幅度和/或显著更低的计算成本实现。GIVT还与强大的潜在扩散基线竞争,尤其是在高分辨率下。
我们为连续情况推导了标准采样技术的变体,如温度采样、束搜索和无分类器引导(CFG)[25],并展示了它们的有效性。
我们证明了GIVT在表示学习方面以显著更低的计算成本匹配或超越了先前的序列图像生成模型。
GIVT在性能上与基于VQ的UViM方法相当[33]在密集预测任务中,如语义分割和单目深度估计。
我们强调,基于Transformer解码器的视觉数据生成模型(如GIVT)的进展直接受益于大型语言模型在扩展和推理效率方面的进步。相反,与扩散模型不同,我们这类模型的改进可以轻松迁移到多模态交错建模,[1, 29, 2]这正变得越来越流行。
相关工作
用于视觉数据标记化的VQ-VAE在像素空间自回归建模成功之后,[71, 59, 50, 43, 8]将自回归建模迁移到VQ-VAE的潜在空间[49, 54]成为一种更高效的替代方案。使用GAN和感知损失进行VQ-VAE训练,以及现代因果[20, 77, 73]和掩码[7, 6, 39]变换器进行潜在建模,带来了显著的质量提升。另一个利用VQ-VAE的活跃领域是图像和文本的交错多模态生成建模[1, 29, 2]. 此外,VQ-VAE 是对密集预测视觉任务的标签空间进行标记化的流行选择。[33, 42]. 最后,一些受语言启发的图像自监督学习技术依赖于 VQ-VAE 表示。[3, 75, 39].
分布的离散混合将分类分布的 logits 的密集预测替换为连续混合模型,随后进行离散化。这种方法在[59]中提出,用于像素空间自回归建模,以减少模型参数数量并提高学习效率,在神经压缩中也很流行。[45, 10, 44].
NLP 中的连续输出处理机器翻译中大型词汇表的一种流行方法是通过连续分布预测语言标记的词嵌入,而不是通过分类分布预测标记 ID。[35, 34, 64, 65, 38]解码通常以贪婪方式进行嵌入查找,因此不会产生多样化的样本。此外,模型消耗并预测来自固定有限集合的词嵌入。
具有学习先验的 VAE大量文献研究了通过学习先验改进 VAE:逆自回归流成为一种流行的选择。[31, 9]. 其他方法使用归一化流[70]或变分后验与伪输入的混合[66]。对于具有离散(非VQ)潜变量的VAE,基于受限玻尔兹曼机的学习先验已被研究[69, 57].
使用Transformer进行时间序列建模最近有各种工作探索了用于时间序列建模/预测的Transformer。这些工作要么使用回归损失[79, 48, 36, 22, 11],要么使用分位数预测[19, 41],或者采用数据离散化/分箱[53]。有些相关的是,[47, 74]从离散标记回归连续语音特征。这些模型都不像GIVT那样预测连续分布,从而允许自回归生成。
生成式无限词汇Transformer
如第1节所述,我们的方法在概念上与近期在VQ-VAE[20, 76, 7, 6]的离散编码上训练仅解码器Transformer模型的工作相似,关键区别在于我们不进行量化(即,不使用VQ)。现在我们描述我们方法的组成部分。
VAE训练
我们首先训练一个连续潜变量 -VAE[24],其编码器和先验为高斯分布,最初由[30]. 给定输入图像,编码器预测多元正态分布的均值和协方差(对角协方差矩阵),并使用重参数化技巧从中采样表示[30]。然后,VAE解码器将潜在序列映射回图像。由于我们使用高斯编码器分布,证据下界(ELBO)[30]中的KL项可以按照[30, 第F.1节] 中描述的闭式形式计算。. 至于 ELBO中的重建/似然项,我们依赖于MSE、 感知损失和GAN损失的混合用于图像生成,遵循[20, 7],或用于密集预测任务的 分类交叉熵[33]。我们的编码器 在空间上对进行下采样,由此我们 得到,其空间维度为,特征维度为,其中,给定一个输入。为了计算KL项,相关的和形状为被展平为向量。
超参数乘以KL项控制着被正则化的强度。正如我们将在 第5节中看到的,这种对VAE的正则化对于 能够很好地建模由此产生的(真实的)潜在分布是重要的。
GIVT训练
接下来,我们训练一个GIVT来预测或(当条件信号可用时,例如,在 类条件生成中)。表示被重塑为长度为的维实值向量(或“软标记”)。请注意,这与标准的 VQ-VAE设置不同,在标准设置中,潜在变换器解码器建模一个-长度的序列,其元素为整数表示码本索引。为适应这一差异,我们对标准的仅解码器变换器架构做了两处 小改动(见图2):在输入处,我们用单个线性层替换嵌入 查找表,以从投影到变换器的隐藏 维度。在输出处,我们不预测类别分布, 而是让变换器预测连续分布的参数。假设混合 分量的通道间独立性,我们用-混合高斯模型(GMM)来建模该连续分布。因此,GIVT模型对每个 软标记预测个参数(均值和混合成分的方差参数,以及混合概率)。实验上,我们发现用softmax激活函数归一化混合概率,用softplus归一化方差参数是有益的。
我们在GIVT预测的分布上使用标准的交叉熵损失(等同于负对数似然),并最小化,假设类别或条件信号均匀分布(详见 附录C关于损失的细节)。我们训练两种类型的GIVT模型,如下所述。
GIVT-Causal。这里,GIVT 在长度为 的潜变量序列中预测每个 维向量,条件为此前全部向量。因此,self-attention 层采用时间因果 mask[20, 72](这使模型能在推理时顺序生成,与 causal inference 无关)。这种训练策略也称为 teacher forcing,与 VQ-GAN 中的潜变量建模类似[20]。对于类别条件图像生成,我们在输入序列前添加一个 [CLS] 向量,即为每个类别 学习一个向量。
GIVT-MaskGIT与MaskGIT一样[7],我们在训练时随机掩蔽输入序列的一个子集,然后在推理时逐步揭示被掩蔽的标记。与相比,唯一的变化是[7]与我们的实值标记相关:由于我们有无限多个标记,因此没有
明显的方法来定义特殊的掩码标记(当使用VQ时,可以
直接扩展词汇表以包含特殊标记,例如[MASK])。相反,给定和一个掩码指示每个位置是否
被掩码,我们首先将中对应于的位置替换为零(以去除信息),然后
用单个密集层对其进行嵌入,如上所述。此外,我们拼接在特征维度上两个学习到的特殊向量之一,一个[MASK]向量用于掩码位置,另一个[UNMASK]向量用于其他情况(我们减半嵌入输入和特殊标记的维度,使得最终隐藏维度
保持不变)。
迈向端到端训练:适配器
使用未量化的VAE并将生成的潜在序列用连续而非分类分布建模的一个有趣结果是,VAE和GIVT可以联合训练或端到端微调(使用重参数化技巧[30])。然而,这种设置带来了其自身的挑战(例如,它包含多个必须适当平衡的损失),我们将其留作未来工作。相反,我们探索了一种简单的替代方案,以更好地匹配VAE的潜在分布和GIVT预测的分布:我们使用一个小型可逆流模型[15, 16],或称为“适配器”,将VAE潜在序列映射到具有相同维度的新潜在空间。我们依赖基于“体积保持”加性耦合层的模型,其具有对角雅可比矩阵[15]。然后,GIVT与适配器联合训练,以预测适配器诱导的变换潜在空间中的序列(使用相同的损失)。在推理时,从GIVT中抽取的样本首先由逆适配器处理,然后由VAE解码器解码为图像。注意,由于可逆性,适配器不需要额外的损失,并且与GIVT模型相比,增加了可忽略的计算和模型参数开销(小于0.1%)(详见第4节和附录B节)。
推理
给定如上训练的VAE和GIVT,在推理过程中,我们按顺序从GIVT中采样(见图3)或如MaskGIT[7]中那样,并将采样序列解码为图像。我们现在研究离散Transformer的各种推理方案,并推导出它们的连续对应方案。
温度采样、核采样、束搜索在文本序列模型(参见[26]的概述和讨论)以及基于VQ-GAN的方法中,通常需要调整和优化采样算法。我们从温度采样开始,对于离散模型,它调整每个解码步骤预测的分类分布的softmax温度。对于GIVT,我们改为缩放预测的高斯分布的协方差矩阵,并将此策略称为“方差缩放”。正如我们将在4中看到的,这个简单的改变可能对样本质量产生显著影响。
核采样[26]提出收集最大的logits,使得归一化后的累积概率超过阈值(例如0.8),并从这种减少支持的分布中采样。在GIVT中,当预测单个混合时,这可以通过截断每个维度的预测分布来近似(从而选择更高密度的支持)。这与方差缩放有类似的效果,因此我们不采用此策略。
我们还考虑束搜索,这与离散Transformer解码器中的束搜索相同。对于每个样本,我们维护光束,并且在每一步中,我们为每个光束采样多个候选(这里称为“扇出”)。然后,我们计算所有光束和扇出到当前采样步骤为止的累积对数概率,并选择具有最高累积对数概率的个光束。最后,在GIVT中没有与top-k采样[21]类似的概念,因为它预测的是连续分布。
基于分布的免分类器引导在扩散文献中,免分类器引导(CFG)[25]已被成功应用。具体来说,条件扩散模型通过额外的空类来学习无条件数据分布。然后,在推理过程中,条件对数密度被“移离”无条件对数密度:给定引导权重,更新后的(扩散)分数估计为其中估计数据分布的对数密度的梯度,(参见[25, 第2节])。由此,我们现在为我们的GIVT推导出一个CFG变体,因为我们直接预测一个密度。我们将这种方法称为“基于密度的CFG”(DB-CFG)。公式1可以写成 即.., 估计密度的对数(见图6)。因此,我们 希望调整我们的模型以从中采样。我们遵循[25]并训练 GIVT,增加一个额外的空类。在推理过程中,我们每一步 评估GIVT两次,一次以实际标签为条件,一次以为条件。为了实现无分类器 引导,我们随后必须从由两个GIVT 预测导出的未归一化版本中采样。为此,我们转向拒绝采样,这需要: 1)一个未归一化的密度;2)一个好的提议分布, 这接近真实的目标分布;以及3)一个缩放因子用于限制与未归一化目标密度之间的似然比。
公式(1)。
公式(2)。
我们混合的分布是GMM,找到好的提议分布可能具有挑战性。相反,我们首先从中采样混合索引,并对和中相应的混合分量应用DB-CFG(这些分量是具有对角协方差的多元高斯分布)。我们经验发现,无条件分量(即.., 分布 使用标签) 预测的分布往往比条件分布具有更大的方差(如 图6所示)。因此,合理的做法是从中选取样本提议,其中是 GIVT 在给定标签时预测的参数。我们凭经验发现,抽取 1000 个样本足以在 99.9% 的情况下找到至少一个有效样本。对于剩余的%,则回退到从.
我们强调,DB-CFG 的额外开销很小:每个推理步只需两次而不是一次前向传播,分别预测条件分布与无条件分布。随后,我们在加速器上从这些分布并行抽取 1000 个样本,速度很快。关于 Python 代码,参见附录 3.4。
实验
图像生成
我们使用ImageNet1k[56]并探索类别条件生成(其中我们将GIVT条件于类别标签)用于256px和512px,以及无条件生成用于256px。
-VAE我们 紧密遵循MaskGIT的设置[7]。我们使用其VAE架构, 由ResBlocks构建(详见附录C;编码器和解码器共 有53.5M参数),移除VQ层及相关损失, 并将其替换为预测的线性层(第3.1节)。我们使用与[7, 46]中相同的重建、感知和GAN损失权重,以及相同的优化器参数; 我们仅改变潜在维度和KL项的权重。默认情况下,我们将token维度设置为 (即..,VAE预测每个token的16个 均值和方差)和。我们注意到,我们的VAE是在张图像上训练的,并且我们也将其用于我们的实验,而无需重新训练(如[7]).
GIVT。对于GIVT-Causal,我们遵循原始的Transformer解码器架构[72]在仅解码器模式下,但 从注意力层、MLP块和LayerNorm中移除偏置,并 按照常见做法将ReLU替换为GELU。对于GIVT-MaskGIT,我们简单 地在训练期间移除注意力掩码,并输入掩码输入而不是 移位输入。我们使用BERT-Large配置[13]默认情况下 (304M参数),并探索更大的主干网络,具有1.67B 参数,用后缀“-L”表示(详见附录B)。对于带有适配器的模型变体(后缀“+A”),我们使用8个双射 iRevNet块的堆叠[28](隐藏通道 维度为,导致额外的112k 参数用于), 应用于表示,然后再将其重塑为序列。我们将GIVT模型配置为预测-混合 GMM,其中混合分量是因子化的(即混合分量是 具有对角协方差的多变量高斯分布),并探索预测 一个单一的多变量高斯分布,建模token的完整协方差矩阵作为替代方案。对于条件生成 实验,我们使用一个学习到的嵌入,将其前置到嵌入的 token序列之前。为了训练GIVT,我们使用带有余弦调度(500 个epoch;线性预热50个epoch)的Adam优化器,将学习率和 权重衰减设置为和,分别,优化器的参数设置为,dropout概率设置为对于GIVT-causal和对于GIVT-MaskGIT,以及批量大小 设置为。我们使用与VAE训练期间相同的数据 增强(参见[7, 46]),并且对于每个批次从VAE编码器分布中采样(除了数据增强之外的另一个随机性来源)。
我们在JAX中实现GIVT[4]并使用distrax[12]来实现候选并计算对数概率。
GIVT-MaskGIT推理遵循[7],我们将推理步数固定为16,并采用余弦调度(即.. 让在步骤,掩码 标记的比例由)给出。 我们还在每一步按似然对标记进行排序,并使用 “选择温度”.
探索VAE潜在空间为了更好地理解 特征维度、KL正则化,VAE的重建质量,以及GIVT的采样质量,我们训练了潜在维度为和的VAE,使用本节开头描述的VAE训练设置。对于每个得到的VAE,我们训练了一个GIVT-Causal,使用较小的BERT-Base维度,以及一系列混合数[13]的值。.
评估对于VAE,我们报告“重建 FID”,即重建50k张ImageNet验证图像时获得的FID。对于我们的GIVT变体和基线,我们报告采样FID[23],即 采样覆盖所有ImageNet类别的50k张平衡图像集时的FID。在 两种情况下,我们都依赖成熟的ADM TensorFlow套件[14],该套件 使用整个ImageNet训练集作为参考。此外,我们 还报告精确率和召回率[58]。最后,我们通过 在无条件GIVT-Causal的平均池化中间表示上训练线性分类器来评估 表示学习能力,如先前工作[8, 76](详见附录F)。
全景 分割和深度估计
我们基于UViM框架[33],该框架使用VQ-VAE来 压缩计算机视觉密集预测任务的标签空间,并使用一个编码器-解码器Transformer,以RGB图像作为输入, 在VQ-VAE潜在空间中将相关的密集标签预测为离散代码。在这里,我们将VQ-VAE替换为一个-VAE并使用GIVT编码器-解码器 来建模连续潜码。对于VAE,我们使用与[33]相同的基于Transformer的自编码器架构(6层编码器和12层 解码器)和交叉熵损失。我们设置, ,以及KL权重用于全景 分割,用于深度估计。为了以与[33]相同的方式构建编码器-解码器GIVT 模型,我们采用为ImageNet生成所描述的因果 变体,并在每个自注意力层之后插入一个交叉注意力层。遵循[33],我们使用来自[62]的ImageNet-21k预训练ViT-L/16作为编码器,将图像 分辨率设置为512像素,并采用来自[33]的预处理和优化 超参数。我们使用没有编码器和解码器上下文的UViM变体[33]。最后,我们考虑 方差缩放和束搜索,在训练集的保留子集上选择参数,如[33].
结果
图像生成
VAE潜在空间在图4中,我们展示了改变KL项的权重如何影响1) VAE 重建FID和2) 在相应潜在序列上训练的Base-size GIVT-Causal的采样FID。对于1),增加会导致重建FID变差 ,因为VAE在潜在空间中存储的信息更少。它将更多的建模工作转移到VAE解码器上,使得解码器逐渐更具生成性,这会影响采样质量(参见[67, 第 7节][55]以获取更多讨论)。
对于2),我们看到了相反的趋势:增加会导致在潜在序列上训练的GIVT模型的采样FID降低(更好)。可以说, 这是因为VAE潜在序列更接近高斯先验,因此更容易被GIVT建模。最后,增加混合数量最初显著降低了采样FID,在处达到平台期。因此我们默认设置和,并对较大的(-L)GIVT模型使用带有的VAE。我们强调,在文献中探索和调整诸如或类似地VQ中的词汇表大小和承诺损失等超参数是常见的[20],表5 [55],表8[51][46,图3].
采样FID在表1我们展示了在类别条件下的四种模型类别的采样FID。ImageNet:GAN、 基于扩散的方法,以及掩码和序列建模 方法。GIVT-MaskGIT优于MaskGIT[7]后者具有相当的模型大小 和推理成本,而DB-CFG带来了额外的改进。在 没有引导技术的情况下,我们的GIVT-Causal模型大幅优于所有 扩散基线以及VQGAN。使用引导 技术,GIVT-Causal获得了3.35的FID,而VQGAN为5.20,且计算量超过更小的 模型(GIVT为0.3B参数,而对比模型为1.4B参数),并且也优于32% 更大的LDM-4-G。我们更大的GIVT-Causal-L+A在无引导和有引导情况下,分别比ViT-VQGAN获得16%和17%的FID降低,而ViT-VQGAN具有相同的生成Transformer大小但更长的序列长度(导致超过较慢的 采样)以及一个更大的 VAE。
我们展示了采样FID用于表2中的 ImageNet。GIVT-MaskGIT 在模型大小和推理成本相当的情况下,FID 比 MaskGIT 低 38%。GIVT-Causal-L+A 在无引导和有引导的情况下均优于 DiT-XL/2(目前最佳的 DiT 模型)(尽管模型更大)。
最后,我们在附录中展示了无E条件结果。这项任务难度显著 更高,但 GIVT-Causal 大幅超越了基于扩散的 ADM[14]。
| 类别 | 模型 | 推理 | Steps | FID | Precision | Recall |
|---|---|---|---|---|---|---|
| GAN | BigGAN-deep [5] | 6.95 | 0.87 | 0.28 | ||
| StyleGAN-XL [60] | 2.30 | 0.78 | 0.53 | |||
| Diffusion | ADM [14] | 250 | 10.94 | 0.69 | 0.63 | |
| ADM-G [14] | CG | 250 | 4.59 | 0.82 | 0.52 | |
| LDM-4 [55] | 250 | 10.56 | 0.71 | 0.62 | ||
| LDM-4-G [55] | CFG | 250 | 3.60 | 0.87 | 0.48 | |
| DiT-XL/2 [51] | 250 | 9.62 | 0.67 | 0.67 | ||
| DiT-XL/2-G [51] | CFG | 250 | 2.27 | 0.83 | 0.57 | |
| Masked modeling | MaskGIT [7] | 16 | 4.916 | 0.836 | 0.489 | |
| GIVT-MaskGIT(本文) | 16 | 4.64 | 0.85 | 0.49 | ||
| GIVT-MaskGIT(本文) | ,DB-CFG | 16 | 4.53 | 0.87 | 0.47 | |
| Sequence model | VQGAN [20] | Top-k mixed | 256 | 17.04 | ||
| VQGAN [20] | Top-k ,CG | 256 | 5.20 | |||
| ViT-VQGAN-L [76] | 1024 | 4.17 | ||||
| ViT-VQGAN-L [76] | CG | 1024 | 3.04 | |||
| GIVT-Causal(本文) | 256 | 5.6729 | 0.7490 | 0.5927 | ||
| GIVT-Causal(本文) | ,DB-CFG | 256 | 3.35 | 0.84 | 0.53 | |
| GIVT-Causal-L+A(本文) | 256 | 3.4556 | 0.7701 | 0.6131 | ||
| GIVT-Causal-L+A(本文) | ,DB-CFG | 256 | 2.5933 | 0.8085 | 0.5695 |
| 模型 | 推理 | Steps | FID | Precision | Recall |
|---|---|---|---|---|---|
| ADM [14] | 250 | 23.2 | 0.73 | 0.60 | |
| ADM-G [14] | CG | 250 | 7.72 | 0.87 | 0.42 |
| DiT-XL/2 [51] | 250 | 12.03 | 0.75 | 0.64 | |
| DiT-XL/2-G [51] | CFG | 250 | 3.04 | 0.84 | 0.54 |
| MaskGIT [7] | 16 | 7.8007 | 0.86 | 0.46 | |
| GIVT-MaskGIT(本文) | 16 | 4.8635 | 0.8785 | 0.4770 | |
| GIVT-Causal-L(本文) | 512 | 8.3493 | 0.7860 | 0.6060 | |
| GIVT-Causal-L+A(本文) | ,DB-CFG | 512 | 2.9163 | 0.8429 | 0.5480 |
| 模型 | COCO Pan.(PQ) | NYU Depth v2(RMSE) | ||
|---|---|---|---|---|
| 重建 | 推理 | 重建 | 推理 | |
| UViM[33] | 66.0 | 39.0 | 0.183 | 0.459 |
| GIVT(我们的) | 71.0 | 40.2 | 0.195 | 0.474 |
消融与可视化图5比较了模型配置(混合数、适配器)和采样算法(方差缩放、束搜索、DB-CFG)对FID的影响。对于每种模型配置,所有采样算法都带来了FID的显著改善,其中DB-CFG在所有配置中效果最佳。将从1增加到16总体上带来的改进略大于保持并添加适配器。此外,将适配器与结合使用会在采样算法上产生叠加改进。
图7(左)展示了方差缩放和CFG参数对采样FID的影响。在图7(右)中,我们可视化了预测标准差作为GIVT-Causal推理步骤的函数。标准差逐渐减小,这意味着采样过程后期的预测变得更加确定。此外,无条件预测通常具有更高的标准差,这符合预期。
对于GIVT-MaskGIT,预测每个潜向量的具有全协方差矩阵的单个高斯分布,而不是假设对角协方差,仅带来了约3%的适度改进。因此,具有因子化分量密度的GMM似乎是更有效的替代方案。此外,全协方差矩阵使得DB-CFG比对角协方差更难处理(因为高维多元分布具有更多低密度区域)。
样本图1展示了来自GIVT-Causal-L+A的十个样本,以及附录H展示了其他GIVT-Causal变体和GIVT-MaskGIT的样本。我们可以看到模型产生了高保真、连贯的样本。为了研究样本多样性,我们展示了固定类别下不同模型的多个样本。
在附录 H 的图 15 中,可以看到本文 VAE 的两个样本(通过解码从先验中采样的潜变量得到),呈现出图像纹理的混合。随后我们展示 GIVT-MaskGIT 推理的不同步骤,并观察到与基于 VQ 的模型类似的行为([7, 图 2])。
| 模型 | 类型 | #令牌 | #参数 | 准确率 |
|---|---|---|---|---|
| BigBiGAN[17] | 344M | 61.3 | ||
| iGPT-L[8] | 仅解码器 | 1024 | 1362M | 60.3 |
| VIM+CNN[76] | 仅解码器 | 1024 | 650M | 61.8 |
| VIM+ViT[76] | 仅解码器 | 1024 | 650M | 65.1 |
| MAGE ViT-L[39] | 编码器-解码器 | 256 | 404M | 78.9 |
| GIVT-Causal(我们的) | 仅解码器 | 256 | 304M | 65.1 |
表示学习表4展示了 无条件的GIVT-Causal和文献中生成模型的 ImageNet线性探针准确率(我们选择了在模型大小和计算量上最接近的模型变体)。GIVT-Causal与VIM+ViT (ViT-VQGAN)[76]相匹配,后者拥有超过倍的 模型参数和倍的 序列长度(因此FLOPs也更多)。GIVT-Causal仅被MAGE[39]超越,其 潜在编码器-解码器架构比仅解码器模型更适合表示学习。关于探针 准确率随层索引变化的调查可在附录F.
全景 分割和深度估计
表3比较了基于GIVT的UViM变体与基于VQ-VAE的基线(两者均无编码器/解码器上下文)在COCO Panoptic 2017上的性能[32]和NYU Depth v2[61]。我们分别报告了全景质量指标(PQ)[32]和RMSE,发现我们的基于GIVT的模型在全景分割上优于基线,在深度估计上略逊一筹。在附录H中,我们展示了可视化结果。
结论
在本文中,我们提出了对标准Transformer仅解码器模型的简单修改,使其能够生成实值向量。据我们所知,这是第一个能够生成实值向量序列的仅解码器模型。在VQ-GAN或Mask-GIT的图像生成背景下,这绕过了训练困难,如VQ-VAE中低码本使用率以及相应的缓解措施(如熵损失或码本分裂算法),因为可以使用标准VAE,这些VAE更容易训练。此外,我们的方法避免了大型嵌入矩阵,因为特征表示可以直接被我们的GIVT模型消费和预测。我们简单、无量化的方法在类条件图像生成和图像表示学习方面优于基于VQ的对应方法,通常有显著优势。GIVT在应用于UViM时,在密集预测任务中也获得了强劲的性能。我们希望未来的工作探索GIVT在其他模态(如音频和时间序列建模)中的应用。
致谢
我们感谢André Susano Pinto、Neil Houlsby、Eirikur Agustsson、Lucas Theis和Basil Mustafa对本项目的启发式讨论和有益反馈。我们也感谢Han Zhang在VAE训练代码方面的支持。
参考文献
参考文献按原论文顺序与书目信息保留。
- Aghajanyan, A., Huang, P.Y., Ross, C., Karpukhin, V., Xu, H., Goyal, N., Okhonko, D., Joshi, M., Ghosh, G., Lewis, M., Zettlemoyer, L.: CM3: A causal masked multimodal model of the internet. arXiv:2201.07520 (2022)
- Aghajanyan, A., Yu, L., Conneau, A., Hsu, W.N., Hambardzumyan, K., Zhang, S., Roller, S., Goyal, N., Levy, O., Zettlemoyer, L.: Scaling laws for generative mixed-modal language models. In: ICML (2023)
- Bao, H., Dong, L., Piao, S., Wei, F.: BEiT: BERT pre-training of image transformers. In: ICLR (2021)
- Bradbury, J., Frostig, R., Hawkins, P., Johnson, M.J., Leary, C.,
Maclaurin, D., Necula, G., Paszke, A., VanderPlas, J.,
Wanderman-Milne, S., Zhang, Q.: JAX:
composable transformations of
Python+NumPy programs (2018),
http://github.com/google/jax - Brock, A., Donahue, J., Simonyan, K.: Large scale gan training for high fidelity natural image synthesis. In: ICLR (2018)
- Chang, H., Zhang, H., Barber, J., Maschinot, A., Lezama, J., Jiang, L., Yang, M., Murphy, K.P., Freeman, W.T., Rubinstein, M., Li, Y., Krishnan, D.: Muse: Text-to-image generation via masked generative transformers. In: ICML (2023)
- Chang, H., Zhang, H., Jiang, L., Liu, C., Freeman, W.T.: MaskGIT: Masked generative image transformer. In: CVPR. pp. 11315–11325 (2022)
- Chen, M., Radford, A., Child, R., Wu, J., Jun, H., Luan, D., Sutskever, I.: Generative pretraining from pixels. In: ICML. pp. 1691–1703 (2020)
- Chen, X., Kingma, D.P., Salimans, T., Duan, Y., Dhariwal, P., Schulman, J., Sutskever, I., Abbeel, P.: Variational lossy autoencoder. In: ICLR (2016)
- Cheng, Z., Sun, H., Takeuchi, M., Katto, J.: Learned image compression with discretized gaussian mixture likelihoods and attention modules. In: CVPR. pp. 7939–7948 (2020)
- Das, A., Kong, W., Sen, R., Zhou, Y.: A decoder-only foundation model for time-series forecasting. arXiv:2310.10688 (2023)
- DeepMind, Babuschkin, I., Baumli, K., Bell, A., Bhupatiraju, S.,
Bruce, J., Buchlovsky, P., Budden, D., Cai, T., Clark, A., Danihelka,
I., Dedieu, A., Fantacci, C., Godwin, J., Jones, C., Hemsley, R.,
Hennigan, T., Hessel, M., Hou, S., Kapturowski, S., Keck, T., Kemaev,
I., King, M., Kunesch, M., Martens, L., Merzic, H., Mikulik, V., Norman,
T., Papamakarios, G., Quan, J., Ring, R., Ruiz, F., Sanchez, A.,
Sartran, L., Schneider, R., Sezener, E., Spencer, S., Srinivasan, S.,
Stanojević, M., Stokowiec, W., Wang, L., Zhou, G., Viola, F.: The
DeepMind JAX
Ecosystem (2020),
http://github.com/deepmind - Devlin, J., Chang, M.W., Lee, K., Toutanova, K.: Bert: Pre-training of deep bidirectional transformers for language understanding. NAACL-HLT (2019)
- Dhariwal, P., Nichol, A.: Diffusion models beat GANs on image synthesis. NeurIPS pp. 8780–8794 (2021)
- Dinh, L., Krueger, D., Bengio, Y.: Nice: Non-linear independent components estimation. In: ICLR (2015)
- Dinh, L., Sohl-Dickstein, J., Bengio, S.: Density estimation using real nvp. In: ICLR (2017)
- Donahue, J., Simonyan, K.: Large scale adversarial representation learning. In: NeurIPS (2019)
- Dosovitskiy, A., Beyer, L., Kolesnikov, A., Weissenborn, D., Zhai, X., Unterthiner, T., Dehghani, M., Minderer, M., Heigold, G., Gelly, S., Uszkoreit, J., Houlsby, N.: An image is worth 16x16 words: Transformers for image recognition at scale. ICLR (2021)
- Eisenach, C., Patel, Y., Madeka, D.: MQTransformer: Multi-horizon forecasts with context dependent and feedback-aware attention. arXiv:2009.14799 (2020)
- Esser, P., Rombach, R., Ommer, B.: Taming transformers for high-resolution image synthesis. In: CVPR. pp. 12868–12878 (2020)
- Fan, A., Lewis, M., Dauphin, Y.: Hierarchical neural story generation. In: ACL. pp. 889–898 (2018)
- Garza, A., Mergenthaler-Canseco, M.: TimeGPT-1. arXiv:2310.03589 (2023)
- Heusel, M., Ramsauer, H., Unterthiner, T., Nessler, B., Hochreiter, S.: GANs trained by a two time-scale update rule converge to a local nash equilibrium. NeurIPS (2017)
- Higgins, I., Matthey, L., Pal, A., Burgess, C., Glorot, X., Botvinick, M., Mohamed, S., Lerchner, A.: Beta-VAE: Learning basic visual concepts with a constrained variational framework. In: ICLR (2016)
- Ho, J., Salimans, T.: Classifier-free diffusion guidance. arXiv:2207.12598 (2022)
- Holtzman, A., Buys, J., Du, L., Forbes, M., Choi, Y.: The curious case of neural text degeneration. In: ICLR (2019)
- Huh, M., Cheung, B., Agrawal, P., Isola, P.: Straightening out the straight-through estimator: Overcoming optimization challenges in vector quantized networks. In: ICML (2023)
- Jacobsen, J.H., Smeulders, A.W., Oyallon, E.: i-revnet: Deep invertible networks. In: ICLR (2018)
- Kim, S., Jo, D., Lee, D., Kim, J.: MAGVLT: Masked generative vision-and-language transformer. In: CVPR. pp. 23338–23348 (2023)
- Kingma, D.P., Welling, M.: Auto-encoding variational bayes. arXiv:1312.6114 (2013)
- Kingma, D.P., Salimans, T., Jozefowicz, R., Chen, X., Sutskever, I., Welling, M.: Improved variational inference with inverse autoregressive flow. NeurIPS (2016)
- Kirillov, A., He, K., Girshick, R., Rother, C., Dollár, P.: Panoptic segmentation. In: CVPR. pp. 9404–9413 (2019)
- Kolesnikov, A., Susano Pinto, A., Beyer, L., Zhai, X., Harmsen, J., Houlsby, N.: UViM: A unified modeling approach for vision with learned guiding codes. NeurIPS pp. 26295–26308 (2022)
- Kumar, S., Anastasopoulos, A., Wintner, S., Tsvetkov, Y.: Machine translation into low-resource language varieties. In: ACL. pp. 110–121 (2021)
- Kumar, S., Tsvetkov, Y.: Von Mises-Fisher loss for training sequence to sequence models with continuous outputs. In: ICLR (2018)
- Kunz, M., Birr, S., Raslan, M., Ma, L., Li, Z., Gouttes, A., Koren, M., Naghibi, T., Stephan, J., Bulycheva, M., Grzeschik, M., Keki’c, A., Narodovitch, M., Rasul, K., Sieber, J., Januschowski, T.: Deep learning based forecasting: a case study from the online fashion industry. arXiv:2305.14406 (2023)
- Łańcucki, A., Chorowski, J., Sanchez, G., Marxer, R., Chen, N., Dolfing, H.J., Khurana, S., Alumäe, T., Laurent, A.: Robust training of vector quantized bottleneck models. In: IJCNN. pp. 1–7 (2020)
- Li, L.H., Chen, P.H., Hsieh, C.J., Chang, K.W.: Efficient contextual representation learning without softmax layer. arXiv:1902.11269 (2019)
- Li, T., Chang, H., Mishra, S., Zhang, H., Katabi, D., Krishnan, D.: Mage: Masked generative encoder to unify representation learning and image synthesis. In: CVPR. pp. 2142–2152 (2023)
- Li, Y., Mao, H., Girshick, R., He, K.: Exploring plain vision transformer backbones for object detection. In: ECCV. pp. 280–296 (2022)
- Lim, B., Arık, S.Ö., Loeff, N., Pfister, T.: Temporal fusion transformers for interpretable multi-horizon time series forecasting. International Journal of Forecasting pp. 1748–1764 (2021)
- Lu, J., Clark, C., Zellers, R., Mottaghi, R., Kembhavi, A.: Unified-IO: A unified model for vision, language, and multi-modal tasks. In: ICLR (2022)
- Menick, J., Kalchbrenner, N.: Generating high fidelity images with subscale pixel networks and multidimensional upscaling. arXiv:1812.01608 (2018)
- Mentzer, F., Agustsson, E., Tschannen, M.: M2T: Masking transformers twice for faster decoding. In: ICCV (2023)
- Mentzer, F., Gool, L.V., Tschannen, M.: Learning better lossless compression using lossy compression. In: CVPR. pp. 6638–6647 (2020)
- Mentzer, F., Minnen, D., Agustsson, E., Tschannen, M.: Finite scalar quantization: VQ-VAE made simple. arXiv:2309.15505 (2023)
- Nachmani, E., Levkovitch, A., Salazar, J., Asawaroengchai, C., Mariooryad, S., Skerry-Ryan, R., Ramanovich, M.T.: Lms with a voice: Spoken language modeling beyond speech tokens. arXiv:2305.15255 (2023)
- Nie, Y., Nguyen, N.H., Sinthong, P., Kalagnanam, J.: A time series is worth 64 words: Long-term forecasting with transformers. In: ICLR (2022)
- van den Oord, A., Vinyals, O., Kavukcuoglu, K.: Neural discrete representation learning. NeurIPS (2017)
- Parmar, N., Vaswani, A., Uszkoreit, J., Kaiser, L., Shazeer, N., Ku, A., Tran, D.: Image transformer. In: ICML. pp. 4055–4064 (2018)
- Peebles, W., Xie, S.: Scalable diffusion models with transformers. arXiv:2212.09748 (2022)
- Radford, A., Narasimhan, K., Salimans, T., Sutskever, I.: Improving language understanding by generative pre-training (2018)
- Rasul, K., Ashok, A., Williams, A.R., Khorasani, A., Adamopoulos, G., Bhagwatkar, R., Bilovs, M., Ghonia, H., Hassen, N., Schneider, A., Garg, S., Drouin, A., Chapados, N., Nevmyvaka, Y., Rish, I.: Lag-Llama: Towards foundation models for time series forecasting. arXiv:2310.08278 (2023)
- Razavi, A., Van den Oord, A., Vinyals, O.: Generating diverse high-fidelity images with VQ-VAE-2. NeurIPS (2019)
- Rombach, R., Blattmann, A., Lorenz, D., Esser, P., Ommer, B.: High-resolution image synthesis with latent diffusion models. In: CVPR (2022)
- Russakovsky, O., Deng, J., Su, H., Krause, J., Satheesh, S., Ma, S., Huang, Z., Karpathy, A., Khosla, A., Bernstein, M., et al.: Imagenet large scale visual recognition challenge. IJCV 115, 211–252 (2015)
- Sadeghi, H., Andriyash, E., Vinci, W., Buffoni, L., Amin, M.H.: PixelVAE++: Improved pixelvae with discrete prior. arXiv:1908.09948 (2019)
- Sajjadi, M.S., Bachem, O., Lucic, M., Bousquet, O., Gelly, S.: Assessing generative models via precision and recall. NeurIPS (2018)
- Salimans, T., Karpathy, A., Chen, X., Kingma, D.P.: PixelCNN++: Improving the PixelCNN with discretized logistic mixture likelihood and other modifications. In: ICLR (2016)
- Sauer, A., Schwarz, K., Geiger, A.: StyleGAN-XL: Scaling StyleGAN to large diverse datasets. In: SIGGRAPH (2022)
- Silberman, N., Hoiem, D., Kohli, P., Fergus, R.: Indoor segmentation and support inference from RGBD images. In: ECCV. pp. 746–760 (2012)
- Steiner, A., Kolesnikov, A., Zhai, X., Wightman, R., Uszkoreit, J., Beyer, L.: How to train your ViT? data, augmentation, and regularization in vision transformers. TMLR (2021)
- Strudel, R., Garcia, R., Laptev, I., Schmid, C.: Segmenter: Transformer for semantic segmentation. In: CVPR. pp. 7262–7272 (2021)
- Tokarchuk, E., Niculae, V.: On target representation in continuous-output neural machine translation. In: ACL (2022)
- Tokarchuk, E., Niculae, V.: The unreasonable effectiveness of random target embeddings for continuous-output neural machine translation. arXiv:2310.20620 (2023)
- Tomczak, J., Welling, M.: Vae with a vampprior. In: AISTATS. pp. 1214–1223 (2018)
- Tschannen, M., Bachem, O., Lucic, M.: Recent advances in autoencoder-based representation learning. arXiv:1812.05069 (2018)
- Tschannen, M., Kumar, M., Steiner, A., Zhai, X., Houlsby, N., Beyer, L.: Image captioners are scalable vision learners too. In: NeurIPS (2023)
- Vahdat, A., Andriyash, E., Macready, W.: Dvae#: Discrete variational autoencoders with relaxed boltzmann priors. NeurIPS (2018)
- Vahdat, A., Kautz, J.: NVAE: A deep hierarchical variational autoencoder. NeurIPS pp. 19667–19679 (2020)
- Van Den Oord, A., Kalchbrenner, N., Kavukcuoglu, K.: Pixel recurrent neural networks. In: ICML. pp. 1747–1756 (2016)
- Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A.N., Kaiser, Ł., Polosukhin, I.: Attention is all you need. NeurIPS (2017)
- Villegas, R., Babaeizadeh, M., Kindermans, P.J., Moraldo, H., Zhang, H., Saffar, M.T., Castro, S., Kunze, J., Erhan, D.: Phenaki: Variable length video generation from open domain textual descriptions. In: ICLR (2022)
- Wang, J., Du, Z., Chen, Q., Chu, Y., Gao, Z., Li, Z., Hu, K., Zhou, X., Xu, J., Ma, Z., Wang, W., Zheng, S., Zhou, C., Yan, Z., Zhang, S.: LauraGPT: Listen, attend, understand, and regenerate audio with GPT. arXiv:2310.04673 (2023)
- Wang, R., Chen, D., Wu, Z., Chen, Y., Dai, X., Liu, M., Jiang, Y.G., Zhou, L., Yuan, L.: Bevt: Bert pretraining of video transformers. In: CVPR. pp. 14733–14743 (2022)
- Yu, J., Li, X., Koh, J.Y., Zhang, H., Pang, R., Qin, J., Ku, A., Xu, Y., Baldridge, J., Wu, Y.: Vector-quantized image modeling with improved VQGAN. ICLR (2022)
- Yu, J., Xu, Y., Koh, J.Y., Luong, T., Baid, G., Wang, Z., Vasudevan, V., Ku, A., Yang, Y., Ayan, B.K., Hutchinson, B.C., Han, W., Parekh, Z., Li, X., Zhang, H., Baldridge, J., Wu, Y.: Scaling autoregressive models for content-rich text-to-image generation. TMLR (2022)
- Zhai, X., Kolesnikov, A., Houlsby, N., Beyer, L.: Scaling vision transformers. In: CVPR. pp. 12104–12113 (2022)
- Zhou, H., Zhang, S., Peng, J., Zhang, S., Li, J., Xiong, H., Zhang, W.: Informer: Beyond efficient transformer for long sequence time-series forecasting. In: AAAI. pp. 11106–11115 (2021)
arXiv版本历史
v1:初始版本。
v2:添加了关于损失的详细信息。
v3:多个模型更新:
GMM:从每通道标量GMM或单个高斯预测更改为-变量GMM,具有因子化分量 (即..,分量分布是-变量高斯分布,具有对角 协方差)。
引入了适配器。
大型GIVT-Causal模型(1.67B参数)及GIVT-Causal在512像素图像生成上的结果。
由于模型分片,将优化器从AdaFactor改为Adam(此更改对性能无影响)。
使用这些新模型变体更新图像生成结果。
v4:ECCV 2024最终版(微小改动)。
额外讨论
为什么图像使用连续潜变量? 连续潜变量自然适合图像表示等固有连续值数据,并避免过度压缩以及VQ引起的信息损失。去除VQ也避免了基于随机梯度的离散表示学习中众所周知的挑战:VQ-VAE[49]及其变体需要联合解决VQ优化问题(该问题本身是NP难的)和嵌入学习问题。大量文献[7, 27, 33, 37, 49]专注于缓解这些优化困难导致的问题,例如低词汇利用率(第1节)。例如,[46]通过使用乘积码本避免向量量化,但仍依赖不可微的标量量化。相比之下,使用GIVT,我们基于-VAE获得了最先进的结果,该VAE不包含任何不可微操作,也不需要VQ文献中的任何高级技巧。[7, 27, 33, 37, 49].
连续潜变量与无限词汇量在
理论上,连续(即..,实值)潜变量/标记在将它们精确映射到
离散代码时总是意味着无限的词汇量。换句话说,无限词汇量是连续潜变量的结果。在实践中,连续潜变量以有限精度表示,例如.., float32,但使用
原始的32位表示作为离散代码仍然意味着对于单个潜变量维度来说,词汇量大小B 不切实际地大。GIVT直接
建模连续潜变量,避免了物化这个词汇量。
架构细节
用于图像生成的不同模型变体的架构细节列于表5。我们依赖标准的pre-LN设置(参见,例如,UViMvtt.py以获取具体实现)。对于UViM实验,我们使用带有默认配置的GIVT-Causal,并在每个自注意力层之后插入一个交叉注意力层,如[33]中所述,以融合从输入RGB图像中提取的视觉特征。
适配器通过堆叠8个卷积、双射的iRevNet块[28]构建,隐藏通道维度为(导致112k 额外参数用于)。 每个块由3个卷积层组成,与GroupNorm和 ReLU交替,并且具有相同的输入和输出形状 (即.., 适配器应用于VAE潜在变量在将其重塑为序列之前)。 我们基于参考实现iRevNet.py构建iRevNet块, 将BatchNorm替换为GroupNorm并移除下采样层。
| 模型 | 大小 | 分辨率 | 宽度 | 深度 | MLP | 头数 | 参数 | 令牌 | 丢弃 | |||
|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 因果 | 基础 | 256 | 768 | 12 | 3072 | 12 | 86M | 256 | 0.1 | |||
| 因果 | 默认 | 256 | 16 | 16 | 1024 | 24 | 4096 | 16 | 304M | 256 | 0.2 | |
| 因果 | 大型 | 256 | 16 | 16 | 1536 | 48 | 8192 | 16 | 1.67B | 256 | 0.3 | |
| 因果 | 大型 | 512 | 32 | 32 | 1536 | 48 | 8192 | 16 | 1.67B | 512 | 0.1 | |
| MaskGIT | 默认 | 256 | 16 | 16 | 1024 | 24 | 4096 | 16 | 304M | 256 | 0.4 | |
| MaskGIT | 默认 | 512 | 16 | 16 | 1024 | 24 | 4096 | 16 | 304M | 1024 | 0.4 |
训练细节
损失函数和实现细节
我们从GIVT对序列中每个特征向量或软标记的分布进行建模的情况开始,条件是序列中所有先前的特征向量,使用一个-混合高斯模型,其分量在通道间独立(即..,这些分量是 具有对角协方差矩阵的多元高斯分布)。 对于大小为的维特征序列,长度为,预测的GIVT输出是一个张量在哪里是张量及是一个张量,全部沿最后一个维度堆叠, 分别包含均值、标准差和混合权重。 我们假设被一个小的正常数下界限制 (例如..,通过对相应的网络输出应用softplus函数并裁剪到),并且的条目为非负。由组成,其中是GMM分量沿最后一个轴堆叠的均值张量;以相同的方式组成。混合权重假定在分量之间归一化(例如..,通过softmax函数),即.., 对于 所有.
然后,特征序列的负对数似然(NLL) (即..训练图像在通过VAE编码器嵌入并应用重参数化之后)相对于GIVT预测的分布可以写成
为其中是高斯密度当使用单个高斯而不是GMM来建模每个输出通道时(即..使得对所有成立)公式2变为虽然这个损失简化为对序列维度的求和,但请注意(在GMM情况下为)
由通过 GIVT(设置到一个学到的[CLS]或[BOS]向量)。此外可以看出,令 的下界为 ,可避免某些 落入 的低密度区域时出现极大损失值。
图8展示了不同KL权重的VAE的训练损失曲线以及,遵循图4的设置。增加并减少导致更低的损失值。但请注意, 较低的损失并不总是导致较低的采样FID (见图4)。重要的是,初始和 最终损失值,以及整个训练过程中训练损失的相对减少量,可能与离散交叉熵或NLL通常观察到的值有显著差异,例如,在 语言建模中。
对于多变量情况,其中特征通道之间的依赖关系被建模,GIVT预测一个协方差矩阵对于每个混合成分,且数据的负对数似然变为其中是一个多元 高斯密度。
对于所有损失变体,MaskGIT版本的GIVT将求和简化为对索引对应于被掩码的位置, 并且损失通过掩码令牌的数量进行归一化。
最后,损失计算和采样可以轻松地使用 专门的深度学习包实现,例如JAX库 distrax[12]:
# With z and m, s, pi as in Eq. (2)
pdf = distrax.MixtureSameFamily(
mixture_distribution=distrax.Categorical(logits=pi),
components_distribution=distrax.MultivariateNormalDiag(
loc=m, scale_diag=s))
# Sample from next token distribution (teacher forcing)
samples = pdf.sample()
# Compute NLL
loss = -pdf.log_prob(z).mean()图像生成
对于ImageNet上的图像生成实验,我们改编了MaskGIT中的基于CNN的VQ-GAN分词器(参见vqgan_tokenizer.py)。我们将VQ层及相关损失替换为高斯重参数化层[30],并对CNN使用表6中给出的超参数。有关GIVT模型架构和变体的详细信息,请参见第B节。
embedding_dim |
32 |
|---|---|
filters |
128 |
num_res_blocks |
2 |
channel_multipliers |
(1, 1, 2, 2, 4) |
conv_downsample |
False |
activation_fn |
“swish” |
norm_type |
“GN” |
ImageNet预处理我们按照先前的工作对ImageNet数据进行预处理,具体如下:[7, 46]:
解码JPEG
随机裁剪,保留源图像的80%至100%
调整大小至目标分辨率(或)使用双三次滤波器 并启用抗锯齿
以0.5的概率随机水平翻转图像
全景 分割与深度估计
对于我们的UViM全景分割和深度估计实验, 我们采用公开的UViM GitHub 仓库并且仅替换阶段 I 中的 VQ 层,调整阶段 II 中的 transformer 解码器,并相应修改损失函数。
DB-CFG 实现
我们在图 中展示了用于 DB-CFG 的拒绝采样器的 JAX 实现。11.
无条件图像生成
在表7中,我们展示了无条件 ImageNet 生成的 FID 结果。
探测中间表示
训练的 GIVT-causal 中提取中间表示。无条件 通过将给定层的输出在序列维度上进行平均池化(得到一个维度为1024的特征向量)来生成ImageNet。因此,输入图像首先用VAE编码器编码,然后像教师强制期间一样馈送到GIVT(即.. 潜在序列首先右移并填充,然后以因果注意力掩码馈送到GIVT)。
遵循先前的工作[8, 76],我们研究了层索引对线性探测准确率的影响。为了加速评估,我们依赖于使用快速评估器形式的25次射击分类。[78]并测试一系列不同的数据集。图9中展示的结果表明,第10至20层对所有数据集都能达到高准确率。
我们选择第18层用于使用完整ImageNet训练集的线性探测实验。为了选择超参数,我们预留了1%的训练数据,并考虑了[68]附录A中描述的调度和超参数选择。结果以及其他文献中的生成模型结果可见于表4. 参见 第5.1节,了解 这些结果的讨论。
方法在FLOPs方面的比较
在表1报告的模型中,只有两个 扩散模型 [14, 51]报告了FLOPs,但没有序列或掩码 模型,而这些模型与我们的GIVT模型最为相关。假设 相同的自回归Transformer优化(例如缓存)和VAE 推理成本,我们得到以下近似比较,使用GIVT-Causal-L+A的FLOPs作为 参考。MaskGIT和GIVT-MaskGIT在推理时所需的FLOPs几乎相同。
| 模型 | 推理 | FID | Precision | Recall |
|---|---|---|---|---|
| ADM [14] | 26.2 | 0.61 | 0.63 | |
| MAGE (ViT-B) | temp | 12.080213022509724 | 0.61988 | 0.6364 |
| MAGE (ViT-L) | temp | 9.9528908191744 | 0.65928 | 0.6589 |
| GIVT-Causal(本文) | , | 19.9226 | 0.6259 | 0.6013 |
| GIVT-Causal(本文) | , | 17.6995 | 0.6587 | 0.6282 |
| GIVT-Causal-L(本文) | , | 11.0204 | 0.7207 | 0.6015 |
| 模型 | 指导 | 步骤 | 参数 | FLOPs | FID |
|---|---|---|---|---|---|
| GIVT-Causal-L+A | – | 256 | 1.67B | 3.46 | |
| VQGAN [20] | – | 256 | 1.4B | 17.04 | |
| ViT-VQGAN-L [76] | – | 1024 | 1.7B | 4.17 | |
| GIVT-Causal-L+A | DB-CFG | 256 | 1.67B | 2.59 | |
| VQGAN [20] | CG | 256 | 1.4B | 5.20 | |
| ViT-VQGAN-L [76] | CG | 1024 | 1.7B | 3.04 |
更广泛的影响
本文描述的方法修改了仅解码器的Transformer模型,使其能够生成特征向量序列,这可以看作是对Transformer模型基础的研究努力,具有许多潜在的下游应用。更广泛的影响将主要取决于下游应用及其受这项工作影响的方式。
在这里,我们专注于视觉数据的生成。对于类条件图像生成,我们依赖ImageNet,所得模型通过类标签提供对生成图像内容的基本控制。这些模型不允许对图像内容进行细粒度控制,或对现有图像进行操作。ImageNet固有的偏见和问题已被充分研究和记录[Uday Prabhu和Birhane,2020]。我们预计我们的图像生成模型可能反映其中一些问题,并建议用户在使用和部署模型时谨慎考虑这些问题。
将我们的图像生成模型扩展到文本到图像界面相对直接,这使得可以在从网络收集的更大且较少策划的图像/文本数据集上进行训练。此类模型允许更细粒度的控制,因此可能被滥用于恶意目的,如传播错误信息或身份盗窃。此外,此类模型可能反映大型非策划网络数据集中的偏见,如有害的刻板印象。因此,此类模型的部署和发布必须更加谨慎。
import jax
import jax.numpy as jnp
import tensorflow_probability as tfp
tfd = tfp.substrates.jax.distributions
# We assume the ps have shape (b, seq_len, c).
def rejection_sample(
seed, p_cond: tfd.Normal, p_uncond: tfd.Normal, w, max_samples=1_000
):
rng_sample, rng_uni = jax.random.split(seed, 2)
scale_simple = jnp.stack([p_cond.scale, p_uncond.scale], -1).max(-1) * 2
simple = tfd.Normal(loc=p_cond.loc, scale=scale_simple) # Proposal.
def unnormalized_pcfg(x):
return jnp.exp((1 + w) * p_cond.log_prob(x) - w * p_uncond.log_prob(x))
# Find scaling factor by checking for max around the conditional loc.
points = p_cond.loc[None, ...] + jnp.linspace(-10, 10, 1001).reshape(
-1, 1, 1, 1)
fac = jnp.max(unnormalized_pcfg(points) / simple.prob(p_cond.loc), axis=0)
# Shape (max_samples, b, seq_len, c)
xs = simple.sample(seed=rng_sample, sample_shape=(max_samples,))
facq = fac * simple.prob(xs)
ys = jax.random.uniform(rng_uni, shape=facq.shape, minval=0.0, maxval=facq)
p = unnormalized_pcfg(xs)
# Mask is True if the sample is valid.
mask = ys < p
# Now we need to do tricks to get the first element in `mask` that is
# True, in a jit-able way. We do this by making a shifted mask that is
# False for every element after the first True. Example:
# mask [0, 1, 0, 1, 0, 0, 1, 0]
# cmask [0, 1, 1, 1, 1, 1, 1, 1]
# shifted_cmask [0, 0, 1, 1, 1, 1, 1, 1]
# keep [0, 1, 0, 0, 0, 0, 0, 0] # <- picks the first valid!
cmask = jnp.cumsum(mask, axis=0).astype(jnp.bool_)
shifted_cmask = jnp.pad(
cmask, [(1, 0), (0, 0), (0, 0), (0, 0)], constant_values=False
)[:-1]
assert shifted_cmask.shape == mask.shape
keep = jnp.logical_and(cmask, jnp.logical_not(shifted_cmask))
sample = jnp.where(keep, xs, 0).sum(0) # Grab the first valid sample.
ok = mask.sum(0) > 0 # Fall back to pdf_c if not ok.
return jnp.where(ok, sample, pdf_c.sample(seed=rng_sample))jax.jit兼容的DB-CFG拒绝采样器实现。额外的视觉示例
在图12中,我们展示了当向VAE输入ImageNet验证集中的图像时,我们的VAE的重建结果。当多次从编码器分布中采样时,我们看到重建中的低层纹理变化。图13显示了固定标签的多个样本,以展示我们模型的多样性并与基线进行比较。在图14以及16我们展示了来自不同GIVT-Causal 变体的样本。在图17中,我们展示了 MaskGIT 对于我们在图1中展示的相同类别的样本。图18展示了我们GIVT-MaskGIT模型的样本。我们 探索在采样中途更改标签(即..,在生成图像上半部分之后)对GIVT-Causal的影响,见图19.
| GIVT-Causal(本文) | GIVT-MaskGIT(本文) | 验证集 |
| BigGAN [5](来自 [7, 图 9]) | MaskGIT (VQ) [7, 图 9] |
| VAE 样本 | 步骤 1 | 步骤 4 | 步骤 8 | 步骤 16 |
| 输入 | 真实值 | GIVT-Causal UViM | 基于VQ的UViM |
| 输入 | 地面真值 | GIVT-Causal UViM | 基于VQ的UViM |
我们在第节讨论连续潜变量与无限词汇表之间的关系。A.↩︎