跳转至

Rethinking Cross-Layer Information Routing in Diffusion Transformers

会议: NeurIPS2026
arXiv: 2605.20708
领域: 图像生成
关键词: 扩散 Transformer、跨层信息路由、时间步条件化、注意力残差、分块聚合

一句话总结

DAR 将扩散 Transformer 的逐层残差累加改为时间步感知的历史子层输出注意力,在 ImageNet 256×256 上以 600K 步取得无 CFG 的 ODE FID 7.56,相比训练 1.75M 步的 SiT 基线降低 2.11,并能与 REPA 的表征对齐损失结合。

研究背景与动机

扩散 Transformer(DiT)的 token 化、注意力、条件注入、训练目标和潜空间编码器已经反复改进,但跨层传递信息的方式仍常沿用标准 Transformer:每个注意力或 MLP 子层都把自己的输出加进同一条残差流。展开后,这相当于给输入嵌入和所有历史子层输出赋予相同的单位权重。浅层信息一旦进入残差流,深层只能继续往上叠加,不能显式重新决定哪些来源值得保留、哪些应该压低。

论文用训练至 600K 步的 SiT-XL/2,在 4096 个 ImageNet 样本上诊断这一问题。以时间步 \(t=1.0\) 的切片为例,块输出 RMS 从约 15.5 增长到约 1576;前五块之后,反向梯度明显衰减;深层相邻块的 token 余弦相似度持续高于 0.9。作者把它们联系到 PreNorm dilution:残差流越来越大,而新分支的输入仍被归一化,深层更新因此更难对总状态产生有效影响。不过,这些是诊断关联,不是证明任何高相似度都等价于无用计算。

扩散还有普通语言模型没有的控制维度:去噪时间步。不同噪声水平需要不同的浅深层信息组合,固定权重无法直接响应这一需求;手工 U-Net 长跳连虽然重新引入浅层特征,却预先固定了层间配对。论文给基线历史来源加入只用于测量、初值为 1 的标量门,保持前向不变,再用损失对门的梯度探测来源的重要性,发现来源偏好随时间步变化。核心 idea:不再把历史信息永久压进一条等权累加的残差流,而让每个子层通过时间步感知的深度注意力重新选择历史来源,并用分块保存控制开销。

方法详解

整体框架

DAR(Diffusion-Adaptive Routing)保留 SiT 的注意力、MLP 和条件化计算,替换的是这些子层之间的残差聚合。给定当前带噪潜变量、时间步和类别条件,网络先获得输入嵌入;每个子层从可用历史来源中做 softmax 加权聚合,再执行原有子层变换,把新输出加入后续可检索的来源集合。这里的“深度注意力”是在不同子层输出之间选择,不是替换空间 token 之间的自注意力。

路由由三组设计共同决定:历史来源如何加权、查询如何获得时间步信息,以及分块后保留哪些来源。普通 DAR 还在预测头前使用专门的最终聚合器,让最后一个 chunk 的原始子层输出直接参与最终预测;与 REPA 联用时,最终聚合的参数共享方式另有调整。推理仍执行正常去噪采样,不需要 DINOv2 教师;教师只在启用 REPA 的训练中提供监督。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["带噪潜变量<br/>时间步与条件"] --> B["深度注意力聚合"]
    B --> C["时间步感知查询"]
    C --> D["分块来源与末端读取"]
    D --> E["原有子层与预测头<br/>输出速度场"]
    E -->|推理:采样器反复调用| F["生成图像"]
    G["训练:速度目标<br/>可选 DINOv2 教师"] -.->|MSE;可选 REPA 对齐| E
    D -.->|历史来源供后续子层检索| B

图中前三个节点表示路由器内的设计关系,而不是三个额外串行网络。查询参与深度注意力的打分,分块决定其候选来源;虚线监督只在训练时出现。

关键设计

1. 深度注意力聚合:让每个子层重新选择历史信息,而非继续等权累加

把注意力和 MLP 分别视为一个子层,记输入嵌入为 \(v_0\),第 \(i\) 个子层的原始输出为 \(v_i\)。标准残差下,进入第 \(l\) 个子层的状态就是此前所有来源之和;DAR 则在深度维度上归一化权重,再混合来源。对每个来源做 RMSNorm 得到 key,用查询与 key 的点积判断相关性,真正被加权的 value 仍是来源输出本身:

\[ h_l=\sum_{i=0}^{l-1}\alpha_{i\to l}(t)v_i,\qquad \alpha_{i\to l}(t)=\frac{\exp\!\left(q_l(t)^\top k_i/\sqrt{d}\right)}{\sum_{j=0}^{l-1}\exp\!\left(q_l(t)^\top k_j/\sqrt{d}\right)},\qquad k_i=\mathrm{RMSNorm}(v_i). \]

softmax 约束权重和为 1,使来源数量增长不再自动意味着单位系数之和增长;同时,路由可以压低无关来源、突出少量有效来源。所谓“非增量聚合”指每次从可访问来源重新构造输入,而不是对上一层累计状态只做一次加法;并不意味着网络不再按层顺序执行,也不意味着彻底取消历史依赖。

这一改动保留同构 Transformer 堆栈,不需要指定浅层与深层一一配对。论文观察到 DAR 路由权重集中在少量来源上,并且随时间步变化,这与前面基线的反事实门梯度诊断呼应。但门梯度和 softmax 权重不是同一量,不能直接把二者的颜色或数值当作可比概率。

2. 时间步感知查询:区分查询参数静态与路由权重静态

DAR 比较三种查询。纯 static 是每个子层一个可学习向量;dynamic 对最近子层输出做线性投影;static 加显式时间步注入则复用现有时间嵌入,不新增查询投影矩阵。三者可写为:

\[ q_l(t)= \begin{cases} w_l,&\text{pure static},\\ W_q^{(l)}v_{l-1},&\text{dynamic},\\ w_l+e(t),&\text{static with explicit timestep injection}. \end{cases} \]

dynamic 查询从带噪输入及 adaLN-Zero 条件通路影响后的子层输出中继承时间步和内容信息;显式注入则直接把时间嵌入加到查询上。时间嵌入器末层采用零初始化,使显式注入版本在训练起点恢复纯 static。这里的 static 指查询的基础参数化,不应理解成固定路由:即使纯 static 查询不直接接收时间步,key 和 value 仍随带噪输入与网络条件变化。

表 2 显示,显式时间步注入已经能接近 dynamic,甚至在 400K 步取得更低 FID,说明主要收益不能简单归因于内容相关查询或更大的参数量。dynamic 的线性探针分析还显示,聚合器输入的时间步解码 \(R^2\) 在前五块超过 0.95,深层接近 1.0。不过,原文探针段落用的是聚合状态 \(h_l\),查询公式用的是最近原始输出 \(v_{l-1}\);笔记保留这个证据对象差异,不把它们默认为完全相同的张量。

3. 分块来源与末端读取:保留当前细节,只压缩已经经过的块

保存全部子层输出会让来源缓存随深度线性增加。DAR 把 \(L\) 个子层划为大小为 \(S\) 的 chunk;每个已完成 chunk 只保留其最后一个子层输出作为摘要,即 \(c_n=v_{nS}\),输入嵌入作为 \(c_0\)。摘要不是 chunk 内输出的求和,也不是对整个 chunk 再做一次单独池化;其对更早信息的依赖来自该末子层已有的路由输入。

当第 \(l\) 个子层位于第 \(n\) 个 chunk 时,可以访问输入嵌入、此前 chunk 的末端摘要,以及当前 chunk 内所有此前的原始子层输出:

\[ \mathcal{S}_l=\{c_0,c_1,\ldots,c_{n-1}\}\cup\{v_{(n-1)S+1},\ldots,v_{l-1}\},\qquad c_n=v_{nS}. \]

候选来源数量至多为 \(S+N\),其中 \(N=L/S\)。论文按单个 token 的特征维度记账,将来源缓存由 \(O(Ld)\) 降至 \(O((S+N)d)\);实际批量和 token 维度仍需计入显存。\(S=1\) 退化为所有历史子层都可访问,增大 \(S\) 则用更少摘要压缩更多旧输出,但可能丢失历史细节。

最终预测头使用的来源略有不同:输入嵌入、前 \(N-1\) 个 chunk 的摘要,以及最后一个 chunk 的全部原始输出,包括最终 \(v_L\)。相比只读取 chunk 摘要的 AttnRes,这给最后几层的任务相关细节保留了直接出口;附录报告 200K 步约 2 点 FID 收益,但没有给出对应完整数值表。与 REPA 联用时,不使用独立最终聚合器参数,而复用最后 chunk 的 MLP 聚合器查询及逐来源 RMSNorm 参数;这是架构实现差别,不是 REPA 损失的新定义。

一个完整示例

SiT-XL/2 有 28 个 Transformer 块,每块含注意力和 MLP 两个子层,所以 \(L=56\)。采用 \(S=4\) 时,共有 14 个 chunk;这里的 c4 是四个子层,而不是四个 Transformer 块。

进入第二个 chunk 的第一个子层时,来源是输入嵌入 \(c_0\) 和前一 chunk 的末子层输出 \(c_1=v_4\)。进入第七个子层时,还能访问当前 chunk 内的 \(v_5\)、\(v_6\);不再单独缓存旧 chunk 的 \(v_1\)、\(v_2\)、\(v_3\)。

预测头前,来源包含 \(c_0\) 到 \(c_{13}\),以及最后 chunk 的 \(v_{53}\) 到 \(v_{56}\),共 18 个来源。时间步改变时,显式注入或 dynamic 查询都会相应改变打分,网络可在同一组可用来源中重新选择;这里不虚构某个来源必然获得的具体权重。

损失函数 / 训练策略

DAR 本身改变跨层聚合,不额外规定新的生成损失。ImageNet 实验沿用 SiT 的速度预测 MSE 和数据处理,global batch size 为 1024,学习率为 \(1\times10^{-4}\),使用 bfloat16;SiT 和 REPA 基线均由作者在相同实验环境重跑。

启用 REPA 时,使用 DINOv2-B 作为预训练视觉编码器,在第八层施加系数为 0.5 的表征对齐损失。DAR 负责“历史状态如何组合”,REPA 负责“中间表征被什么监督塑形”;二者兼容不等于完全没有实现交互,上述末端参数复用正是联用配置的特别之处。

实现上,作者用 fused Triton kernel 合并 RMSNorm、查询与 key 点积、softmax 和加权和,并在反向阶段重计算部分中间量。它减少算子启动和 HBM 访问,并没有消除来源保存本身的显存成本;因此训练迭代减少、单步吞吐和峰值显存应分开评价。

实验关键数据

主实验

ImageNet-1K,256×256,50,000 张生成样本,默认 250 次函数评估;CFG 的 guidance scale 为 1.5。表中采用原文表 1 的 XL/2 系统结果,static c4 的显式时间步注入定义由附录表 5 进一步说明。

方法 训练步数 参数量 无 CFG ODE FID 无 CFG SDE FID 有 CFG ODE FID 有 CFG SDE FID
SiT 1.75M 675M 9.67 8.61 2.15 2.06
SiT-Plus 1M 752M 10.85 10.02 2.36 2.34
DAR Static c4 600K 675M 7.56 6.92 2.08 2.23
DAR Dynamic c4 500K 751M 8.07 7.39 2.05 2.17

这些是不同训练预算下的系统比较,不是统一步数下的纯架构消融。无 CFG 时 DAR 的 ODE/SDE 都改善;有 CFG 的 SDE 上,SiT 的 2.06 反而优于 DAR static 的 2.23 和 dynamic 的 2.17,不能写成所有采样配置都占优。

消融实验

下面按原文表 2 比较时间步查询,指标为不同训练步数的 FID。它最直接区分“参数量增加”和“时间步可见性”的作用。

查询配置 100K 200K 400K
Static,不显式注入时间步 22.36 15.47 11.51
Dynamic 13.95 9.29 8.10
Static,显式注入时间步 17.39 10.12 7.97

原文第 5.3 节对应段落把时间步消融引用为“表 3”,但实际数据在表 2;表 3 是下面的 REPA 联用结果。这里保留并指出交叉引用矛盾,不据此混用两张表。

REPA 联用配置 100K FID 200K FID 300K FID
SiT + REPA 9.89 6.89 6.29
DAR + REPA 7.09 5.92 5.68

DAR+REPA 在 100K 的 7.09 接近、但不优于 REPA 在 200K 的 6.89;论文的约 2 倍早期加速是相近质量阈值的描述,不是两个 FID 完全相等。

关键发现

  • 时间步显式注入在 400K 把纯 static 的 FID 从 11.51 降至 7.97,dynamic 为 8.10;支持时间步感知的重要性,但没有证明 dynamic 在所有阶段最优。
  • 分块消融在 300K、无 CFG 下报告 \(S=1/4/8\) 的 FID 为 10.41/8.39/11.14。中间 chunk size 最好;该表的 \(S=4\) 数值与附录表 5 的 300K static c4 8.62 不同,原文未在表 4 清楚说明查询配置,不能直接当作同一检查点。
  • 作者报告达到基线 FID 9.67 所需迭代减少 8.75 倍。融合实现下 SiT/DAR static c4 为 1.83/1.73 steps/s,对应估计墙钟加速 8.27 倍,而不是 8.75 倍;峰值显存为 54.56/69.97 GB。该质量阈值的精确交点未在离散表格中直接给出,不能把 static 200K 的 10.12 说成已经达到 9.67。
  • 附录直接移植 AttnRes 的最好已报告结果为 700K 的 FID 8.71;DAR static c4 在 600K 为 7.56,低 1.15。这个比较支持扩散专用调整,但不能逐项归因到某一个修改。
  • Qwen-Image 四步 DMD 在 GenEval2 prompts 上将教师相对 LPIPS 从 0.538 降到 0.512,RAPSD 高频占比偏差从 +0.340 降到 +0.215。后者比较频率阈值大于 0.2 的高频内容,原文没有给出完整计算公式;这不是 GenEval2 语义正确率。

亮点与洞察

  • 残差流不仅是优化稳定性的工具,也可以看成一个隐含的历史检索器。把固定求和改成可学习选择,给深层提供了重新组织信息的自由度,而不必改动空间注意力主体。
  • static 参数不等于 static 行为。显式时间嵌入让低参数开销的查询响应噪声阶段,同时 key 本身也依赖当前输入;讨论路由动态性时应区分这两层含义。
  • 分块的关键不只是少存几个张量,而是决定哪些信息允许被压缩。旧 chunk 保留末端输出、最后 chunk 保留原始输出,使成本控制与最终细节出口同时进入架构设计。
  • 可迁移的研究方向是把其他条件生成模型的控制信号加入深度查询,再检查与原有表征损失是否互补。这是由结果引出的假设,并非本文已验证的跨任务结论。

局限与展望

  • 作者承认,大规模多十亿参数 T2I/T2V 预训练仍需系统研究;Qwen-Image 后训练只提供初步证据,不能外推到视频或所有偏好优化流程。
  • chunk size 的理论分析依赖一个假设性的路由熵与压缩失真代价模型。其最优大小随深度平方根增长的预测不是实际 FID 的无条件定理,也缺少跨深度实验证明。
  • fused kernel 相对 naive DAR 的算子级加速不能当作相对 SiT 的端到端加速。真实训练中 DAR 仍有约 5.8% 单步开销,且上述 static 配置显存更高。
  • DMD 采用 LoRA rank 64、学生/假分支学习率 \(5\times10^{-6}\)/\(2\times10^{-6}\)、四步去噪、guidance 4.0、1024² 分辨率和每 GPU batch size 1;教师相对指标不能替代独立的感知质量、语义对齐或用户偏好评价。
  • 后续最有价值的实验是统一质量阈值与墙钟预算,报告不同深度和噪声阶段的路由行为,并把时间步注入、chunk 摘要、最终聚合与 REPA 参数共享分别消融。

相关工作与启发

  • vs Attention Residuals:共同使用深度 softmax 注意力;DAR 的区别在扩散时间步查询、末子层摘要、chunk size 选择和最终原始输出读取。它是扩散场景的机制适配与实证诊断,不应把深度注意力本身称为首次提出。
  • vs U-ViT / U-DiT:长跳连通过人工浅深层配对恢复多层特征;DAR 用可学习来源选择保持同构堆栈。系统参数、训练步数和采样协议不同,跨模型结果不能全部归因于连接拓扑。
  • vs REPA:REPA 用视觉教师监督中间表征,DAR 改变表征跨层流动。联用收益说明两条轴能互补,但原文的特殊最终聚合实现意味着兼容性仍需具体架构验证。
  • vs 层缓存加速方法:缓存方法主要复用相邻去噪步或层的特征来降低推理计算;DAR 在每次网络调用内重构深度信息来源,重点是训练收敛与生成质量,不是直接减少采样步数。

评分

  • 新颖性: 4/5 — 深度注意力继承 AttnRes,但时间步诊断及扩散适配形成明确增量。
  • 实验充分度: 4/5 — 有查询、分块、REPA、直接 AttnRes 和 DMD 证据,缺少大规模预训练与完整因素拆解。
  • 写作质量: 3/5 — 机制整体清晰,但表格交叉引用及探针张量表述存在歧义。
  • 价值: 4/5 — 为 DiT 提供一个与训练目标互补的架构轴,收益需结合吞吐、显存和采样条件判断。