Back to blog

MOD-DiT: Mixture of Distributions for Dynamic Sparse Attention in Video Diffusion Transformers

训练自由、无采样的动态稀疏注意力框架,通过线性近似模型预测三种核心注意力模式(块对角、平行对角、垂直)的强度演化,在 HunyuanVideo 上实现 2.05× 加速,在 Wan2.1 上实现 1.75× 加速。

MOD-DiT: 视频扩散 Transformer 的动态稀疏注意力

一、论文概述

项目内容
标题Mixture of Distributions Matters: Dynamic Sparse Attention for Efficient Video Diffusion Transformers
作者Yuxi Liu, Yipeng Hu (北京大学), Zekun Zhang (UESTC), Kunze Jiang (USTC), Kun Yuan (北京大学)
机构北京大学、电子科技大学、中国科学技术大学
论文arXiv:2601.11641
版本v3 (2026-07-01)
学科cs.CV, cs.LG
许可arXiv 永久非独占许可

二、核心思想

问题定义

视频扩散 Transformer(vDiT)通过 3D 全注意力机制联合建模时空动态,推动了 CogVideoX、HunyuanVideo、Wan2.1 等 SOTA 模型的发展。然而,自注意力的二次计算复杂度仍然是长序列视频生成实际部署的关键瓶颈。

现有稀疏注意力方法存在两大缺陷:

  1. 静态模式:依赖简化的固定稀疏模式(如固定垂直网格或全局指数衰减),无法捕捉 vDiT 中注意力分布的动态混合特性
  2. 采样开销:动态稀疏方法依赖计算昂贵的采样操作,反而抵消了效率收益,且忽略了去噪步对稀疏注意力的影响

解决方案概述

MOD-DiT(Mixture-Of-Distribution DiT)是一种训练自由、无采样的动态稀疏注意力框架,通过两阶段过程精确建模演化的注意力模式:

  1. 线性近似模型:利用早期去噪步的先验信息,通过分布式混合方法建模高效的线性近似模型,预测特定去噪区间的 mask 模式
  2. 在线块掩码策略:动态应用预测的 mask,同时保留历史稀疏信息,无需重复采样操作

关键指标:HunyuanVideo 2.05× 加速,Wan2.1 1.75× 加速,同时在 VBench 上保持顶级质量。

三、技术架构

整体框架

MOD-DiT 整体框架

图 1. MOD-DiT 在 HunyuanVideo 上的可视化对比:始终实现 2.05× 加速,且保持与原始视频几乎相同的生成质量。

核心公式

稀疏注意力定义:

SA(Q,K,V)=softmax(A+M)V(1)\text{SA}(Q, K, V) = \text{softmax}(A + M)V \tag{1}

其中 A=QKT/dA = QK^T/d 为缩放注意力分数,M∈RN×NM \in \mathbb{R}^{N \times N} 为二元 mask(Mp,q=−∞M_{p,q} = -\infty 表示丢弃交互)。

注意力稀疏图构建:

Sh,(i,j)=1B2∑x=0B−1∑y=0B−1I(Ah,(iB+x, jB+y)<η)(2)S_{h,(i,j)} = \frac{1}{B^2} \sum_{x=0}^{B-1} \sum_{y=0}^{B-1} \mathbb{I}(A_{h,(iB+x,\, jB+y)} < \eta) \tag{2}

其中 η\eta 为稀疏阈值,I(⋅)\mathbb{I}(\cdot) 为指示函数。Sh,(i,j)S_{h,(i,j)} 越小表示该块包含的信息越丰富。

广义线性近似模型:

S(t)=∑k=12n−1ck(t)Ck+∑k=1ndk(t)Dk+∑k∈Aek(t)Ek+R(t)(3)S(t) = \sum_{k=1}^{2n-1} c_k(t) C_k + \sum_{k=1}^{n} d_k(t) D_k + \sum_{k \in \mathcal{A}} e_k(t) E_k + R(t) \tag{3}

其中 n=N/Bn = N/B 为每维块数。三项基矩阵分别对应:

  • 平行对角模式 {Ck}\{C_k\}:建模帧间空间相关性,偏移量 δk=k−(n−1)\delta_k = k - (n-1)
  • 垂直模式 {Dk}\{D_k\}:通过列状结构捕获跨 token 全局依赖
  • 块对角模式 {Ek}\{E_k\}:维持帧内时间相干性,A\mathcal{A} 为有效块索引集

系数求解(最小二乘):

X(t)=arg⁡min⁡X∥vec(S(t))−MX∥22(4)X(t) = \arg\min_X \| \text{vec}(S(t)) - MX \|_2^2 \tag{4}

其中 M=[vec(C1),…,vec(E∣A∣)]M = [\text{vec}(C_1), \ldots, \text{vec}(E_{|\mathcal{A}|})] 为设计矩阵。归一化近似误差 NAE(N)=∥R(t)∥2/∥S(t)∥2\text{NAE}(N) = \|R(t)\|_2 / \|S(t)\|_2 随序列长度增大而减小。

代理密集注意力重建:

A^(tp(i+1))(p,q)={Amasked(tp(i+1))(p,q),M(tp(i+1))(p,q)=0A^(tp(i))(p,q),M(tp(i+1))(p,q)=−∞(5)\hat{A}(t_p(i+1))_{(p,q)} = \begin{cases} A_{\text{masked}}(t_p(i+1))_{(p,q)}, & M(t_p(i+1))_{(p,q)} = 0 \\ \hat{A}(t_p(i))_{(p,q)}, & M(t_p(i+1))_{(p,q)} = -\infty \end{cases} \tag{5}

线性预测(中段-末段去噪步):

c^k(t)=c^k(tp(i+1))+c^k(tp(i+1))−c^k(tp(i))tp(i+1)−tp(i)⋅(t−tp(i+1))\hat{c}_k(t) = \hat{c}_k(t_p(i+1)) + \frac{\hat{c}_k(t_p(i+1)) - \hat{c}_k(t_p(i))}{t_p(i+1) - t_p(i)} \cdot (t - t_p(i+1)) d^k(t)=d^k(tp(i+1))+d^k(tp(i+1))−d^k(tp(i))tp(i+1)−tp(i)⋅(t−tp(i+1))\hat{d}_k(t) = \hat{d}_k(t_p(i+1)) + \frac{\hat{d}_k(t_p(i+1)) - \hat{d}_k(t_p(i))}{t_p(i+1) - t_p(i)} \cdot (t - t_p(i+1))

动态 Mask 生成:

Mh(t)(p,q)={0,(p,q)∈Kt∪KE−∞,otherwise(6)M_h(t)_{(p,q)} = \begin{cases} 0, & (p,q) \in \mathcal{K}_t \cup \mathcal{K}_E \\ -\infty, & \text{otherwise} \end{cases} \tag{6}

其中 Kt\mathcal{K}_t 为 Top-K 动态模式集,KE=supp(E)\mathcal{K}_E = \text{supp}(E) 为块对角支撑集(通过静态阈值 emin⁡>τee_{\min} > \tau_e 条件保留)。

稀疏掩码保真度分数(SMFS):

SMFS=∥A⊙M∥F∥A∥F(7)\text{SMFS} = \frac{\|A \odot M\|_F}{\|A\|_F} \tag{7}

值越低表示稀疏 mask 以更少的计算保留了更多关键注意力信息。

注意力模式可视化

四种注意力模式

图 2. CogVideoX-v1.5 中的四种注意力模式可视化:(a) 块对角、(b) 平行对角、(c) 垂直、(d) 混合。

稀疏图演化

稀疏图演化

图 3. CogVideoX-v1.5 中层 0、头 9 的稀疏图在去噪步 12→21→31→41 的演化过程,展示了注意力稀疏模式在不同去噪步间的显著差异。

线性近似误差验证

线性近似误差

图 4. 线性近似模型(式3)在不同序列长度上的归一化近似误差(NAE),每条曲线包含均值和 95% 置信区间。误差随序列长度增加而降低。

模式强度演化

模式强度演化

图 5. 垂直和平行对角模式强度随去噪步的演化,显示中段-末段去噪阶段收敛到分段线性趋势。

推理时间对比

推理时间对比

图 6. Full Attention 与 MOD-DiT 在不同序列长度下的推理时间对比,MOD-DiT 加速效果随序列长度增加而更显著。

超参数消融

Top-K 消融:

Top-K 消融

图 7. Hunyuan 模型上 Top-K 超参数的消融实验:Subject Consistency (S.C.) 和 Imaging Quality (I.Q.) 随 K 增大先升后稳。

稀疏阈值消融:

阈值消融

图 8. 稀疏计算阈值 η\eta(式2)的消融:阈值对性能影响微弱,仅在极端大值时略有退化;默认选择 η=10−4\eta = 10^{-4}。

重建误差分析:

NRE 结果

图 9. CogVideoX-v1.5 上 6 个随机采样注意力头的归一化重建误差(NRE)在各去噪步的分布,确认 MOD-DiT 重建引入的误差极小。

Warm-up 步骤消融:

Warm-up 消融

图 10. CogVideoX-v1.5 上 warm-up 步数 mm 的消融:m=12m=12 在避免冗余全注意力成本的同时实现近最优的 S.C. 和 I.Q.

SMFS 对比:

SMFS 对比

图 11. HunyuanVideo 上不同方法的 SMFS 对比:MOD-DiT 在所有序列长度下均显著低于 SVG 和 Radial。

NRE 直方图:

NRE 直方图

图 14. 垂直(a)和平行对角(b)模式的 NRE 直方图,基于 300 个数据点(5 prompts × 10 layers × 6 heads),确认分段线性是普遍属性。

Mask 视觉对比:

Mask 对比

图 13. 各稀疏注意力方法在 CogVideo 推理中为特定注意力头生成的 mask 视觉对比:MOD-DiT 的 mask 更忠实地保留了真实注意力图的结构特征。

FastWan 加速对比:

FastWan 加速

图 18. HunyuanVideo 上不同稀疏注意力方法的生成可视化对比:MOD-DiT 实现 2.05× 加速且质量与原始视频几乎一致。

Wan2.1 加速对比:

Wan2.1 加速1

Wan2.1 加速2

图 19-20. Wan2.1 上不同稀疏方法的生成对比:MOD-DiT 实现 1.75× 加速,保持全注意力水平的主题一致性和场景真实感。

硬件加速策略

  1. 优化最小二乘核:基于 CUDA 的优化核,相比 PyTorch 原生 torch.lstsq 加速 100×
  2. 混合注意力执行:Warm-up 阶段用 FlashAttention-2,稀疏阶段切换到 SageAttention
  3. 块级注意力计算:将 Q/K 张量划分为 128×128 非重叠块,利用 GPU warp 级并行

MOD-DiT 引入的计算开销仅占全注意力延迟的 1-2%。

四、核心创新

创新点说明理论/实验依据
三模式混合假设揭示 vDiT 注意力图呈现块对角、平行对角、垂直三种结构的动态混合图2-3 可视化;式(3) 线性近似;NAE < 0.14
训练自由、无采样动态稀疏完全在线推理,无需训练成本,消除采样开销式(5)-(6) 在线预测-校正流水线
三级动态路由输入级、头级、去噪步级同时优化 mask 模式和稀疏比Algorithm 1 三阶段流程
硬件优化核CUDA 加速最小二乘求解,减少 100× 计算时间Remark 1;Appendix B
分段线性预测中段-末段去噪步模式强度呈稳定分段线性,可线性外推图5;图14 300数据点验证

五、实验结果

主基准测试(Table 1)

HunyuanVideo (13B, 117帧, 768×1280, A100):

方法稀疏度PSNR↑SSIM↑LPIPS↓S.C.↑B.C.↑M.S.↑T.F.↑I.Q.↑A.Q.↑D.D.↑延迟加速
Full0%---0.95820.94780.97660.97230.66930.63810.92156978s1.00×
MInference67.1%18.210.6380.4900.93000.91800.95800.95000.64500.56500.90805286s1.32×
Radial75.55%26.720.8850.1250.93120.92900.97150.93220.64670.57150.84233731s1.87×
SVG71.22%26.440.8610.1700.83010.87110.95440.90180.59280.53300.90113834s1.82×
Sparge68.00%25.430.8420.1950.93390.91560.95280.95100.64320.56030.91204105s1.70×
MOD-DiT83.23%27.730.8790.1190.93980.93200.96870.95270.65870.58990.91043405s2.05×

Wan2.1 (14B, 69帧, 768×1280, A100):

方法稀疏度PSNR↑SSIM↑LPIPS↓S.C.↑B.C.↑M.S.↑T.F.↑I.Q.↑A.Q.↑D.D.↑延迟加速
Full0%---0.96230.96550.98440.97890.67220.60120.93783375s1.00×
MInference60.5%15.810.6750.3430.91000.92000.97500.96500.65500.57000.91502557s1.32×
Radial71.33%21.570.8180.1670.91520.93390.98600.97500.66200.57830.85002021s1.67×
SVG69.08%20.930.7950.2220.79860.88590.95560.93470.61330.52920.91672109s1.60×
Sparge50.10%18.670.7350.1980.90330.92750.97980.96320.66450.56100.90432296s1.47×
MOD-DiT81.37%22.750.8210.1520.94270.94880.98230.97520.66740.58270.92891929s1.75×

CogVideoX-v1.5 结果(Table 5)

方法稀疏度PSNR↑SSIM↑LPIPS↓S.C.↑I.Q.↑延迟加速
Full0%---0.92300.6255987s1×
MInference64.9%15.010.6010.3340.86790.5580696s1.42×
Radial70.7%22.890.8660.1720.92140.6167611s1.62×
SVG75.0%21.150.8180.1830.91580.5948596s1.65×
Sparge67.3%20.340.7730.2550.90430.5966661s1.49×
MOD-DiT80.1%25.770.8680.1330.92660.6239542s1.82×

公平对比(等稀疏度)

方法S.C.↑B.C.↑M.S.↑T.F.↑I.Q.↑A.Q.↑D.D.↑延迟
Radial (匹配稀疏度)0.91520.93390.98600.97500.66200.57830.85001224s
MOD-DiT0.95180.95940.98670.97760.66300.60700.97891265s

SVG-2 对比

方法S.C.↑B.C.↑M.S.↑T.F.↑I.Q.↑A.Q.↑D.D.↑延迟
SVG-20.93250.93690.96360.95180.64310.59010.90761056s
MOD-DiT0.93980.93200.96870.95270.65870.58990.91041044s

蒸馏模型适配(FastWan 8-step)

方法S.C.↑B.C.↑M.S.↑T.F.↑I.Q.↑A.Q.↑D.D.↑延迟
FastWan0.93170.94890.97660.96580.65310.59100.9201327s
FastWan + MOD-DiT0.92870.94700.96920.95430.65020.59040.9176210s

MOD-DiT 在蒸馏模型上额外获得 1.56× 加速,VBench 质量指标几乎无损。

消融实验摘要

实验发现
重建间隔 Δt所有测试值(1,2,3,5,10)表现几乎相同;“None” 方案显著下降
块大小64→128 降低延迟但质量几乎不变,128 为最优选择
Warm-up 步数 mm=12 达到近优质量且避免冗余全注意力成本
Top-K性能随 K 先增后稳,存在最优平衡点
稀疏阈值 η对性能影响微弱,默认 η=10⁻⁴ 稳定有效

六、关键设计洞察

6.1 注意力流形的三基分解

在任何帧优先的 token 布局下,高能量注意力流形可严格分解为三个正交互点子空间:

  • 块对角 (Pspatial\mathcal{P}_{\text{spatial}}):帧内空间相干性
  • 平行对角 (Ptemporal\mathcal{P}_{\text{temporal}}):帧间时间连续性
  • 垂直全局 (Pglobal\mathcal{P}_{\text{global}}):跨帧语义锚定

这三个子空间的并集构成视频生成主高交互流形的理论完备基。

6.2 预测-校正流水线

通过短期区间(Δt\Delta t)线性外推 + 周期性全局重构(式5)的双重机制:

  • 早期高度非线性动态由全注意力 warm-up 安全处理
  • 中段-末段稳定 regime 下激活线性预测
  • 周期性重构防止长期外推误差累积

6.3 块对角模式的静态阈值化

由于块对角模式编码跨去噪步的结构不变性帧内空间交互,其存在可通过静态阈值高效验证,避免了持续线性预测的需求。

七、相关工作

视频稀疏注意力:Sparse-vDiT(静态垂直+块对角)、SVG/SVG-2(时空分类)、Radial(指数衰减)、MInference(LLM迁移)、SpargeAttn(交叉模态)、LiteAttention(时序连贯跳过)、DFSAttn(动态细粒度)

高效扩散:TeaCache/FasterCache/AdaCache(时序缓存)、蒸馏方法(减少步数)、高比率 VAE(潜空间压缩)

稀疏注意力:Longformer(滑动窗口)、Swin Transformer(层次 shifted windows)、FlexPrefill(上下文感知)

八、总结

核心贡献

  1. 三模式混合假设:首次系统揭示 vDiT 注意力图呈现块对角、平行对角、垂直三种结构的动态混合
  2. 训练自由动态稀疏:插件式算法,完全在线推理,无训练成本和采样开销
  3. 三级动态路由:输入级、头级、去噪步级同时优化 mask 模式和稀疏比
  4. 硬件优化核:CUDA 最小二乘核加速 100×,计算开销仅 1-2%
  5. 广泛实验验证:HunyuanVideo 2.05×、Wan2.1 1.75×、CogVideoX 1.82× 加速,VBench 顶级质量
  6. 蒸馏模型兼容:FastWan 8-step 额外 1.56× 加速

局限性

  • 全注意力 warm-up 阶段在短序列场景可能引入额外计算开销
  • 块对角模式通过静态阈值而非动态预测,可能遗漏少数变化场景

参考资源