想象一下,你正在尝试绘制一幅持续 10 秒的超高清巨型壁画。为此,你拥有一支由 120 亿名微小艺术家(即“令牌”)组成的团队,他们需要彼此交流,以确保色彩完美融合、动作流畅自然。
在当前的技术状态下,每一位艺术家都必须停下来,向每一位其他艺术家耳语一个秘密,以决定下一步该画什么。这被称为“全注意力机制”。虽然它能产生精美的效果,但速度极慢且成本高昂。如果你将壁画规模扩大一倍,协调所需的时间并非仅仅翻倍,而是会变为四倍。这就像试图组织一场派对,要求所有人在音乐开始前必须与彼此一一握手。
“走捷径”的问题
研究人员曾尝试通过让艺术家只与少数邻居交流(稀疏注意力)来加速这一过程。然而,以往的方法就像一位笨拙的经理,只是告诉艺术家们:“只和你同一排的邻居交谈。”这导致壁画看起来怪异:画中的水波会泛起奇异的涟漪,人脸会发生扭曲,视频会出现闪烁。艺术家们错过了维持画面连贯性所必需的重要对话。
解决方案:Veda(智能工头)
该论文介绍了Veda,这是一个新系统,它像一位敏锐、观察入微的工头。Veda 并非猜测谁需要与谁交谈,而是利用一种“蒸馏”技巧。
以下是其工作原理,使用一个简单的类比:
- “神谕”地图:想象一下,缓慢而完美的“全注意力”方法就像一位总建筑师,他绘制了一张包含所有应该发生的对话的完整地图。这张地图完美无缺,但绘制起来耗时极长。
- 聪明的学生:Veda 训练一个轻量级的“学生”(一个小型、快速的 AI),让它查看总建筑师的地图,并学习谁与谁交谈的模式。
- “三池”技巧:以往的学生试图通过取对话的“平均值”来总结地图。但在视频中,最重要的往往是一个单一的、响亮的呼喊(峰值信号),而不是平均的嘈杂声。Veda 的学生更聪明:它会查看对话的最大值、最小值和平均值。这确保了它不会遗漏那些保持视频稳定性的关键“响亮”时刻。
- 头感知分块:视频包含许多不同的“头”(专门团队)。有些团队关注事物随时间的变化(时间维度),而另一些则关注事物在空间中的外观(空间维度)。Veda 意识到“一刀切”的规则行不通。它为每个团队提供定制大小的拼图块(分块),以契合其特定任务,确保它们不会遗漏重要细节。
结果:速度提升且无模糊
一旦 Veda 学会了这张地图,它就会告诉艺术家们:“你们只需要与这特定的 5% 的人交谈。忽略其余的人。”
因为 Veda 确切知道谁可以忽略而不丢失重要连接,视频生成变得极其迅速:
- 速度:它将生成高清 10 秒视频的速度提高了5.1 倍。
- 质量:与以往导致视频出现故障或扭曲的方法不同,Veda 生成的视频看起来与缓慢的完美版本一样好。
- 可扩展性:视频越长、细节越丰富,Veda 的优势就越明显。它几乎以线性方式处理额外工作,而旧方法则会陷入交通堵塞。
总结
Veda 就像雇佣了一位超级聪明的工头,他研究完美计划,学习哪些对话至关重要,然后指挥团队跳过无聊的寒暄。这使得他们能够在极短的时间内绘制出一幅巨大的高质量壁画,而不会让画面支离破碎。
技术摘要:Veda - 通过蒸馏稀疏注意力实现可扩展的视频扩散
1. 问题陈述
将扩散 Transformer(DiT)扩展用于高分辨率、长时程视频生成,目前受限于自注意力机制相对于时空 Token 序列长度的二次方计算和内存复杂度(O(N2))。尽管稀疏注意力提供了一条理论上缓解此问题的路径,但现有方法面临一个关键权衡:
- 静态模式:使用预定义稀疏模式(如滑动窗口)的方法缺乏对 DiT 学习到的动态、特定于头的注意力结构的适应性。
- 动态选择:学习动态掩码的方法(如 VSA、VMOBA)往往遭受估计不准确的问题。它们依赖扩散目标进行的隐式监督以及粗略统计量(如平均池化),导致与完整注意力的真实几何结构不一致。
- 核心问题:实证证据表明,在高稀疏度(≥90%)下,这些现有方法会降低生成质量,产生结构化伪影,如空间扭曲、“水波纹”图案和时间闪烁。本文认为,质量下降并非由稀疏性本身引起,而是由稀疏掩码未能与完整注意力的分块几何结构对齐所导致。
2. 方法论:Veda 框架
Veda 是一个蒸馏稀疏注意力框架,旨在即使在极端稀疏度下也能保留完整注意力的分块结构和排序。它将分块选择视为一个显式的重建问题,而非扩散目标的隐式副产品。
2.1. 蒸馏分块评分
Veda 不依赖扩散损失来塑造稀疏性,而是采用一个轻量级估计器,从完整注意力骨干网络中显式学习分块级注意力分数。
- 目标构建:通过对查询 - 键分块区域上的完整注意力矩阵应用最大池化,构建参考“Oracle"掩码。选择最大池化而非平均池化,是为了保留通常被背景噪声稀释的显著高频信号峰值。
- 统计感知估计器:为了重建这些分数,Veda 为每个分块使用TripPool描述符,连接平均、最大和最小统计量。这丰富了分块表示,超越了简单的平均池化。
- 特定于头的投影:估计器为每个注意力头使用不同的 MLP 投影,以捕捉异质的依赖模式。
- 优化:模型使用逐行KL 散度蒸馏目标(Ldistill)进行训练,以使预测的分块分数分布与完整注意力参考对齐。关键在于,对馈送到估计器的骨干特征应用了停止梯度操作。这将掩码学习与特征学习解耦,防止破坏预训练的生成流形。
2.2. 感知头的分块
认识到注意力头在时空依赖方面表现出显著的异质性(有些捕捉局部空间交互,有些捕捉长程时间依赖),Veda 放弃了均匀分块策略。
- 配置搜索:对于每一层和每一个头,Veda 搜索最优分块配置 πl,h=(pt,ph,pw),该配置对硬件分块大小进行因式分解。
- 离线选择:最优配置通过在校准集上最小化完整注意力输出与稀疏注意力输出之间的近似误差来离线选择,确保分块策略匹配每个头的特定结构需求。
2.3. 硬件高效实现
为了将理论上的 FLOP 减少转化为实际的挂钟时间加速,Veda 使用 ThunderKittens DSL 实现了自定义的分块跳过稀疏注意力内核。
- 异步执行:该内核利用 NVIDIA Hopper 的张量内存访问(TMA)和 Warp 专用化。它使用生产者 - 消费者范式将数据移动与计算解耦,其中生产者 Warp 将非连续的键/值分块获取到共享内存中,而消费者 Warp 执行张量核心操作。
- 效率:该设计将内存延迟隐藏在稠密矩阵运算之后,实现了约 80% 的 FlashAttention-3 内存 FLOPs 利用率(MFU)。
- 训练效率:专用的 TileLang 内核通过两次传递生成真实值热力图,允许使用稀疏监督(查询分块的随机子集)来减少训练开销,同时不损害性能。
3. 主要贡献
- 实证洞察:本文证明,稀疏视频扩散中的生成质量取决于稀疏掩码与完整注意力分块几何结构的对齐程度,而非稀疏率本身。
- 蒸馏稀疏注意力:Veda 引入了一个框架,将分块选择表述为显式重建问题,利用统计感知估计器和特定于头的投影来最小化估计误差。
- 感知头的分块:一种新颖策略,为不同的注意力头分配不同的时空分块因式分解,解决了视频 DiT 中依赖模式的异质性问题。
- 硬件优化内核:一个自定义的分块跳过内核,实现了与序列长度的近线性扩展和高硬件利用率,将算法稀疏性转化为可感知的延迟收益。
4. 实验结果
实验在大规模视频扩散模型上进行,包括 Waver-T2V-12B 和 Wan2.1-T2V-14B,分辨率高达 720P,序列长度高达 241 帧。
- 加速比:在 Waver-T2V-12B(720P,241 帧)上,Veda 实现了 5.1 倍的端到端加速(将采样时间从 19.4 分钟减少到 3.8 分钟)和 10.5 倍的自注意力加速。注意力开销从 92% 降低到 50%。
- 可扩展性:加速比随序列长度增加而提升。在 245K 个 Token 时,Veda 呈现近线性扩展(309.1 毫秒),而完整注意力则呈二次方增长(1576.5 毫秒),实现了 5.1 倍的加速。
- 质量保持:
- 人类评估:在 Waver-bench 1.0 上,90% 稀疏度下的 Veda 实现了与完整注意力的感知对等。在 95% 稀疏度下,它显著优于最先进的 VSA 方法(即使 VSA 在较低稀疏度下运行)。
- 伪影减少:与在高稀疏度下遭受水波纹图案和时间闪烁的基线不同,Veda 保持了高视觉保真度和时间连贯性。
- 定量指标:VBench 评估显示,即使在 95% 稀疏度下,Veda 在主体一致性和运动平滑度方面仍保持与完整注意力基线相当的水平。
5. 意义与主张
本文声称,Veda 成功解决了高分辨率视频生成的可扩展性瓶颈,同时不损害结构完整性。通过将稀疏注意力重新框架化为蒸馏问题,而非启发式搜索或隐式学习任务,Veda 实现了以前无法达到的激进稀疏度(高达 95%),而不会出现严重的质量下降。
作者强调,这些收益不仅仅是理论上的;自定义内核确保了稀疏性转化为实际的挂钟时间加速。该框架被呈现为一种模块化、感知硬件的解决方案,能够随时空分辨率有利地扩展,使得在当前硬件约束下实现高保真、长视频生成成为可能。这项工作表明,未来的改进可以集中在更紧密的内核融合、跨时间步的自适应稀疏策略以及分块分数的时间缓存上。
每周获取最佳 computer science 论文。
受到斯坦福、剑桥和法国科学院研究人员的信赖。
请查收邮箱确认订阅。
出了点问题,再试一次?
无垃圾邮件,随时退订。