Back to blog

VSA: 可训练稀疏注意力加速视频扩散模型

VSA 提出可训练稀疏注意力机制,加速视频扩散模型推理。

VSA: 可训练稀疏注意力加速视频扩散模型

一、论文概述

项目内容
标题VSA: Faster Video Diffusion with Trainable Sparse Attention
作者Peiyuan Zhang, Yongqi Chen, Haofeng Huang, Will Lin, Zhengzhong Liu, Ion Stoica, Eric P. Xing, Hao Zhang
机构UC San Diego, MBZUAI, UC Berkeley
论文https://arxiv.org/abs/2505.13389
代码https://github.com/hao-ai-lab/FastVideo
发布2025年5月19日 (v5: 2025年10月28日)
许可CC BY 4.0
领域视频生成、扩散模型、稀疏注意力、高效推理

二、核心思想

问题定义

视频扩散Transformer (DiT) 的扩展受限于二次方3D注意力的计算瓶颈:

  • 一段5秒720p视频展开后超过100K tokens
  • 现有SOTA视频DiT(HunyuanVideo、Wan-2.1、CogVideoX)在全分辨率长序列训练上大部分计算消耗在注意力
  • 训练后的DiT推理仍然非常缓慢

现有方法的局限

方法类型代表问题
后验稀疏STA, Sparge Attention, SVG仅在推理时应用稀疏,训练成本未减少;训练-测试不匹配
固定模式空间-时间、滑动窗口无法适应数据依赖的动态稀疏模式
启发式方法压缩KV、步幅窗口在充分训练后性能下降

解决方案概述

VSA提出可训练、硬件高效的稀疏注意力,在训练和推理时都替代全注意力:

核心观察:注意力矩阵中大部分权重集中在少量”关键token”上,其余接近零。

关键创新:

  1. 分层粗细注意力:粗阶段定位关键token,细阶段仅在关键区域内计算
  2. 端到端可训练:通过数据学习稀疏模式,而非启发式规则
  3. 硬件对齐:将空间-时间立方体映射到GPU kernel的tile,确保实际加速

三、技术架构

整体框架

输入视频 latent (T, H, W)
         │
         ▼
┌─────────────────────────────────────────────────────────────┐
│                    VSA 分层稀疏注意力                         │
│                                                              │
│  ┌─────────────────────────────────────────────────────┐    │
│  │ 粗阶段 (Coarse Stage)                               │    │
│  │ - 将 (Ct, Ch, Cw) = (4,4,4) 立方体 mean pooling     │    │
│  │ - 生成 Qc, Kc, Vc ∈ R^(L/B × d)                    │    │
│  │ - 计算立方体级注意力 Ac                              │    │
│  │ - Top-K 选择关键立方体 → 生成块稀疏掩码 M            │    │
│  └─────────────────────────────────────────────────────┘    │
│                          │                                    │
│                          ▼                                    │
│  ┌─────────────────────────────────────────────────────┐    │
│  │ 细阶段 (Fine Stage)                                 │    │
│  │ - 使用掩码 M 指导 token 级注意力                     │    │
│  │ - 仅在 Top-K 立方体内计算 Q·K^T 和 A·V              │    │
│  │ - 块稀疏布局对齐 GPU 内存访问模式                    │    │
│  └─────────────────────────────────────────────────────┘    │
│                          │                                    │
│                          ▼                                    │
│  ┌─────────────────────────────────────────────────────┐    │
│  │ 门控融合                                            │    │
│  │ O = Oc ⊙ Gc + Of ⊙ Gf                              │    │
│  │ Gc, Gf 为线性投影生成的门控向量                      │    │
│  └─────────────────────────────────────────────────────┘    │
└─────────────────────────────────────────────────────────────┘
         │
         ▼
    输出 O

立方体分区与索引映射

视频latent形状 (T, H, W) 被划分为多个立方体,每个形状 (Ct, Ch, Cw) = (4, 4, 4):

Tile大小 B = Ct × Ch × Cw = 64

1D索引映射公式:
n = (⌊t/Ct⌋·Nh·Nw + ⌊h/Ch⌋·Nw + ⌊w/Cw⌋)·B
  + (t mod Ct)·Ch·Cw + (h mod Ch)·Cw + (w mod Cw)

确保同一立方体内的token在1D序列中连续排列
→ 映射到GPU SM的同一tile

关键设计选择

设计维度选择原因
Tile大小 B64 (4×4×4)表达性与效率的最佳平衡
Top-K3287.5%稀疏率,跨序列长度表现一致
池化方式Mean Pooling优于Max Pooling和Conv
局部先验不需要粗阶段+细阶段已足够
粗阶段贡献需要提供全局上下文,通过门控融合

核心公式

注意力计算:

S = Q·K^T / √dk
A = Softmax(S + M)
O = A·V

门控融合:

O = Oc ⊙ Gc + Of ⊙ Gf

其中 Gc, Gf 是从输入隐藏状态线性投影得到的门控向量。

近似稀疏率:

sparsity ≈ K·B / L

粗阶段计算成本可忽略(<1% 总FLOPS)。

四、核心创新

创新点说明效果
可训练稀疏注意力端到端学习关键token位置,而非后验启发式训练FLOPS降低2.53×,无质量损失
分层粗细架构粗阶段定位+细阶段精算,门控融合兼顾全局上下文和局部精度
硬件对齐设计立方体→tile映射,块稀疏布局保留85% FlashAttention3 MFU
稀疏蒸馏首次将稀疏注意力与蒸馏结合50.9×加速,无质量损失
稀疏适配策略退火策略平滑过渡全注意力→VSA成功适配Wan-2.1预训练模型

五、实验结果

消融实验 (Table 1)

实验设置:

  • 模型:120M参数Wan架构
  • 训练FLOPS:4.5×10^20
  • 视频latent:(16, 32, 32)
  • 数据集:Vchitect-T2V-Dataverse

(a) VSA vs 其他注意力方法

方法Loss (Opt)Loss (Over)
Compress KV0.152810.14282
Spatial Temporal0.135740.13034
Spatial Full0.135550.12811
Strided Window0.132710.12716
Full Attention0.138770.12703
VSA0.131620.12687

发现:VSA在计算最优和过训练设置下都优于全注意力和其他稀疏方法。

(b) 注意力设计消融

设计Loss
L (局部窗口)0.13330
F (仅细阶段)0.13296
C & L0.13220
C & F0.13162
C & F & L0.13194

发现:粗阶段(C)+细阶段(F)的组合最优,局部先验(L)无显著帮助。

(d) Tile大小消融

Tile大小TFLOPSLoss
256×2564780.13375
128×1284440.13244
64×644080.13162
16×641810.13155

发现:更小tile性能更好但速度更慢,64×64是最佳平衡点。

扩展性实验 (Figure 2)

模型规模训练FLOPSVSA LossFull Loss注意力FLOPS节省
410M4×10^21≈0.126≈0.1268×
60M-1.4B最高4×10^21Pareto优-2.53× (总训练FLOPS)

关键发现:

  • VSA在410M模型上达到与全注意力几乎相同的loss曲线
  • 60M到1.4B的扩展实验确认VSA持续产生更好的Pareto前沿
  • VSA是首个在严格扩展性评估下优于全注意力的可训练稀疏注意力

稀疏适配与蒸馏 (Section 3.3)

Wan-1.3B VBench结果

模型质量分语义分总分
原始Wan83.71%77.98%82.56%
全注意力微调84.07%81.85%83.63%
VSA微调83.60%79.47%82.77%

推理加速:

模型全注意力VSA加速比
Wan-1.3B (480P)31s18s1.7×
Wan-14B (720P)1274s576s2.2×

稀疏蒸馏:

  • 学生模型使用VSA + DMD2蒸馏
  • 教师模型保持全注意力
  • 结果:50.9×加速,无质量损失

内核性能 (Section 3.4)

指标数值
VSA fine kernel MFU85% of FlashAttention3
长序列加速比7× (理论极限8×)
包含粗阶段后加速6×+
FlexAttention同配置加速仅2×
Wan-1.3B注意力加速6×
Hunyuan注意力加速2-3×

关键Token预测准确性 (Section 3.5)

  • 随机选择32个立方体:仅捕获**8%**注意力分数
  • VSA粗阶段选择:大部分层和时间步达到**60-90%**准确率
  • 准确率随时间步单调递增
  • 跨层呈锯齿模式,暗示未来可优化方向

六、相关工作

稀疏注意力在LLM中的发展

方法特点与VSA的区别
MoBA可训练动态稀疏VSA直接贡献粗阶段输出,使用更小块
NSA原生稀疏注意力VSA避免分组查询约束(视频双向注意力)
Mistral滑动窗口固定模式,非数据依赖
H2O, Minference推理时加速VSA同时优化训练和推理

稀疏注意力在视频DiT中的发展

方法特点局限
STA滑动tile注意力后验应用,训练成本未减少
Sparge Attentionprofile驱动稀疏训练-测试不匹配
SVG无训练稀疏化稀疏率受限,质量下降
DSV训练时稀疏多阶段设计复杂
VSA端到端可训练-

VSA的独特性

视频DiT对可训练稀疏注意力的需求比LLM更紧迫:

  1. 视频DiT序列更长(100K tokens仅5秒视频)
  2. 视频DiT将大部分计算预算用于长序列训练
  3. LLM的”短训练-长适配”范式在视频DiT中不适用

七、代码实现分析

项目信息

内核实现细节

粗阶段内核:

融合操作:softmax → Top-K选择 → 掩码转索引
- 序列长度降低64×(100K → 1.5K)
- 内存开销可忽略
- 计算成本 <0.2% 总注意力FLOPS
- 运行时间占比 ~14%

细阶段内核:

基于ThunderKittens的块稀疏注意力
- Tile大小:64
- 输入:块索引(非二进制掩码)
- 对齐FlashAttention的块稀疏布局
- MFU:85% of FA3

稀疏适配策略

1. 初始化 Gc = 0 (粗门控为零)
2. 移除 Gf (细门控等效为1)
3. 设置 K = B/L (等效全注意力)
4. 逐步减小 K 到目标稀疏率
5. 训练过程中 Gc 自动学习

八、总结

核心贡献

  1. 首个可训练稀疏注意力:在严格扩展性评估下优于全注意力

    • 训练FLOPS降低2.53×
    • 注意力FLOPS降低8×
    • 无扩散损失下降
  2. 硬件高效设计:

    • 保留85% FlashAttention3 MFU
    • 长序列**7×**加速(接近理论极限8×)
    • 块稀疏布局对齐GPU内存访问
  3. 实际部署价值:

    • Wan-1.3B推理:31s → 18s (1.7×)
    • Wan-14B推理:1274s → 576s (2.2×)
    • 首个注意力仅占**20%**运行时的视频DiT
  4. 稀疏蒸馏突破:

    • 首次将稀疏注意力与蒸馏结合
    • **50.9×**加速,无质量损失

技术影响

  • 范式转变:从”后验稀疏”到”训练时稀疏”
  • 扩展性验证:60M到1.4B模型、最高4×10^21 FLOPS
  • 实用性:成功适配SOTA开源模型Wan-2.1

局限性

  1. 固定立方体大小:(4,4,4)要求latent维度可被4整除
  2. 最优稀疏率未定:Top-K选择依赖序列长度和训练预算
  3. 仅验证视频DiT:未扩展到图像生成或多模态

未来方向

  1. 自适应稀疏率:根据层、时间步、序列长度动态调整Top-K
  2. 扩展Scaling Law:将稀疏率作为独立维度纳入扩展定律
  3. 更长序列验证:在>100K token序列上全面评估
  4. 跨模态扩展:将VSA应用于图像、3D、多模态生成

九、参考资源