← 总目录 / 板块四 · 多模态模型的发展
板块四 · 多模态模型的发展

第35篇 · DiT

用Transformer替换U-Net:扩散模型的"架构解放"与Gflops标度律
Scalable Diffusion Models with Transformers · William Peebles (UC Berkeley) & Saining Xie (NYU),工作完成于Meta AI实习期间 · 2022 · arXiv:2212.09748

来源声明:本页所有数字均复述自论文原文(arXiv),未做外部验证;标注【解读者补充】【解读者观点】的内容不来自论文。

一、全局大图

摘要拆解对照表(兼导航)

摘要短语对应原文章节本页解读位置
"训练latent diffusion,用作用于latent patches的transformer替换U-Net"§3预备知识 / §4实验设置第三章·逐章精读 A
"以Gflops衡量前向复杂度分析可扩展性"§4.2扩展分析(12模型×Table 4)第三章·逐章精读 C
"更高Gflops的DiT始终更低FID"§4.2核心发现第三章·逐章精读 C
"DiT-XL/2在ImageNet 256²与512²超越所有先前扩散模型"§4.3 SOTA结果(FID 2.27 / 3.04)第三章·逐章精读 D
"四种DiT block设计对比"§4.1消融(adaLN-Zero胜出)第三章·逐章精读 B

主要贡献与证据强度预评

#论文声称的贡献预评证据强度
1纯transformer(无U-Net归纳偏置)可作扩散backbone达到SOTA:256² FID 2.27、512² FID 3.04强实验支撑 表2/表3主结果
2Gflops是比参数量更好的复杂度代理;Gflops与FID-50K强负相关强实验支撑 Table 4十二个模型系统扫描
3adaLN-Zero条件注入显著优于in-context/cross-attn/vanilla adaLN强实验支撑 四变体同预算消融
4扩大采样算力无法弥补模型算力不足较强但样本少 少量配对比较(L/2千步 vs XL/2百步)
5"U-Net归纳偏置对扩散性能不关键"这一反直觉主张由贡献1间接支撑 无直接机制分析
知识依赖主干线:ViT最佳实践 → LDM框架(8×下采样VAE,z=32×32×4)→ patchify(p决定token数,枢纽节点:p=2让token翻4倍,是"加算力不加参数"的关键机关)→ 四种条件注入block → Gflops-FID标度律 → SOTA结果。前置依赖:第31篇DDPM、第32篇ViT、第34篇LDM——本文站在三者肩膀上。
推荐阅读路线。必读主线:§3.1扩散预备(可快读)→ §3.2设计空间(patchify+四种block)→ §4.1 block消融 → §4.2扩展分析(全文灵魂)→ 表2/3主结果 → §6结论。可跳读支线:附录实现细节(复现才需要)、512²部分(结论一句话)。跳过§3.2 patchify的代价:看不懂"Gflops翻倍而参数几乎不变"这个全文最重要的观察。

二、失效模式先行:U-Net的一统天下与新问题

【论文声称,§1】过去五年各领域架构已被transformer主导,唯独扩散模型仍死守卷积U-Net——它源自PixelCNN++,主体是ResNet块加低分辨率处的空间自注意力。Dhariwal & Nichol虽然消融过归一化层等细节,但高层设计从未被质疑。由此产生一个悬而未决的问题:扩散模型的性能究竟有多少来自U-Net特有的归纳偏置?如果换成"无偏见"的标准transformer会怎样?更进一步:NLP里验证过的scaling行为在扩散模型上是否存在?本文的回答方式不是辩论而是实测:搭一个严格遵循ViT实践、只换backbone的对照体系,扫出12个配置看趋势。

三、逐章精读

A. 设计空间:patchify与模型规格(原文§3.1–3.2)

来源:论文§3.1、§3.2。学习目标:学完你能(a)写出token数公式并解释"p减半→Gflops至少×4";(b)背出四个模型规格的层宽配置;(c)说明为什么用现成Stable Diffusion VAE。

维度流:256²图像经f=8 VAE编码为 z∈ℝ32×32×4;patch大小p把空间网格切成 T=(I/p)² 个token(I=32):p=8→16个token,p=4→64个,p=2→256个。每个token线性嵌入到维度d,加标准ViT频率位置编码。

T = (I/p)²  总Gflops ≳ O(T²·d)(自注意力)+ O(T·d²)(MLP)

关键机关【原文】p减半使T翻4倍,从而至少使总Gflops翻4倍,但对参数量几乎无影响——参数主要花在d和层数上,token数只改变激活规模。这就是后文"参数量不能预测质量"的机制根源。

模型层数N隐藏维d头数Gflops(I=32,p=4)参数量(p=2时)
DiT-S1238461.433M
DiT-B12768125.6130M
DiT-L2410241619.7458M
DiT-XL2811521629.1675M

配置覆盖0.3–118.6 Gflops。XL是作者新增的最大档位。末层把每个token解码为 p×p×2C 张量(同时输出噪声预测和对角协方差预测)。

训练配方【原文§4】:AdamW恒定学习率1×10⁻⁴、无权重衰减、batch 256、唯一增强为水平翻转、EMA衰减0.9999、所有规模共用同一套超参且未做调优;VAE直接用Stable Diffusion的现成预训练权重;保留ADM的1000步线性方差调度。JAX实现,TPU-v3 pods,XL/2约5.7迭代/秒(TPU v3-256)。

常见误读:"DiT发明了新架构"。错——它刻意什么都不发明:ViT原封不动、VAE借来、扩散超参抄ADM。它的贡献是把"什么都不改也能SOTA"这件事本身变成了证据。【解读者观点:这种克制正是它能干净回答架构问题的原因。】

一句话记住本节:p是这台机器的油门——拧小它,算力涨、参数不涨。

闭卷自检:32×32×4的latent在p=2下有多少token?四个规格中哪个是新加的?

B. 四种条件注入block:adaLN-Zero为何碾压(原文§3.2.2、§4.1)

来源:论文§3.2.2与§4.1消融。学习目标:学完你能说出四种设计的注入位置、各自的Gflops开销排序,并用零初始化原理解释adaLN-Zero的优势。

失效模式先行:t(时间步)和c(类别标签)只是两个向量,而transformer的主战场是token序列——怎么把这两个"全局旋钮"接进28层网络?接得笨,要么贵(cross-attention +15% Gflops)要么弱(in-context垫底)。

设计机制Gflops开销FID-50K(400K步,无cfg)
In-contextt,c嵌入作为额外两个token追加到序列,末块后移除≈119.37 G(可忽略)35.24
Cross-attentiont,c拼成长度2的序列,自注意力后加交叉注意力层137.62 G(最贵,约+15%)26.14
AdaLN用t+c嵌入回归LayerNorm的γ、β118.56 G(最省)25.21
AdaLN-ZeroadaLN基础上再回归残差缩放α;MLP输出零初始化,每个块初始为恒等函数118.64 G(可忽略)19.47

全部对比基于DiT-XL/2同预算400K步。adaLN-Zero的FID几乎只有in-context的一半;零初始化相对vanilla adaLN再降近6分。

直觉类比:adaLN-Zero像给每个残差支路装了一个"从静音渐起"的推子——训练开始时所有支路音量为零(整网=恒等映射),随训练按需拉起。【类比在哪里失效】推子是人工平滑控制的;这里的α是由条件嵌入回归出来的、每层每通道独立,且初始为零只是初始化而非约束,训练中完全自由。

来源考据【解读者补充,论文有引用】:零初始化残差分支的经验来自ResNet(如Goyal et al.的zero-init gamma)和GPT-2的zero-init输出投影——DiT把它搬进了扩散条件化的语境。adaLN-Zero的调节网络输出神经元数是隐藏维度的6倍(γ,β,α各两组,对应attention与MLP两个子层)。附录还给出一个彩蛋:guidance只施加于latent前三个通道时scale 1.5等效于全四通道的1.375,后者FID可达2.20

一句话记住本节:条件注入的最优解不是更贵的attention,而是"自适应归一化+恒等初始化"——便宜且快。

闭卷自检:四种设计的FID排序?adaLN-Zero比adaLN多回归了什么参数?为什么初始为恒等函数有帮助?

C. 扩展分析:Gflops才是真预言家(原文§4.2,Table 4)

来源:论文§4.2。学习目标:学完你能(a)从Table 4中挑出"Gflops相近、FID相近"的证据对;(b)解释为什么参数量失效而Gflops有效;(c)说出"大模型计算效率更高"的具体表现。

数字对账:12个模型全景(400K步,无cfg,FID-50K):

模型GflopsParams(M)FID备注
DiT-S/80.3633153.60S系列参数恒为33M
FID却从153→68
DiT-S/41.4133100.41
DiT-S/26.063368.40
DiT-XL/2118.6467519.47

三组"算力等价对":DiT-S/2(6.06G)与DiT-B/4(5.56G)FID分别为68.40/68.38——几乎逐位一致;DiT-L/4(19.70G)与DiT-XL/4(29.05G)、DiT-B/2(23.01G)落在43~46区间。这是"Gflops决定论"最漂亮的一组证据。

手算验证:S系列参数不变(33M),仅靠p从8缩到2就把FID砍掉85分——若按参数量外推,这三次改进都"不该发生";按Gflops外推(0.36→6.06,约17倍),落在log-log趋势线上。反过来,DiT-L/8参数458M是DiT-B/4的3.5倍,Gflops却更高(5.01 vs 5.56?不对——L/8为5.01G低于B/4的5.56G),FID 118.87反而远差于B/4的68.38:参数多3.5倍、算力略低、质量崩坏——参数量假说在此直接证伪。

第二发现:更大的DiT计算效率更高——估算训练算力=Gflops×batch×steps×3(前向+反向近似),XL/4在大约10¹⁰总Gflops后被XL/2反超;即小模型先出发、大模型后劲足且曲线更低。

常见误读①:"Gflops-FID负相关=无限加大算力就好"。注意所有点都在固定400K步测得,且作者强调单任务曲线噪声大、未观察到FID饱和意味着尚未到上限,不代表不存在。
常见误读②:"patch越小越好"。p=2确实最好,但它纯粹靠烧推理算力换来——XL/2的118.6G已是XL/4的4倍,部署成本必须计入(见第四章)。

一句话记住本节:想预测扩散模型的质量,别问它有多少参数,问它每次前向花多少算力。

闭卷自检:我能复述S/2 vs B/4这对"孪生实验"吗?能解释训练总算力的估算式里为什么要乘3吗?

D. SOTA结果与采样算力问题(原文§4.3–4.4)

来源:论文§4.3表2、表3及§4.4。学习目标:学完你能报出256²/512²两个基准的前后FID变化,并解释"模型算力vs采样算力"实验的设计逻辑。

ImageNet 256²(DiT-XL/2训7M步):无cfg FID 9.62(Recall高达0.67全场最高);cfg=1.25→3.22;cfg=1.50→2.27(sFID 4.60,IS 278.24),超过StyleGAN-XL的2.30与此前扩散最佳LDM-4-G的3.60。仅训2.35M步(与ADM预算相当)时已达2.55。算力坐标轴:XL/2为118.6 Gflops,与潜空间U-Net的LDM-4(103.6G)同档,远低于像素空间ADM的1120G、ADM-U的742G。

ImageNet 512²(训3M步,超参不变):64×64×4 latent、p=2→1024 token、524.6 Gflops;cfg=1.50达FID 3.04,超ADM-G/U组合的3.85(后者2813 Gflops)。

模型算力 vs 采样算力:测试小模型能否靠加采样步数[16…1000]翻身:DiT-L/2跑满1000步需80.7 Tflops/图、FID-10K 25.9;DiT-XL/2只用128步、15.2 Tflops/图(少5倍)反而更好(23.7)。结论:采样端补偿不了模型端的差距——这与语言模型领域"小模型多采样轮次追不上大模型"的直觉互证。

VAE解码器消融(表5):解码器可在训完扩散模型后互换重选:original 2.46 → ft-MSE 2.30 → ft-EMA 2.27,无需重训扩散主干——两阶段解耦的红利延伸到了评测环节。

一句话记住本节:2.27这个数字的意义不在破纪录本身,而在它是"纯transformer+别人家的VAE+别人的调度器"凑出来的——架构假设被干净地隔离了。

闭卷自检:256²与512²各自的最佳FID与对手数字?1024 token是怎么算出来的?

四、批判性阅读

4.1 benchmark到底测什么

类条件ImageNet生成测的是"给定标签合成多样、逼真图像",FID-50K依赖Inception特征的统计距离——它对多样性敏感(这正是cfg压Recall换Precision时FID仍能改善的原因之一)、对语义正确性不敏感。论文同时报告sFID/IS/Precision/Recall并使用ADM的TensorFlow评估套件保证可比性,这是加分项;但所有结论限于单一数据集、单一任务类型,没有文本条件的实证(只在结论里作为未来方向提及)。

4.2 比较是否公平

表2中各对手的训练预算差异巨大:DiT-XL/2训了7M步,而对比行ADM系预算明显更低(论文自己注明2.35M步时已有2.55,算是主动披露);StyleGAN-XL是判别式路线,FID接近但在Recall(0.53 vs DiT的0.57@cfg1.5口径不同)与采样速度上另有优势,论文未给任何方法的采样耗时对比——延迟这条第二坐标轴整体缺席。"超越所有先前扩散模型"成立,但"超越所有生成模型"的说法要打折:GAN的单步采样优势未被讨论。

4.3 成本与效率第二坐标轴

论文的效率论证集中在前向Gflops(训练/采样单步成本),这是它选定的坐标系。但p=2策略的本质是把算力从参数挪到激活:显存占用、通信、以及实际墙钟时间未必与Gflops同比例。XL/2的5.7迭代/秒(TPU v3-256)是全文唯一的吞吐数字。

4.4 论文没有告诉你什么

五、综合考核

5.1 重建因果链

  1. L4仅从"z=32×32×4"、"p=2"、"675M参数"、"118.6 Gflops"四个数字出发,推演这篇论文必须解决的两个核心设计问题,并指出每个数字逼出了什么。
    参考答案与评分标准(1)"32×32×4+p=2"逼出条件注入设计问题:256个token的全局序列里,t/c两个标量条件如何高效接入——四种block消融由此而来,答案是adaLN-Zero。(2)"675M参数vs 118.6 Gflops"逼出复杂度度量问题:p可变导致参数与算力脱钩,必须选择能预测质量的代理指标——扩展分析(Gflops标度律)由此而来。另外"32×32×4"也隐含了对LDM两阶段方案的依赖决策(借用SD的VAE)。评分:两个问题各3分,须写明数字→问题的因果方向;提到"参数与算力脱钩机制(p减半token×4)"者满分。
  2. L3构造反例或边界:举出一个场景,其中"Gflops越高FID越低"的趋势不再成立。(提示:想想固定算力预算下的分配、以及数据上限)
    参考答案场景一:固定训练总算力时,超大模型欠训练——XL/2若只训极少步数,其FID会劣于充分训练的小模型(论文图示的"XL/4早期领先XL/2"就是雏形,极端化即可)。场景二:数据信息上限——若数据集简单/极小,模型容量超过数据的可学习熵后继续加大Gflops只会过拟合或原地踏步。场景三:p过小导致token过多时attention的平方成本挤占有效训练步数。评分要点:须说明趋势成立的隐含前提(充分训练、数据充足);任一具体场景说清机制即可满分。

5.2 数字总对账

  1. L1256²图像、f=8 VAE、p=2:token总数是多少?每个token解码输出多少个数(C=4)?
    标准答案T=(32/2)²=256个token;每个token输出p×p×2C=2×2×8=32个数(含噪声与协方差两组预测)。解析:256×32=8192=32×32×2×4×2的一半是纯latent元素4096的两倍,因双头输出,自洽。
  2. L2用"Gflops相近⇒FID相近"检验:DiT-L/8(5.01G)与DiT-B/4(5.56G)的FID分别是多少?"相近⇒相近"在这里成立吗?这说明什么?
    标准答案118.87 vs 68.38——不成立!同为约5.5G,FID差50分。这说明Gflops并非充分统计量:同等算力下"宽而浅+大patch"(L/8)劣于"窄而合适+小patch"(B/4)。论文的负相关说的是总体趋势,个别配置偏离正是其价值所在。解析题意:此题考察对"强负相关≠逐点函数关系"的理解。
  3. L2采样算力对比:DiT-L/2千步80.7 Tflops/图对XL/2百步15.2 Tflops/图。前者每图的算力是后者的几倍?FID谁更好?
    标准答案80.7/15.2≈5.3倍;XL/2的23.7优于L/2的25.9。多花的5倍采样算力不仅没补齐反而买了个更差的结果。

5.3 设计决策答辩

  1. L4答辩四连:(a)为什么沿用ADM的协方差参数化而不学完整协方差?(b)为什么用现成SD的VAE而不自己训?(c)为什么所有规模共用一套超参而不分别调优?(d)为什么XL/2最终选p=2而不是更小的p?
    参考答案与评分标准(a)遵循Nichol & Dhariwal:ℒ_simple训ε_θ、完整ℒ训Σ_θ的对角参数化已在ADM验证,改动它会污染"只换backbone"的对照设计。(b)控制变量:自训VAE会引入第二变量,且SD VAE已被大规模验证;代价是与LDM/SD生态共享潜在偏置。(c)未调优的超参在所有12个配置上稳定工作且无loss尖峰,本身就是"transformer扩散训练很皮实"的证据,也让scaling对比更干净;风险是最优点可能没找到,XL档或许被低估。(d)p=2已把token推到256、XL/2达118.6G;再小则计算平方爆炸、工程上不现实,且论文未观察到FID饱和但预算有限——p=2是该代硬件下的务实终点。评分:每问2.5分,要求"理由+代价/反事实"齐全。
  2. L3有人说:"DiT证明了transformer全面优于U-Net。"找出这句话的两个漏洞。
    参考答案漏洞一:实验仅在类条件ImageNet、latent空间、特定分辨率的设定下进行,未覆盖像素空间扩散、文生图、视频等场景,外推无据(论文自己也只敢列为未来方向)。漏洞二:XL/2的胜利依赖p=2的高Gflops配置,同算力下并未与精心调优的U-Net逐一配对比较(对比对象ADM/LDM用的是各自原始配置)。此外"证明"一词过强:单个数据集上的FID优势是证据而非证明。评分要点:每个漏洞须落到具体实验设定;泛泛说"可能不公平"不给分。

5.4 证据审计(回看第一章预评)