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”上,其余接近零。
关键创新:
- 分层粗细注意力:粗阶段定位关键token,细阶段仅在关键区域内计算
- 端到端可训练:通过数据学习稀疏模式,而非启发式规则
- 硬件对齐:将空间-时间立方体映射到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大小 B | 64 (4×4×4) | 表达性与效率的最佳平衡 |
| Top-K | 32 | 87.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 KV | 0.15281 | 0.14282 |
| Spatial Temporal | 0.13574 | 0.13034 |
| Spatial Full | 0.13555 | 0.12811 |
| Strided Window | 0.13271 | 0.12716 |
| Full Attention | 0.13877 | 0.12703 |
| VSA | 0.13162 | 0.12687 |
发现:VSA在计算最优和过训练设置下都优于全注意力和其他稀疏方法。
(b) 注意力设计消融
| 设计 | Loss |
|---|---|
| L (局部窗口) | 0.13330 |
| F (仅细阶段) | 0.13296 |
| C & L | 0.13220 |
| C & F | 0.13162 |
| C & F & L | 0.13194 |
发现:粗阶段(C)+细阶段(F)的组合最优,局部先验(L)无显著帮助。
(d) Tile大小消融
| Tile大小 | TFLOPS | Loss |
|---|---|---|
| 256×256 | 478 | 0.13375 |
| 128×128 | 444 | 0.13244 |
| 64×64 | 408 | 0.13162 |
| 16×64 | 181 | 0.13155 |
发现:更小tile性能更好但速度更慢,64×64是最佳平衡点。
扩展性实验 (Figure 2)
| 模型规模 | 训练FLOPS | VSA Loss | Full Loss | 注意力FLOPS节省 |
|---|---|---|---|---|
| 410M | 4×10^21 | ≈0.126 | ≈0.126 | 8× |
| 60M-1.4B | 最高4×10^21 | Pareto优 | - | 2.53× (总训练FLOPS) |
关键发现:
- VSA在410M模型上达到与全注意力几乎相同的loss曲线
- 60M到1.4B的扩展实验确认VSA持续产生更好的Pareto前沿
- VSA是首个在严格扩展性评估下优于全注意力的可训练稀疏注意力
稀疏适配与蒸馏 (Section 3.3)
Wan-1.3B VBench结果
| 模型 | 质量分 | 语义分 | 总分 |
|---|---|---|---|
| 原始Wan | 83.71% | 77.98% | 82.56% |
| 全注意力微调 | 84.07% | 81.85% | 83.63% |
| VSA微调 | 83.60% | 79.47% | 82.77% |
推理加速:
| 模型 | 全注意力 | VSA | 加速比 |
|---|---|---|---|
| Wan-1.3B (480P) | 31s | 18s | 1.7× |
| Wan-14B (720P) | 1274s | 576s | 2.2× |
稀疏蒸馏:
- 学生模型使用VSA + DMD2蒸馏
- 教师模型保持全注意力
- 结果:50.9×加速,无质量损失
内核性能 (Section 3.4)
| 指标 | 数值 |
|---|---|
| VSA fine kernel MFU | 85% 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 Attention | profile驱动稀疏 | 训练-测试不匹配 |
| SVG | 无训练稀疏化 | 稀疏率受限,质量下降 |
| DSV | 训练时稀疏 | 多阶段设计复杂 |
| VSA | 端到端可训练 | - |
VSA的独特性
视频DiT对可训练稀疏注意力的需求比LLM更紧迫:
- 视频DiT序列更长(100K tokens仅5秒视频)
- 视频DiT将大部分计算预算用于长序列训练
- LLM的”短训练-长适配”范式在视频DiT中不适用
七、代码实现分析
项目信息
- 仓库: https://github.com/hao-ai-lab/FastVideo
- 内核框架: ThunderKittens
- 支持模型: Wan-2.1 (1.3B, 14B), HunyuanVideo
内核实现细节
粗阶段内核:
融合操作: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 自动学习
八、总结
核心贡献
-
首个可训练稀疏注意力:在严格扩展性评估下优于全注意力
- 训练FLOPS降低2.53×
- 注意力FLOPS降低8×
- 无扩散损失下降
-
硬件高效设计:
- 保留85% FlashAttention3 MFU
- 长序列**7×**加速(接近理论极限8×)
- 块稀疏布局对齐GPU内存访问
-
实际部署价值:
- Wan-1.3B推理:31s → 18s (1.7×)
- Wan-14B推理:1274s → 576s (2.2×)
- 首个注意力仅占**20%**运行时的视频DiT
-
稀疏蒸馏突破:
- 首次将稀疏注意力与蒸馏结合
- **50.9×**加速,无质量损失
技术影响
- 范式转变:从”后验稀疏”到”训练时稀疏”
- 扩展性验证:60M到1.4B模型、最高4×10^21 FLOPS
- 实用性:成功适配SOTA开源模型Wan-2.1
局限性
- 固定立方体大小:(4,4,4)要求latent维度可被4整除
- 最优稀疏率未定:Top-K选择依赖序列长度和训练预算
- 仅验证视频DiT:未扩展到图像生成或多模态
未来方向
- 自适应稀疏率:根据层、时间步、序列长度动态调整Top-K
- 扩展Scaling Law:将稀疏率作为独立维度纳入扩展定律
- 更长序列验证:在>100K token序列上全面评估
- 跨模态扩展:将VSA应用于图像、3D、多模态生成
九、参考资源
- 论文: https://arxiv.org/abs/2505.13389
- 代码: https://github.com/hao-ai-lab/FastVideo
- FlashAttention: https://arxiv.org/abs/2205.14135
- ThunderKittens: https://github.com/Hagedorn-Lab/ThunderKittens
- Wan-2.1: https://github.com/Wan-Video/Wan2.1
- NSA (Native Sparse Attention): 受VSA启发的LLM稀疏注意力
- MoBA: 混合块注意力,VSA的灵感来源之一