UDT:用数据自适应Token缩减统一U-Net与扩散Transformer

快速导读:提出U-Net扩散Transformer (UDT),通过数据自适应token合并实现无参数下/上采样,在保持token特征维度的同时显著加速训练收敛并提升生成性能。

扩散Transformer (DiT) 已成为生成建模的主流架构,但其各向同性的结构存在一个根本性问题:编码阶段过深而有效解码阶段过短,导致表示学习不均衡。现有U-Net风格的DiT试图通过空间下采样引入多尺度建模来解决这一问题,但依赖固定局部邻域压缩token,破坏了token间的全局交互,导致表示质量下降。

本文提出U-Net扩散Transformer (UDT),核心创新在于用数据自适应token合并 (ToMe) 替代传统的空间下采样。在编码器阶段,UDT根据token间的语义相似性逐步合并token,减少序列长度;在解码器阶段,利用记录的合并索引精确恢复token。整个过程保持token隐藏维度不变,仅压缩空间维度。这一设计使UDT在ImageNet 256×256上以80个epoch达到7.7 FID(无CFG),超越了SiT-XL在1400个epoch的7.9 FID,训练收敛速度提升近20倍。

英文题目:UDT: Reconciling U-Nets and Diffusion Transformers with Data-Adaptive Token Reduction

论文出处:arXiv 每日论文精选 · arXiv:2608.01298

原始论文:PDF / 论文页面

这篇论文解决什么问题?

扩散Transformer的瓶颈在于其各向同性架构。以SiT为例,24层Transformer中前18层用于编码,仅6层用于解码。这种不对称导致模型在编码阶段过度压缩信息,而解码阶段缺乏足够的容量来恢复细节。PCA可视化和线性探测实验(第4页)清晰地展示了这一问题:SiT的中间层特征呈现模糊的团状分布,线性探测准确率在深层急剧下降。

U-DiT等先前工作尝试通过空间下采样构建U形架构,但存在两个关键缺陷。第一,固定邻域的下采样(如2×2合并)忽略了token间的语义相似性,可能将语义相关的token强行分离,或将无关token合并在一起。第二,空间下采样通常伴随通道维度的变化(如增加通道数补偿信息损失),这破坏了token特征维度的一致性,导致与cross-attention、REPA等组件的兼容性问题。

数据自适应token合并 (ToMe) 提供了一条新路径。ToMe最初用于加速预训练ViT的推理,通过计算token间的余弦相似度,将最相似的token对合并。但将其引入扩散Transformer的训练阶段面临独特挑战:训练时token分布随噪声水平动态变化,且需要可微分的合并/解合并操作以支持端到端训练。

UDT的关键洞察在于:ToMe的合并操作天然适合作为Transformer兼容的下采样方案。它不改变token的特征维度,仅减少序列长度,因此可以无缝集成到任何Transformer架构中。同时,合并决策基于语义相似性而非空间位置,能够自适应地保留细粒度细节(如前景物体)同时压缩冗余区域(如背景)。

核心创新

  • 首次将数据自适应token合并机制引入扩散Transformer训练,作为Transformer兼容的下/上采样方案。
  • 设计了一种保持token特征维度的U-Net架构,避免了现有方法中通道维度变化带来的表示不一致和工程开销。
  • 证明了该架构可作为即插即用的替代方案,无缝集成到DiT、JiT、MMDiT等多种变体中。

方法概览

提出UDT,一种U-Net形状的扩散Transformer。核心是在编码器阶段通过数据自适应token合并 (ToMe) 逐步减少token序列长度,在解码器阶段利用记录的合并索引精确恢复token。整个过程保持token隐藏维度不变,仅压缩空间维度。

  • 架构设计:UDT采用对称的U形结构,编码器阶段通过ToMe逐步将token序列长度减半(如从256降至16),解码器阶段通过记录的合并索引逐步恢复。跳跃连接将编码器各层的未合并token直接传递到解码器对应层,确保信息不丢失。
  • Token合并机制:在每个合并层,计算所有token对之间的余弦相似度,选择最相似的r对进行合并。合并方式为对两个token的特征取加权平均,权重由它们各自包含的原始token数量决定。合并索引被存储用于后续的解合并操作。
  • Token解合并:解码器阶段,根据存储的合并索引,将合并后的token复制回其原始位置。这一操作完全可逆且无参数,确保解码器能够精确恢复编码器压缩前的token结构。
  • 与现有组件的兼容性:由于UDT保持token特征维度不变,它可以无缝集成REPA表示对齐、cross-attention条件注入、SwiGLU激活函数、RoPE位置编码等高级技术。唯一的限制是RoPE只能在未合并的层使用,因为合并操作破坏了原始位置索引。
  • 高级技术集成 (UDT+):在基础UDT上,论文探索了log-normal时间步采样、SwiGLU激活、RoPE位置编码、VA-VAE等技术的组合效果。消融实验(第9页表5)表明,这些技术各自带来约0.1-0.3 FID的提升,组合使用效果更佳。
  • 多架构即插即用:UDT可作为DiT(ϵ预测)、JiT(像素级预测)、MMDiT(文生图)等变体的直接替代方案。实验表明(第10页表9),在相同训练设置下,UDT版本始终优于原始各向同性版本。

逐图理解论文

空间下采样的困境

空间下采样的困境
Figure 3: Visualization of token reduction: (a) Data-adaptive token reduction (ours) gradually merges semantically similar tokens (e.g., background regions), allowing fine-grained details to be retained. (b) Spatial token reduction with factor-2 downsampling, which relies on fixed local neighborhoods and ignores similarity between tokens across the image, resulting in the loss of fine details.

图3对比数据自适应token合并与空间下采样的效果差异。(a) UDT的合并策略智能地保留前景细节(如鸟的羽毛),合并背景冗余区域。(b) 空间下采样无差别地压缩所有区域,导致细粒度信息丢失。

研究空白与核心区分

研究空白与核心区分
Figure 2: Representation Analysis. (a) PCA visualization of intermediate layers. Features are extracted from the bottleneck layer (e.g., layer 18 of 24 for SiT-L/2, layer 16 of 32 for SiT↓, and layer 12 of 24 for UDT). Features with 112 tokens in our method are unmerged for visualization. (b) Linear probing evaluation on pretrained models across layers. All experiments are conducted at noise level t = 0.1 using Large models trained for 80 epochs on ImageNet 256 × 256.

图2通过PCA可视化和线性探测评估中间层表示质量。(a) UDT的特征呈现清晰的类别聚类,而SiT和SiT↓的特征则模糊分散。(b) 线性探测准确率曲线显示UDT在各层均保持较高的表示质量,尤其在深层优势明显。

UDT的架构设计与验证

UDT的架构设计与验证
Figure 1: (a) Architecture of the proposed method UDT. It preserves the token hidden dimension (D) while gradually reducing and restoring the token sequence length via data-adaptive token merging and unmerging. (b) FID vs. Epoch on XL models for ImageNet 256 × 256 without classifier-free guidance. UDT (our baseline, marked in light pink) and UDT+ (baseline + advanced techniques, marked in pink), achieve 7.7 FID at 80 epochs and 7.0 FID at 60 epochs, without additional regularization or VAE modification, outperforming SiT-XL’s 7.9 FID at 1400 epochs. Our method with REPA achieves 7.6 FID at just 40 epochs, converging significantly faster than the SiT baseline and REPA.

图1展示UDT的整体架构和训练效率对比。(a) UDT保持token隐藏维度D不变,通过数据自适应合并/解合并逐步减少和恢复token序列长度。(b) FID曲线显示UDT在80个epoch即超越SiT在1400个epoch的性能,收敛速度提升显著。

结论与意义

结论与意义
Figure 4: FID vs. Epoch on ImageNet 256 × 256 without CFG.

图4展示FID随训练epoch的变化曲线。UDT和UDT+的收敛速度远超SiT基线,且UDT+在60个epoch即达到7.0 FID,体现了架构与高级技术的协同效果。

实验如何设计?

  • 主要实验在ImageNet 256×256上进行类别条件图像生成,使用B/2、L/2、XL/2三种模型规模。
  • 评估指标包括FID、Inception Score、Precision、Recall,分别在无CFG和有CFG条件下测试。
  • 表示分析实验使用PCA可视化和线性探测,在噪声水平t=0.1下评估中间层特征质量(第4页)。
  • 消融实验研究token合并策略(瓶颈token数、合并方式)、高级技术组件的影响(第6页表1,第9页表5)。
  • 泛化实验包括:有限数据(10% ImageNet)、更长token序列(patch size 1)、高分辨率(512×512)、以及作为DiT/JiT/MMDiT的即插即用替代(第10页)。

关键结果与论文证据

  • 训练效率:UDT-XL/2在80个epoch达到7.7 FID(无CFG),而SiT-XL需要1400个epoch才达到7.9 FID,训练加速约17.5倍(第7页表2)。
  • 表示质量:PCA可视化显示UDT的中间层特征呈现清晰的类别聚类结构,线性探测准确率显著高于SiT和U-DiT(第4页图2)。
  • Token合并可视化:数据自适应合并能够智能地保留前景物体的细粒度token,同时合并背景中的冗余token;而空间下采样则无差别地丢失细节(第5页图3)。
  • 与REPA结合:UDT+REPA在40个epoch达到7.6 FID(无CFG),在320个epoch配合CFG达到1.38 FID(第8页表4,第7页表3)。
  • 高分辨率生成:在512×512 ImageNet上,UDT+仅需200个微调epoch即达到1.58 FID(第9页表7)。
  • 有限数据场景:在10% ImageNet子集上,UDT-L/2在500个epoch达到10.6 FID,显著优于SiT-L/2的14.2 FID(第10页表8)。
  • 即插即用:在DiT、JiT、MMDiT三种架构上,UDT版本均取得一致的性能提升,验证了其通用性(第10页表9)。
  • 计算效率:尽管UDT引入了跳跃连接的轻微参数增加,但token序列长度的减少使总体GFLOPs降低,实际训练墙钟时间更短(第18页图5)。

阅读时需要注意

  • 位置编码限制:Token合并/解合并操作破坏了原始位置索引,因此RoPE等位置编码只能在未合并的层使用,限制了位置信息在深层网络中的利用。
  • 参数开销:跳跃连接导致参数量轻微增加,在极轻量级部署场景(如移动端)可能成为瓶颈。
  • 合并策略的固定性:当前合并比例(如每层合并50%的token)是预先设定的超参数,缺乏根据输入复杂度自适应调整的机制。
  • 评估指标局限:FID比较需注意不同模型使用的VAE(SD-VAE vs. VA-VAE)、CFG权重和引导区间可能不同,直接数值对比需谨慎。

关联工作

  • DiT (Peebles & Xie, 2023):基线各向同性扩散Transformer架构,UDT作为其U-Net替代方案。
  • U-DiT:通过空间下采样构建U形扩散Transformer的先前工作,UDT在表示质量和训练效率上均显著超越。
  • ToMe (Bolya et al., 2023):UDT所采用的数据自适应token合并机制的基础方法,原用于ViT推理加速。
  • REPA (Yu et al., 2024):通过表示对齐提升扩散Transformer训练效率的方法,UDT与其结合后效果进一步提升。

发表评论