CONTINUOUS TOKENS · TRANSFORMERS · 中文全文译稿

GIVT:
Generative Infinite-Vocabulary Transformers

Michael Tschannen · Cian Eastwood* · Fabian Mentzer°Google DeepMind* 在 GDM 担任 Student Researcher 期间完成。° 作出重要技术贡献。代码与模型 checkpoint:google-research/big_vision。
InputLookup table → linear projection
Sequencefinite tokens → real vectors
Outputcategorical → Gaussian mixture

Transformer 的自回归结构并不依赖有限词表;只要重写输入投影与输出分布,它同样可以生成连续向量序列。

PDF
32 页
Figures
21
Tables
8
Equations
2
References
79

摘要

我们提出 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].

Figure 1原论文图面与译注
选定的512×512512\times512来自GIVT-Causal-L的样本 涵盖10个ImageNet类别(130、130、138、144、933、145、360、207、829、 248)。

为了将这些能力用于图像,最近的工作[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不是预测有限词汇上的类别分布,而是预测dd元高斯混合模型(GMM)的参数。 我们以与标准变换器解码器相同的方式训练GIVT:使用因果注意力掩码和教师强制[72],并且还探索了如MaskGIT中的快速渐进掩码双向建模[13, 7, 6].

Figure 2原论文图面与译注
我们比较了标准的离散令牌生成Transformer(左)与我们的连续、无限词汇变体(GIVT,右),两者使用相同的仅解码器架构。在输入时,GIVT线性嵌入一系列实值向量,而不是通过查找离散令牌。在输出时,GIVT预测多元连续分布的参数,而不是分类分布。
Figure 3原论文图面与译注

GIVT因果训练和推理。左:在训练期间,我们从VAE编码器中采样一个实值潜向量序列,并通过教师强制训练GIVT。右:在推理期间,我们从左到右采样一个向量序列,并将其输入VAE解码器。我们注意到,我们还探索了类似MaskGIT的GIVT模型,此处未显示。没有任何组件使用量化器。

与使用VQ-VAE的两阶段方法类似,并且类似于潜在扩散模型的两阶段方法[55, 51],我们首先使用高斯先验β\beta-VAE[30, 24]学习一个较低维度的潜空间,然后使用GIVT对其进行建模。我们强调,训练β\beta-VAE和GIVT仅依赖于深度学习工具箱中的标准技术,而不是VQ-VAE文献中的高级训练技术,如辅助损失[49, 7]在潜在表示、码本重新初始化[37]或专用优化算法上[33, 27].

我们的主要贡献可总结如下:

  1. 我们展示了GIVT在类条件图像生成中优于VQGAN[55](及其后续变体)和MaskGIT[7],通常以较大幅度和/或显著更低的计算成本实现。GIVT还与强大的潜在扩散基线竞争,尤其是在高分辨率下。

  2. 我们为连续情况推导了标准采样技术的变体,如温度采样、束搜索和无分类器引导(CFG)[25],并展示了它们的有效性。

  3. 我们证明了GIVT在表示学习方面以显著更低的计算成本匹配或超越了先前的序列图像生成模型。

  4. 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训练

我们首先训练一个连续潜变量 β\beta-VAE[24],其编码器和先验为高斯分布,最初由[30]. 给定输入图像xx,编码器EE预测多元正态分布的均值μ\mu和协方差σ\sigma(对角协方差矩阵),并使用重参数化技巧zzN(μ,σ)\mathcal N(\mu, \sigma)中采样表示[30]。然后,VAE解码器将潜在序列映射回图像。由于我们使用高斯编码器分布,证据下界(ELBO)[30]中的KL项可以按照[30, 第F.1节] 中描述的闭式形式计算。. 至于 ELBO中的重建/似然项,我们依赖于MSE、 感知损失和GAN损失的混合用于图像生成,遵循[20, 7],或用于密集预测任务的 分类交叉熵[33]。我们的编码器 在空间上对xx进行下采样,由此我们 得到zz,其空间维度为h×wh \times w,特征维度为dd,其中h=H/16,w=W/16h{=}\lceil H/16\rceil, w{=}\lceil W/16\rceil,给定一个H×WH{\times}W输入xx。为了计算KL项,相关的μ\muσ\sigma形状为w×h×dw \times h \times d被展平为whdwhd向量。

超参数β\beta乘以KL项控制着zz被正则化的强度。正如我们将在 第5节中看到的,这种对VAE的正则化对于 能够很好地建模由此产生的(真实的)潜在分布p(z)p(z)是重要的。

GIVT训练

接下来,我们训练一个GIVT来预测p(z)p(z)p(zc)p(z | c)(当条件信号cc可用时,例如,在 类条件生成中)。表示zz被重塑为长度为hwhwdd实值向量(或“软标记”)。请注意,这与标准的 VQ-VAE设置不同,在标准设置中,潜在变换器解码器建模一个hwhw-长度的序列,其元素为整数表示码本索引。为适应这一差异,我们对标准的仅解码器变换器架构做了两处 小改动(见图2):在输入处,我们用单个线性层替换嵌入 查找表,以从dd投影到变换器的隐藏 维度。在输出处,我们不预测类别分布, 而是让变换器预测连续分布的参数。假设混合 分量的通道间独立性,我们用kk-混合高斯模型(GMM)来建模该连续分布。因此,GIVT模型对每个 软标记预测2kd+k2 k d + k个参数(kdkd均值和kdkd混合成分的方差参数,以及kk混合概率)。实验上,我们发现用softmax激活函数归一化混合概率,用softplus归一化方差参数是有益的。

我们在GIVT预测的分布p~\tilde p上使用标准的交叉熵损失(等同于负对数似然),并最小化LT=cEz[logp~(zc)]\mathcal L_\text{T} = \sum_c \mathbb E_z \left[ -\log \tilde p(z | c) \right],假设类别或条件信号cc均匀分布(详见 附录C关于损失的细节)。我们训练两种类型的GIVT模型,如下所述。

GIVT-Causal。这里,GIVT 在长度为 hwhw 的潜变量序列中预测每个 dd 维向量,条件为此前全部向量。因此,self-attention 层采用时间因果 mask[20, 72](这使模型能在推理时顺序生成,与 causal inference 无关)。这种训练策略也称为 teacher forcing,与 VQ-GAN 中的潜变量建模类似[20]。对于类别条件图像生成,我们在输入序列前添加一个 [CLS] 向量,即为每个类别 cc 学习一个向量。

GIVT-MaskGIT与MaskGIT一样[7],我们在训练时随机掩蔽输入序列的一个子集,然后在推理时逐步揭示被掩蔽的标记。与相比,唯一的变化是[7]与我们的实值标记相关:由于我们有无限多个标记,因此没有 明显的方法来定义特殊的掩码标记(当使用VQ时,可以 直接扩展词汇表以包含特殊标记,例如[MASK])。相反,给定zz和一个掩码MM指示每个位置是否 被掩码,我们首先将zz中对应于MM的位置替换为零(以去除信息),然后 用单个密集层对其进行嵌入,如上所述。此外,我们拼接在特征维度上两个学习到的特殊向量之一,一个[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解码器中的束搜索相同。对于每个样本,我们维护BB光束,并且在每一步中,我们为每个光束采样多个候选(这里称为“扇出”)。然后,我们计算所有光束和扇出到当前采样步骤为止的累积对数概率,并选择具有最高累积对数概率的BB个光束。最后,在GIVT中没有与top-k采样[21]类似的概念,因为它预测的是连续分布。

Figure 4原论文图面与译注
β\beta-VAE 消融:训练 VAE 时,KL 权重 β\beta、通道数 dd 与 mixture 数 kk 之间的关系。圆形标记表示使用 Base-size GIVT-Causal 得到的采样 FID。随着 β\betakk 增大,采样 FID 改善,但重建 FID 也随之升高,从而限制了可达到的最佳采样 FID。
Figure 5原论文图面与译注
不同采样策略和模型变体(GIVT-Causal-L)对样本质量的影响。增加混合成分的数量kk以及添加适配器(+A)带来了叠加的改进。DB-CFG是所有模型配置中最有效的采样策略。

基于分布的免分类器引导在扩散文献中,免分类器引导(CFG)[25]已被成功应用。具体来说,条件扩散模型通过额外的空类\emptyset来学习无条件数据分布。然后,在推理过程中,条件对数密度被“移离”无条件对数密度:给定引导权重ww,更新后的(扩散)分数估计为其中ϵ\epsilon估计数据分布的对数密度的梯度,ϵ(z,c)zlogp~(zc)\epsilon(z,c) \propto \nabla_{z} \log \tilde p(z|c)(参见[25, 第2节])。由此,我们现在为我们的GIVT推导出一个CFG变体,因为我们直接预测一个密度。我们将这种方法称为“基于密度的CFG”(DB-CFG)。公式1可以写成 .., ϵ~\tilde \epsilon估计密度的对数pCFG(zc)p~(zc)1+wp~(z)wp_\text{CFG}(z|c) \propto \tilde p(z|c)^{1+w}\tilde p(z|\emptyset)^{-w}(见图6)。因此,我们 希望调整我们的模型以从pCFGp_\text{CFG}中采样。我们遵循[25]并训练 GIVT,增加一个额外的空类\emptyset。在推理过程中,我们每一步 评估GIVT两次,一次以实际标签cc为条件,一次以\emptyset为条件。为了实现无分类器 引导,我们随后必须从由两个GIVT 预测导出的未归一化版本pCFG(z)p_\text{CFG}(z)中采样。为此,我们转向拒绝采样,这需要: 1)一个未归一化的密度;2)一个好的提议分布pp', 这接近真实的目标分布;以及3)一个缩放因子KK用于限制pp'与未归一化目标密度之间的似然比。

公式(1)。

ϵ~(z,c)=(1+w)ϵ(z,c)wϵ(z,),\begin{aligned} \tilde \epsilon(z, c) = (1 + w) \epsilon(z, c) - w \epsilon(z, \emptyset), \end{aligned}

公式(2)。

ϵ~(z,c)(1+w)zlogp~(zc)wzlogp~(z)\tilde \epsilon(z,c) \propto (1+w) \nabla_{z} \log \tilde p(z|c) - w \nabla_z \log \tilde p(z|\emptyset)

我们混合的分布是GMM,找到好的提议分布可能具有挑战性。相反,我们首先从p~(zc)\tilde p(z|c)中采样混合索引,并对p~(zc)\tilde p(z|c)p~(z)\tilde p(z)中相应的混合分量应用DB-CFG(这些分量是具有对角协方差的多元高斯分布)。我们经验发现,无条件分量(.., 分布 使用\emptyset标签) 预测的分布往往比条件分布具有更大的方差(如 图6所示)。因此,合理的做法是从N(μc,2σc)\mathcal N(\mu_c, 2\sigma_c)中选取样本提议,其中μc,σc\mu_c, \sigma_c是 GIVT 在给定标签cc时预测的参数。我们凭经验发现,抽取 1000 个样本足以在 99.9% 的情况下找到至少一个有效样本。对于剩余的<0.1{<}0.1%,则回退到从N(μc,σc)\mathcal N(\mu_c, \sigma_c).

我们强调,DB-CFG 的额外开销很小:每个推理步只需两次而不是一次前向传播,分别预测条件分布与无条件分布。随后,我们在加速器上从这些分布并行抽取 1000 个样本,速度很快。关于 Python 代码,参见附录 3.4

Figure 6原论文图面与译注
我们的 Density-Based Classifier-Free Guidance(DB-CFG)示意图。图中展示 GIVT 预测的条件与无条件概率密度函数,以及两个 ww 取值下得到的 CFG 概率密度。可以看到 CFG 分布更加尖锐。我们使用 rejection sampling 从 pCFGp_{\mathrm{CFG}} 中采样。
Figure 7原论文图面与译注

左:DB-CFG的影响(第3.4节)和方差缩放(第3.4节)对我们类条件256×256256\times256GIVT-Causal模型采样FID的影响。DB-CFG值在[0.3,0.8][0.3, 0.8]中,方差缩放参数tt[0.9,1.0][0.9, 1.0]导致较低的FID。右:GIVT-Causal对类别130预测的GMM的平均标准差,在128个样本上取平均:条件预测具有较低的标准差;当光栅扫描潜在特征向量时,线条变化处可观察到尖峰。

实验

图像生成

我们使用ImageNet1k[56]并探索类别条件生成(其中我们将GIVT条件于类别标签)用于256px和512px,以及条件生成用于256px。

β\beta-VAE我们 紧密遵循MaskGIT的设置[7]。我们使用其VAE架构, 由ResBlocks构建(详见附录C;编码器和解码器共 有53.5M参数),移除VQ层及相关损失, 并将其替换为预测μ,σ\mu,\sigma的线性层(第3.1节)。我们使用与[7, 46]中相同的重建、感知和GAN损失权重,以及相同的优化器参数; 我们仅改变潜在维度dd和KL项的权重β\beta。默认情况下,我们将token维度设置为d=16d=16 (..,VAE预测每个token的16个 均值和方差)和β=5105\beta = 5 \cdot 10^{-5}。我们注意到,我们的VAE是在256×256256 \times 256张图像上训练的,并且我们也将其用于我们的512×512512 \times 512实验,而无需重新训练(如[7]).

GIVT。对于GIVT-Causal,我们遵循原始的Transformer解码器架构[72]在仅解码器模式下,但 从注意力层、MLP块和LayerNorm中移除偏置,并 按照常见做法将ReLU替换为GELU。对于GIVT-MaskGIT,我们简单 地在训练期间移除注意力掩码,并输入掩码输入而不是 移位输入。我们使用BERT-Large配置[13]默认情况下 (304M参数),并探索更大的主干网络,具有1.67B 参数,用后缀“-L”表示(详见附录B)。对于带有适配器的模型变体(后缀“+A”),我们使用8个双射 iRevNet块的堆叠[28](隐藏通道 维度为4d4d,导致额外的112k 参数用于d=16d=16), 应用于w×h×dw \times h \times d表示,然后再将其重塑为序列。我们将GIVT模型配置为预测1616-混合 GMM,其中混合分量是因子化的(即混合分量是 具有对角协方差的多变量高斯分布),并探索预测 一个单一的多变量高斯分布,建模token的完整协方差矩阵作为替代方案。对于条件生成 实验,我们使用一个学习到的嵌入,将其前置到嵌入的 token序列之前。为了训练GIVT,我们使用带有余弦调度(500 个epoch;线性预热50个epoch)的Adam优化器,将学习率和 权重衰减设置为10310^{-3}10410^{-4},分别,优化器的β2\beta_2参数设置为0.950.95,dropout概率设置为0.20.2对于GIVT-causal和0.40.4对于GIVT-MaskGIT,以及批量大小 设置为81928192。我们使用与VAE训练期间相同的数据 增强(参见[7, 46]),并且对于每个批次从VAE编码器分布中采样(除了数据增强之外的另一个随机性来源)。

我们在JAX中实现GIVT[4]并使用distrax[12]来实现候选并计算对数概率。

GIVT-MaskGIT推理遵循[7],我们将推理步数固定为16,并采用余弦调度(.. 让r=i/16r=i/16在步骤ii,掩码 标记的比例由cos(π/2r)\cos(\pi/2 r))给出。 我们还在每一步按似然对标记进行排序,并使用 “选择温度”tCt_C.

探索VAE潜在空间为了更好地理解 特征维度dd、KL正则化β\beta,VAE的重建质量,以及GIVT的采样质量,我们训练了潜在维度为{4,8,16,32}\{4, 8, 16, 32\}β\beta的VAE,使用本节开头描述的VAE训练设置。对于每个得到的VAE,我们训练了一个GIVT-Causal,使用较小的BERT-Base{2.5105,5105,104,2104}\{2.5 \cdot 10^{-5}, 5 \cdot 10^{-5}, 10^{-4}, 2 \cdot 10^{-4}\}维度,以及一系列混合数[13]的值。kk.

评估对于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替换为一个β\beta-VAE并使用GIVT编码器-解码器 来建模连续潜码。对于VAE,我们使用与[33]相同的基于Transformer的自编码器架构(6层编码器和12层 解码器)和交叉熵损失。我们设置d=16d=16, k=1k=1,以及KL权重β=2.5104\beta=2.5 \cdot 10^{-4}用于全景 分割,β=2104\beta=2 \cdot 10^{-4}用于深度估计。为了以与[33]相同的方式构建编码器-解码器GIVT 模型,我们采用为ImageNet生成所描述的因果 变体,并在每个自注意力层之后插入一个交叉注意力层。遵循[33],我们使用来自[62]的ImageNet-21k预训练ViT-L/16作为编码器,将图像 分辨率设置为512像素,并采用来自[33]的预处理和优化 超参数。我们使用没有编码器和解码器上下文的UViM变体[33]。最后,我们考虑 方差缩放和束搜索,在训练集的保留子集上选择参数,如[33].

结果

图像生成

VAE潜在空间在图4中,我们展示了改变KL项的权重β\beta如何影响1) VAE 重建FID和2) 在相应潜在序列上训练的Base-size GIVT-Causal的采样FID。对于1),增加β\beta会导致重建FID变差 ,因为VAE在潜在空间中存储的信息更少。它将更多的建模工作转移到VAE解码器上,使得解码器逐渐更具生成性,这会影响采样质量(参见[67, 第 7节][55]以获取更多讨论)。

对于2),我们看到了相反的趋势:增加β\beta会导致在潜在序列上训练的GIVT模型的采样FID降低(更好)。可以说, 这是因为VAE潜在序列更接近高斯先验,因此更容易被GIVT建模。最后,增加混合数量kk最初显著降低了采样FID,在k=16k=16处达到平台期。因此我们默认设置k=16k=16β=5105\beta=5 \cdot 10^{-5},并对较大的(-L)GIVT模型使用带有β=105\beta=10^{-5}的VAE。我们强调,在文献中探索和调整诸如β\beta或类似地VQ中的词汇表大小和承诺损失等超参数是常见的[20],表5 [55],表8[51][46,图3].

采样FID在表1我们展示了在类别条件下的四种模型类别的采样FID。256×256256\times256ImageNet:GAN、 基于扩散的方法,以及掩码和序列建模 方法。GIVT-MaskGIT优于MaskGIT[7]后者具有相当的模型大小 和推理成本,而DB-CFG带来了额外的改进。在 没有引导技术的情况下,我们的GIVT-Causal模型大幅优于所有 扩散基线以及VQGAN。使用引导 技术,GIVT-Causal获得了3.35的FID,而VQGAN为5.20,且计算量超过4.5×4.5\times更小的 模型(GIVT为0.3B参数,而对比模型为1.4B参数),并且也优于32% 更大的LDM-4-G。我们更大的GIVT-Causal-L+A在无引导和有引导情况下,分别比ViT-VQGAN获得16%和17%的FID降低,而ViT-VQGAN具有相同的生成Transformer大小但4×4\times更长的序列长度(导致超过4×4\times较慢的 采样)以及一个10×10\times更大的 VAE。

我们展示了采样FID用于512×512512\times5122中的 ImageNet。GIVT-MaskGIT 在模型大小和推理成本相当的情况下,FID 比 MaskGIT 低 38%。GIVT-Causal-L+A 在无引导和有引导的情况下均优于 DiT-XL/2(目前最佳的 DiT 模型)(尽管模型更大)。

最后,我们在附录中展示了E条件结果。这项任务难度显著 更高,但 GIVT-Causal 大幅超越了基于扩散的 ADM[14]

Table 1语义 HTML 转录
类别条件 256×256256\times256 ImageNet 上的结果。GIVT-Causal 在更小的模型规模(相对 VQGAN)或显著更短的序列长度(相对 ViT-VQGAN)下优于基于量化的对应方法。表中报告 FID、Precision 与 Recall;FID 按标准 ADM 评测套件相对训练集计算。+A 表示带 adapter 的 GIVT;CG 表示 classifier guidance 的接受率或尺度;CFG =w=w 表示权重为 ww 的 classifier-free guidance [25];DB-CFG =w=w 表示本文基于分布的 CFG 变体;Top-k 表示 Top-k sampling [21](mixed 表示多个 kk);tt 表示缩放模型预测 σ\sigma 的采样温度;tCt_C 是 MaskGIT 的选择温度。Steps 为推理步数。{}^\dagger 数值由作者根据公开代码获得;{}^\star 表示推理使用 activation caching。
类别模型推理StepsFID \downarrowPrecision \uparrowRecall \uparrow
GANBigGAN-deep [5]6.950.870.28
StyleGAN-XL [60]2.300.780.53
DiffusionADM [14]25010.940.690.63
ADM-G [14]CG =1.0=1.02504.590.820.52
LDM-4 [55]25010.560.710.62
LDM-4-G [55]CFG =1.5=1.52503.600.870.48
DiT-XL/2 [51]2509.620.670.67
DiT-XL/2-G [51]CFG =1.5=1.52502.270.830.57
Masked modelingMaskGIT [7]tC=4.5t_C=4.5164.916 {}^\dagger0.836 {}^\dagger0.489 {}^\dagger
GIVT-MaskGIT(本文)tC=35t_C=35164.640.850.49
GIVT-MaskGIT(本文)tC=60t_C=60,DB-CFG =0.1=0.1164.530.870.47
Sequence modelVQGAN [20]Top-k == mixed256 {}^\star17.04
VQGAN [20]Top-k =600=600,CG =0.05=0.05256 {}^\star5.20
ViT-VQGAN-L [76]1024 {}^\star4.17
ViT-VQGAN-L [76]CG =0.5=0.51024 {}^\star3.04
GIVT-Causal(本文)t=0.9t=0.9256 {}^\star5.67290.74900.5927
GIVT-Causal(本文)t=0.95t=0.95,DB-CFG =0.4=0.4256 {}^\star3.350.840.53
GIVT-Causal-L+A(本文)t=0.9t=0.9256 {}^\star3.45560.77010.6131
GIVT-Causal-L+A(本文)t=0.95t=0.95,DB-CFG =0.4=0.4256 {}^\star2.59330.80850.5695
Table 2语义 HTML 转录
类别条件 512×512512\times512 ImageNet 上的结果。指标采用标准 ADM 评测套件,FID 相对训练集计算。GIVT-MaskGIT 只需 16 个推理步即可取得有竞争力的 FID,并超过其 VQ 对应模型;GIVT-Causal-L+A 超过最佳 DiT 变体 DiT-XL/2-G。{}^\dagger 数值由作者根据公开代码获得;{}^\star 表示推理使用 activation caching。
模型推理StepsFID \downarrowPrecision \uparrowRecall \uparrow
ADM [14]25023.20.730.60
ADM-G [14]CG =1.0=1.02507.720.870.42
DiT-XL/2 [51]25012.030.750.64
DiT-XL/2-G [51]CFG =1.5=1.52503.040.840.54
MaskGIT [7]tC=4.5t_C=4.5167.8007 {}^\dagger0.86 {}^\dagger0.46 {}^\dagger
GIVT-MaskGIT(本文)tC=140t_C=140164.86350.87850.4770
GIVT-Causal-L(本文)t=0.9t=0.9512 {}^\star8.34930.78600.6060
GIVT-Causal-L+A(本文)t=0.9t=0.9,DB-CFG =0.9=0.9512 {}^\star2.91630.84290.5480
Table 3语义 HTML 转录
基于 GIVT-Causal 与 VQ-VAE 的 UViM 在全景分割(COCO Panoptic 2017)和深度估计(NYU Depth v2)上的评测。表中分别报告验证集标签图的 VAE/VQ-VAE 重建指标(recon.)以及实际稠密预测任务的推理指标(inference):全景质量 PQ 与 RMSE。GIVT 的结果与基于 VQ 的 UViM 相当。
模型COCO Pan.(PQ\uparrowNYU Depth v2(RMSE\downarrow
重建推理重建推理
UViM[33] 66.0 39.0 0.183 0.459
GIVT(我们的) 71.0 40.2 0.195 0.474

消融与可视化5比较了模型配置(混合数kk、适配器)和采样算法(方差缩放、束搜索、DB-CFG)对FID的影响。对于每种模型配置,所有采样算法都带来了FID的显著改善,其中DB-CFG在所有配置中效果最佳。将kk从1增加到16总体上带来的改进略大于保持k=1k=1并添加适配器。此外,将适配器与k=16k=16结合使用会在采样算法上产生叠加改进。

7(左)展示了方差缩放和CFG参数对采样FID的影响。在图7(右)中,我们可视化了预测标准差作为GIVT-Causal推理步骤的函数。标准差逐渐减小,这意味着采样过程后期的预测变得更加确定。此外,无条件预测通常具有更高的标准差,这符合预期。

对于GIVT-MaskGIT,预测每个潜向量的具有全协方差矩阵的单个高斯分布,而不是假设对角协方差,仅带来了约3%的适度改进。因此,具有因子化分量密度的GMM似乎是更有效的替代方案。此外,全协方差矩阵使得DB-CFG比对角协方差更难处理(因为高维多元分布具有更多低密度区域)。

样本1展示了来自GIVT-Causal-L+A的十个512×512512\times512样本,以及附录H展示了其他GIVT-Causal变体和GIVT-MaskGIT的样本。我们可以看到模型产生了高保真、连贯的样本。为了研究样本多样性,我们展示了固定类别下不同模型的多个样本。

在附录 H 的图 15 中,可以看到本文 VAE 的两个样本(通过解码从先验中采样的潜变量得到),呈现出图像纹理的混合。随后我们展示 GIVT-MaskGIT 推理的不同步骤,并观察到与基于 VQ 的模型类似的行为([7, 图 2])。

Table 4语义 HTML 转录
无条件 GIVT-Causal 与已有生成模型在 ImageNet 上的线性探测准确率。GIVT-Causal 与 VIM+ViT(ViT-VQ-GAN)[76] 持平;后者参数量超过其 2 倍、序列长度为其 4 倍(因而 FLOPs 也更高)。Type:潜变量生成模型类型;#Param.:完整潜变量生成模型的参数量。
模型 类型 #令牌 #参数 准确率\uparrow
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]相匹配,后者拥有超过2×2\times倍的 模型参数和4×4\times倍的 序列长度(因此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训练代码方面的支持。

参考文献

参考文献按原论文顺序与书目信息保留。

  1. 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)
  2. 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)
  3. Bao, H., Dong, L., Piao, S., Wei, F.: BEiT: BERT pre-training of image transformers. In: ICLR (2021)
  4. 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
  5. Brock, A., Donahue, J., Simonyan, K.: Large scale gan training for high fidelity natural image synthesis. In: ICLR (2018)
  6. 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)
  7. Chang, H., Zhang, H., Jiang, L., Liu, C., Freeman, W.T.: MaskGIT: Masked generative image transformer. In: CVPR. pp. 11315–11325 (2022)
  8. Chen, M., Radford, A., Child, R., Wu, J., Jun, H., Luan, D., Sutskever, I.: Generative pretraining from pixels. In: ICML. pp. 1691–1703 (2020)
  9. Chen, X., Kingma, D.P., Salimans, T., Duan, Y., Dhariwal, P., Schulman, J., Sutskever, I., Abbeel, P.: Variational lossy autoencoder. In: ICLR (2016)
  10. 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)
  11. Das, A., Kong, W., Sen, R., Zhou, Y.: A decoder-only foundation model for time-series forecasting. arXiv:2310.10688 (2023)
  12. 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
  13. Devlin, J., Chang, M.W., Lee, K., Toutanova, K.: Bert: Pre-training of deep bidirectional transformers for language understanding. NAACL-HLT (2019)
  14. Dhariwal, P., Nichol, A.: Diffusion models beat GANs on image synthesis. NeurIPS pp. 8780–8794 (2021)
  15. Dinh, L., Krueger, D., Bengio, Y.: Nice: Non-linear independent components estimation. In: ICLR (2015)
  16. Dinh, L., Sohl-Dickstein, J., Bengio, S.: Density estimation using real nvp. In: ICLR (2017)
  17. Donahue, J., Simonyan, K.: Large scale adversarial representation learning. In: NeurIPS (2019)
  18. 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)
  19. Eisenach, C., Patel, Y., Madeka, D.: MQTransformer: Multi-horizon forecasts with context dependent and feedback-aware attention. arXiv:2009.14799 (2020)
  20. Esser, P., Rombach, R., Ommer, B.: Taming transformers for high-resolution image synthesis. In: CVPR. pp. 12868–12878 (2020)
  21. Fan, A., Lewis, M., Dauphin, Y.: Hierarchical neural story generation. In: ACL. pp. 889–898 (2018)
  22. Garza, A., Mergenthaler-Canseco, M.: TimeGPT-1. arXiv:2310.03589 (2023)
  23. 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)
  24. 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)
  25. Ho, J., Salimans, T.: Classifier-free diffusion guidance. arXiv:2207.12598 (2022)
  26. Holtzman, A., Buys, J., Du, L., Forbes, M., Choi, Y.: The curious case of neural text degeneration. In: ICLR (2019)
  27. Huh, M., Cheung, B., Agrawal, P., Isola, P.: Straightening out the straight-through estimator: Overcoming optimization challenges in vector quantized networks. In: ICML (2023)
  28. Jacobsen, J.H., Smeulders, A.W., Oyallon, E.: i-revnet: Deep invertible networks. In: ICLR (2018)
  29. Kim, S., Jo, D., Lee, D., Kim, J.: MAGVLT: Masked generative vision-and-language transformer. In: CVPR. pp. 23338–23348 (2023)
  30. Kingma, D.P., Welling, M.: Auto-encoding variational bayes. arXiv:1312.6114 (2013)
  31. Kingma, D.P., Salimans, T., Jozefowicz, R., Chen, X., Sutskever, I., Welling, M.: Improved variational inference with inverse autoregressive flow. NeurIPS (2016)
  32. Kirillov, A., He, K., Girshick, R., Rother, C., Dollár, P.: Panoptic segmentation. In: CVPR. pp. 9404–9413 (2019)
  33. 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)
  34. Kumar, S., Anastasopoulos, A., Wintner, S., Tsvetkov, Y.: Machine translation into low-resource language varieties. In: ACL. pp. 110–121 (2021)
  35. Kumar, S., Tsvetkov, Y.: Von Mises-Fisher loss for training sequence to sequence models with continuous outputs. In: ICLR (2018)
  36. 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)
  37. Ł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)
  38. Li, L.H., Chen, P.H., Hsieh, C.J., Chang, K.W.: Efficient contextual representation learning without softmax layer. arXiv:1902.11269 (2019)
  39. 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)
  40. Li, Y., Mao, H., Girshick, R., He, K.: Exploring plain vision transformer backbones for object detection. In: ECCV. pp. 280–296 (2022)
  41. 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)
  42. 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)
  43. Menick, J., Kalchbrenner, N.: Generating high fidelity images with subscale pixel networks and multidimensional upscaling. arXiv:1812.01608 (2018)
  44. Mentzer, F., Agustsson, E., Tschannen, M.: M2T: Masking transformers twice for faster decoding. In: ICCV (2023)
  45. Mentzer, F., Gool, L.V., Tschannen, M.: Learning better lossless compression using lossy compression. In: CVPR. pp. 6638–6647 (2020)
  46. Mentzer, F., Minnen, D., Agustsson, E., Tschannen, M.: Finite scalar quantization: VQ-VAE made simple. arXiv:2309.15505 (2023)
  47. 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)
  48. Nie, Y., Nguyen, N.H., Sinthong, P., Kalagnanam, J.: A time series is worth 64 words: Long-term forecasting with transformers. In: ICLR (2022)
  49. van den Oord, A., Vinyals, O., Kavukcuoglu, K.: Neural discrete representation learning. NeurIPS (2017)
  50. Parmar, N., Vaswani, A., Uszkoreit, J., Kaiser, L., Shazeer, N., Ku, A., Tran, D.: Image transformer. In: ICML. pp. 4055–4064 (2018)
  51. Peebles, W., Xie, S.: Scalable diffusion models with transformers. arXiv:2212.09748 (2022)
  52. Radford, A., Narasimhan, K., Salimans, T., Sutskever, I.: Improving language understanding by generative pre-training (2018)
  53. 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)
  54. Razavi, A., Van den Oord, A., Vinyals, O.: Generating diverse high-fidelity images with VQ-VAE-2. NeurIPS (2019)
  55. Rombach, R., Blattmann, A., Lorenz, D., Esser, P., Ommer, B.: High-resolution image synthesis with latent diffusion models. In: CVPR (2022)
  56. 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)
  57. Sadeghi, H., Andriyash, E., Vinci, W., Buffoni, L., Amin, M.H.: PixelVAE++: Improved pixelvae with discrete prior. arXiv:1908.09948 (2019)
  58. Sajjadi, M.S., Bachem, O., Lucic, M., Bousquet, O., Gelly, S.: Assessing generative models via precision and recall. NeurIPS (2018)
  59. Salimans, T., Karpathy, A., Chen, X., Kingma, D.P.: PixelCNN++: Improving the PixelCNN with discretized logistic mixture likelihood and other modifications. In: ICLR (2016)
  60. Sauer, A., Schwarz, K., Geiger, A.: StyleGAN-XL: Scaling StyleGAN to large diverse datasets. In: SIGGRAPH (2022)
  61. Silberman, N., Hoiem, D., Kohli, P., Fergus, R.: Indoor segmentation and support inference from RGBD images. In: ECCV. pp. 746–760 (2012)
  62. 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)
  63. Strudel, R., Garcia, R., Laptev, I., Schmid, C.: Segmenter: Transformer for semantic segmentation. In: CVPR. pp. 7262–7272 (2021)
  64. Tokarchuk, E., Niculae, V.: On target representation in continuous-output neural machine translation. In: ACL (2022)
  65. Tokarchuk, E., Niculae, V.: The unreasonable effectiveness of random target embeddings for continuous-output neural machine translation. arXiv:2310.20620 (2023)
  66. Tomczak, J., Welling, M.: Vae with a vampprior. In: AISTATS. pp. 1214–1223 (2018)
  67. Tschannen, M., Bachem, O., Lucic, M.: Recent advances in autoencoder-based representation learning. arXiv:1812.05069 (2018)
  68. Tschannen, M., Kumar, M., Steiner, A., Zhai, X., Houlsby, N., Beyer, L.: Image captioners are scalable vision learners too. In: NeurIPS (2023)
  69. Vahdat, A., Andriyash, E., Macready, W.: Dvae#: Discrete variational autoencoders with relaxed boltzmann priors. NeurIPS (2018)
  70. Vahdat, A., Kautz, J.: NVAE: A deep hierarchical variational autoencoder. NeurIPS pp. 19667–19679 (2020)
  71. Van Den Oord, A., Kalchbrenner, N., Kavukcuoglu, K.: Pixel recurrent neural networks. In: ICML. pp. 1747–1756 (2016)
  72. Vaswani, A., Shazeer, N., Parmar, N., Uszkoreit, J., Jones, L., Gomez, A.N., Kaiser, Ł., Polosukhin, I.: Attention is all you need. NeurIPS (2017)
  73. 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)
  74. 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)
  75. 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)
  76. 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)
  77. 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)
  78. Zhai, X., Kolesnikov, A., Houlsby, N., Beyer, L.: Scaling vision transformers. In: CVPR. pp. 12104–12113 (2022)
  79. 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:从dd每通道标量GMM或单个高斯预测更改为dd-变量GMM,具有因子化分量 (..,分量分布是dd-变量高斯分布,具有对角 协方差)。

    • 引入了适配器。

    • 大型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,我们基于β\beta-VAE获得了最先进的结果,该VAE不包含任何不可微操作,也不需要VQ文献中的任何高级技巧。[7, 27, 33, 37, 49].

连续潜变量与无限词汇量在 理论上,连续(..,实值)潜变量/标记在将它们精确映射到 离散代码时总是意味着无限的词汇量。换句话说,无限词汇量是连续潜变量的结果。在实践中,连续潜变量以有限精度表示,例如.., float32,但使用 原始的32位表示作为离散代码仍然意味着对于单个潜变量维度来说,词汇量大小2324.32^{32} \approx 4.3B 不切实际地大。GIVT直接 建模连续潜变量,避免了物化这个词汇量。

架构细节

用于图像生成的不同模型变体的架构细节列于表5。我们依赖标准的pre-LN设置(参见,例如,UViMvtt.py以获取具体实现)。对于UViM实验,我们使用带有默认配置的GIVT-Causal,并在每个自注意力层之后插入一个交叉注意力层,如[33]中所述,以融合从输入RGB图像中提取的视觉特征。

适配器通过堆叠8个卷积、双射的iRevNet块[28]构建,隐藏通道维度为4d4d(导致112k 额外参数用于d=16d=16)。 每个块由3个卷积层组成,与GroupNorm和 ReLU交替,并且具有相同的输入和输出形状w×h×dw \times h \times d (.., 适配器应用于VAE潜在变量zz在将其重塑为序列之前)。 我们基于参考实现iRevNet.py构建iRevNet块, 将BatchNorm替换为GroupNorm并移除下采样层。

Table 5语义 HTML 转录
不同图像生成模型变体的架构细节。我们还列出了训练相应 VAE 时使用的 KL 权重 β。Base 架构用于探索特征维度 d、mixture 数 k 与 β(图 5)。默认使用 β=5·10^-5,Large 模型使用 β=10^-5。* 512px 的 GIVT-Causal-L 采用 space-to-depth 变换:把相邻的两个 16 维特征向量堆叠成 32 维向量,使序列长度为 512 而不是 1024。
模型 大小 分辨率 dd kk 宽度 深度 MLP 头数 参数 令牌 丢弃 β\beta
因果 基础 256 4324\ldots32 1321\ldots32 768 12 3072 12 86M 256 0.1 0.2521040.25\ldots2\cdot10^{-4}
因果 默认 256 16 16 1024 24 4096 16 304M 256 0.2 51055\cdot10^{-5}
因果 大型 256 16 16 1536 48 8192 16 1.67B 256 0.3 10510^{-5}
因果{}^* 大型 512 32 32 1536 48 8192 16 1.67B 512 0.1 10510^{-5}
MaskGIT 默认 256 16 16 1024 24 4096 16 304M 256 0.4 51055\cdot10^{-5}
MaskGIT 默认 512 16 16 1024 24 4096 16 304M 1024 0.4 51055\cdot10^{-5}
Figure 8原论文图面与译注
训练Base大小的GIVT-causal模型时的损失(NLL)曲线,来自图5,其中潜在维度为d=16d=16,作为混合数kk和KL权重β\beta的函数,在β\beta-VAE。减少β\beta或增加kk会导致损失减少,这并不总是转化为更低的FID。损失是使用distrax计算的,如第C节所述。 替代实现可能导致此处缩放1d\frac{1}{d} (116\frac{1}{16})。

训练细节

损失函数和实现细节

我们从GIVT对序列中每个特征向量或软标记的分布进行建模的情况开始,条件是序列中所有先前的特征向量,使用一个kk-混合高斯模型,其分量在通道间独立(..,这些分量是 具有对角协方差矩阵的多元高斯分布)。 对于大小为BBdd维特征序列,长度为LL,预测的GIVT输出是一个B×L×(2dk+k)B \times L \times (2dk + k)张量y=[m;s;π]y = [m; s; \pi]在哪里m,sm, sB×L×dkB \times L \times dk张量及π\pi是一个B×L×kB \times L \times k张量,全部沿最后一个维度堆叠, 分别包含均值、标准差和混合权重。 我们假设ss被一个小的正常数下界限制ϵ=105\epsilon = 10^{-5} (例如..,通过对相应的网络输出应用softplus函数并裁剪到ϵ\epsilon),并且π\pi的条目为非负。mm[m(1);;m(k)][m^{(1)}; \ldots; m^{(k)}]组成,其中m(n)m^{(n)}是GMM分量沿最后一个轴堆叠的B×L×dB \times L \times d均值张量;ss以相同的方式组成。混合权重假定在分量之间归一化(例如..,通过softmax函数),.., n=1kπb,(n)=1\sum_{n=1}^k \pi_{b,\ell}^{(n)} = 1对于 所有b,b,\ell.

然后,特征序列的负对数似然(NLL){zb}b=1B\{z_b\}_{b=1}^B (..训练图像在通过VAE编码器嵌入并应用重参数化之后)相对于GIVT预测的分布p~\tilde p可以写成 为b=1Blog(p~(zb))=b=1Blog(=1L(n=1kπb,(n)c=1dN(zb,,cmb,,c(n),sb,,c(n))))=b=1B=1Llog(n=1kπb,(n)c=1dN(zb,,cmb,,c(n),sb,,c(n))),\begin{aligned} &-\sum_{b=1}^B \log(\tilde p(z_b)) \\ &=-\sum_{b=1}^B \log\left(\prod_{\ell = 1}^L \left(\sum_{n=1}^k \pi_{b, \ell}^{(n)} \prod_{c = 1}^d \mathcal N(z_{b, \ell, c} | m_{b, \ell, c}^{(n)}, s_{b, \ell, c}^{(n)})\right)\right) \\ &=-\sum_{b=1}^B \sum_{\ell = 1}^L \log\left(\sum_{n=1}^k \pi_{b, \ell}^{(n)} \prod_{c = 1}^d \mathcal N(z_{b, \ell, c} | m_{b, \ell, c}^{(n)}, s_{b, \ell, c}^{(n)})\right), \end{aligned}其中N(zm,s)\mathcal N(z | m, s)是高斯密度N(zm,s)=1s2πe12(zms)2.\begin{aligned} \mathcal N(z | m, s) = \frac{1}{s \sqrt{2 \pi}} e^{-\frac12 \left(\frac{z - m}{s}\right)^2}. \end{aligned}当使用单个高斯而不是GMM来建模每个输出通道时(..k=1k=1使得πb,(1)=1\pi_{b, \ell}^{(1)} = 1对所有b,b, \ell成立)公式2变为b=1B=1Lc=1d12(zb,,cmb,,csb,,c)2+log(sb,,c)+12log(2π).\begin{aligned} \sum_{b=1}^B \sum_{\ell = 1}^L \sum_{c = 1}^d \frac12 \left(\frac{z_{b, \ell, c} - m_{b, \ell, c}}{s_{b, \ell, c}}\right)^2 + \log(s_{b, \ell, c}) + \frac12 \log (2\pi). \end{aligned}虽然这个损失简化为对序列维度\ell的求和,但请注意mb,,c,sb,,cm_{b, \ell, c}, s_{b, \ell, c}(在GMM情况下为πb,,c\pi_{b, \ell, c}) 由zb,1,,zb,1z_{b,1},\ldots, z_{b,\ell-1}通过 GIVT(设置zb,0z_{b,0}到一个学到的[CLS][BOS]向量)。此外可以看出,令 sb,,cs_{b,\ell,c} 的下界为 ϵ>0\epsilon>0,可避免某些 zb,z_{b,\ell} 落入 p~\tilde p 的低密度区域时出现极大损失值。

8展示了不同KL权重的VAE的训练损失曲线β\beta以及kk,遵循图4的设置。增加kk并减少β\beta导致更低的损失值。但请注意, 较低的损失并不总是导致较低的采样FID (见图4)。重要的是,初始和 最终损失值,以及整个训练过程中训练损失的相对减少量,可能与离散交叉熵或NLL通常观察到的值有显著差异,例如,在 语言建模中。

对于多变量情况,其中特征通道之间的依赖关系被建模,GIVT预测一个d×dd \times d协方差矩阵Sb,(n)S^{(n)}_{b,\ell}对于每个混合成分nn,且数据的负对数似然变为b=1B=1Llog(n=1kπb,(n)N~(zb,mb,(n),Sb,(n))),\begin{aligned} -\sum_{b=1}^B \sum_{\ell = 1}^L \log\left(\sum_{n=1}^k \pi_{b, \ell}^{(n)} \tilde{ \mathcal N}(z_{b, \ell} | m_{b, \ell}^{(n)}, S_{b, \ell}^{(n)})\right), \end{aligned}其中N~\tilde{\mathcal N}是一个多元 高斯密度。

对于所有损失变体,MaskGIT版本的GIVT将求和LL简化为对索引ll对应于被掩码的位置, 并且损失通过掩码令牌的数量进行归一化。

最后,损失计算和采样可以轻松地使用 专门的深度学习包实现,例如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节。

Table 6语义 HTML 转录
ImageNet CNN-based VAE编码器/解码器的超参数。请注意,嵌入的32个特征被分为16个均值和16个尺度,因此重参数化后我们的实际表示有d=16d=16个通道。
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%

  • 调整大小至目标分辨率(256×256256{\times}256512×512512{\times}512)使用双三次滤波器 并启用抗锯齿

  • 以0.5的概率随机水平翻转图像

全景 分割与深度估计

对于我们的UViM全景分割和深度估计实验, 我们采用公开的UViM GitHub 仓库并且仅替换阶段 I 中的 VQ 层,调整阶段 II 中的 transformer 解码器,并相应修改损失函数。

DB-CFG 实现

我们在图 中展示了用于 DB-CFG 的拒绝采样器的 JAX 实现。11.

无条件图像生成

在表7中,我们展示了无条件 ImageNet 生成的 FID 结果。

Figure 9原论文图面与译注
在无条件 256×256256\times256 ImageNet 上训练的 GIVT-Causal,其各层的 25-shot 线性探测准确率。第 10–20 层在所有数据集上都有较高准确率。
Figure 10原论文图面与译注
不同采样策略和模型变体(GIVT-Causal-L)对 Precision 与 Recall 的影响,补充图 5 的 FID 结果。增加 mixture 数量 kk 并加入 adapter 会产生叠加效应。DB-CFG 是提高所有模型变体 Precision 最有效的采样策略。

探测中间表示

训练的 GIVT-causal 中提取中间表示。256×256256 \times 256无条件 通过将给定层的输出在序列维度上进行平均池化(得到一个维度为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 推理成本,我们得到以下近似比较,使用NNGIVT-Causal-L+A的FLOPs作为 参考。MaskGIT和GIVT-MaskGIT在推理时所需的FLOPs几乎相同。

Table 7语义 HTML 转录
无条件 256×256256\times256 ImageNet 生成结果。指标采用标准 ADM 评测套件,FID 相对训练集计算。MAGE [39] 样本由作者使用其公开 GitHub 代码生成;MAGE Large 包含 latent encoder-decoder,因此参数量显著高于 decoder-only 的 GIVT-Causal。
模型推理FID \downarrowPrecision \uparrowRecall \uparrow
ADM [14]26.20.610.63
MAGE (ViT-B)temp =6.0=6.012.0802130225097240.619880.6364
MAGE (ViT-L)temp =6.0=6.09.95289081917440.659280.6589
GIVT-Causal(本文)k=1k=1t=0.9t=0.919.92260.62590.6013
GIVT-Causal(本文)k=16k=16t=0.95t=0.9517.69950.65870.6282
GIVT-Causal-L(本文)k=16k=16t=0.95t=0.9511.02040.72070.6015
Table 8语义 HTML 转录
GIVT与文献中自回归基线在FLOPs上的近似比较。
模型 指导 步骤 参数 FLOPs FID
GIVT-Causal-L+A 256 1.67B NN 3.46
VQGAN [20] 256 1.4B N\simeq N 17.04
ViT-VQGAN-L [76] 1024 1.7B 4N\simeq 4 N 4.17
GIVT-Causal-L+A DB-CFG 256 1.67B 2N\simeq2N 2.59
VQGAN [20] CG=0.05=0.05 256 1.4B 20N\simeq 20 N 5.20
ViT-VQGAN-L [76] CG=0.5=0.5 1024 1.7B 8N\simeq 8 N 3.04

更广泛的影响

本文描述的方法修改了仅解码器的Transformer模型,使其能够生成特征向量序列,这可以看作是对Transformer模型基础的研究努力,具有许多潜在的下游应用。更广泛的影响将主要取决于下游应用及其受这项工作影响的方式。

在这里,我们专注于视觉数据的生成。对于类条件图像生成,我们依赖ImageNet,所得模型通过类标签提供对生成图像内容的基本控制。这些模型不允许对图像内容进行细粒度控制,或对现有图像进行操作。ImageNet固有的偏见和问题已被充分研究和记录[Uday Prabhu和Birhane,2020]。我们预计我们的图像生成模型可能反映其中一些问题,并建议用户在使用和部署模型时谨慎考虑这些问题。

将我们的图像生成模型扩展到文本到图像界面相对直接,这使得可以在从网络收集的更大且较少策划的图像/文本数据集上进行训练。此类模型允许更细粒度的控制,因此可能被滥用于恶意目的,如传播错误信息或身份盗窃。此外,此类模型可能反映大型非策划网络数据集中的偏见,如有害的刻板印象。因此,此类模型的部署和发布必须更加谨慎。

Figure 11原论文图面与译注
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拒绝采样器实现。
Figure 12原论文图面与译注
来自VAE编码器分布的样本:第一列是ImageNet验证集中的图像。我们用VAE编码器对其进行编码,获得近似后验分布,并从中采样4次。然后将得到的4个潜在序列解码为图像(最后4列)。输入图像的语义布局保持得很好,只有低层纹理发生变化。

额外的视觉示例

在图12中,我们展示了当向VAE输入ImageNet验证集中的图像时,我们的VAE的重建结果。当多次从编码器分布中采样时,我们看到重建中的低层纹理变化。图13显示了固定标签的多个样本,以展示我们模型的多样性并与基线进行比较。在图14以及16我们展示了来自不同256×256256 \times 256GIVT-Causal 变体的样本。在图17中,我们展示了 MaskGIT 对于我们在图1中展示的相同类别的样本。图18展示了我们512×512512\times512GIVT-MaskGIT模型的样本。我们 探索在采样中途更改标签(..,在生成图像上半部分之后)对GIVT-Causal的影响,见图19.

对于本文的 GIVT-Casual UViM 模型和 VQ baseline,视觉输出见图 20 和图 21

Figure 13原论文图面与译注
GIVT-Causal(本文)GIVT-MaskGIT(本文)验证集
BigGAN [5](来自 [7, 图 9]MaskGIT (VQ) [7, 图 9]
类别 980 的多个样本,用于展示样本多样性(不使用 DB-CFG)。作为参考,此处复制了 [7, 图 9] 的示例。
Figure 14原论文图面与译注
GIVT-Causal(t=0.95t=0.95,DB-CFG =0.4=0.4)在与图 1 相同的 10 个 ImageNet 类别上的 256×256256\times256 样本。
Figure 15原论文图面与译注
VAE 样本 步骤 1 步骤 4 步骤 8 步骤 16
第一列:从 VAE 先验中采样得到的两张图像,导致低层图像纹理的混合。其余列:可视化 GIVT-MaskGIT 对两个 ImageNet 类别(947, 94)的输出,经过 1、4、8、16 次推理步骤。正如预期,随着推理步骤的增加,样本开始变得更加连贯。
Figure 16原论文图面与译注
GIVT-Causal-L+A(t=0.95t=0.95,DB-CFG =0.4=0.4)在与图 1 相同的 10 个 ImageNet 类别上的 256×256256\times256 样本。
Figure 17原论文图面与译注
GIVT-MaskGIT(tC=60t_C=60,DB-CFG =0.1=0.1)在与图 1 相同的 10 个 ImageNet 类别上的 256×256256\times256 样本。
Figure 18原论文图面与译注
GIVT-MaskGIT(tC=140t_C=140)在与图 1 相同的 10 个 ImageNet 类别上的 512×512512\times512 样本。
Figure 19原论文图面与译注
在采样过程中途 更改标签(..,在生成上半部分之后)用于 GIVT-Causal:第一列对上半部分和下半部分使用相同的标签,其他列切换为不同的标签。顶行标签: 上半部分为“金毛寻回犬”(207);下半部分为“水獭”(360)、“大猩猩” (366)、“骆驼”(355)。底行标签:上半部分为“鸟屋”(448);下半部分为“船屋”(449)、“灯塔”(437)、 “面包店”(415)。请注意,对于每一行,顶行潜在变量始终相同,但整体色彩平衡在RGB输出中可能不同,因为VAE解码器 (可能由于GroupNorm层)。
Figure 20原论文图面与译注
本文 UViM 模型的‘COCO Panoptic Segmentation’可视化示例:COCO Panoptic 2017 上基于 GIVT-Causal 的 UViM 与标准 VQ-based UViM 的全景分割结果。
Figure 21原论文图面与译注
输入 真实值 GIVT-Causal UViM 基于VQ的UViM
输入 地面真值 GIVT-Causal UViM 基于VQ的UViM
基于GIVT-Causal和标准VQ的UViM在NYU Depth v2上的深度估计视觉示例。
  1. 我们在第节讨论连续潜变量与无限词汇表之间的关系。A.↩︎

LLM WIKI · CONTEXT READER

AI 论文解读

DeepSeek V4 Flash

Enter 发送 · Shift + Enter 换行 · Esc 关闭