通过时空稀疏注意力路由配合 Splash Attention 内核优化,在 1440p 视频生成上实现端到端推理加速 1.69 倍。
Video diffusion 为什么会慢?
Video diffusion 模型生成视频通常很慢,原因有两个:去噪步骤数量多,以及单个去噪步骤的计算成本高。在较长的视频序列长度下,self-attention 可能成为单个去噪步骤延迟的主要贡献因素。仅 81 帧 720p 视频的序列长度,在目前流行的开源视频生成模型中就达到了 50K 到 400K 不等。
举例来说,从 720p(HD)扩展到 1440p(2K)会使序列长度翻两番。由于全量 attention 与序列长度呈二次关系,其在每层延迟中所占的比例会从 55.5% 上升到 88.2%。
既然 attention 是单步延迟的最大驱动因素,任何激进的推理优化都必须首先针对 attention。幸运的是,视频 diffusion 中的 attention 具有高度结构化特性:许多 query-key 之间的交互承载的关注量很小,因此通常可以跳过很大一部分两两计算。全量 attention 不论重要性如何都会计算每一次交互,而稀疏 attention 则使用掩码来只保留最重要的交互、丢弃其余部分。下面的注意力矩阵示例来自一个视频 diffusion 模型,用以说明注意力集中的程度。
然而,attention 的模式在不同 heads、不同层之间,甚至在不同步骤之间都有所不同。Sparse VideoGen(SVG)是一项开创性的工作,它识别并利用了一种 attention 模式:不同的 head 通常可以被描述为 spatial head 或 temporal head。对于实现而言重要的是,这些并非任意的稀疏模式,两种都具有高度规则的几何结构。在 spatial head 内,一个 patch 主要关注同一帧或相邻帧中的其他 patch;在 temporal head 内,一个 patch 关注跨越大量帧的较小空间区域。下图中展示了一个 query 对应的注意力权重,分别来自 spatial head(第一行)和 temporal head(第二行)。在 spatial head 中,query 弥散地关注其自身帧中的所有 token 以及相邻帧。在 temporal head 中,其注意力被限制在更窄的空间区域内,但覆盖了更多的帧。
基于这种结构,SVG 在推理时动态地对 attention heads 进行性能剖析,并将每个 head 路由到 spatial 或 temporal 掩码。它通过采样少量 query,计算它们在全量、spatial 和 temporal attention 下的输出,然后选择与全量基线偏差最小的稀疏掩码来实现这一点。这使得 SVG 能够在给定的稀疏水平下保持更高的质量。我们建议读者参阅 SVG 论文以了解性能剖析和路由算法的详细信息。如图 3 所示,两种掩码也都保留了对第一帧(F00)的全量注意力——作为注意力汇点来锚定全局场景外观——以及帧间交互的局部带区。
下表总结了我们比较的自定义 JAX 和 Pallas Splash Attention 内核实现。所有时间测量均使用单 TPU v6e 设备上的合成 BF16 输入,75.6K tokens,10 个 heads,head dimension 为 128。参考 query 和 key/value 的 tile 大小分别为 3328 和 1536。稀疏变体保留约 38.87% 的 query-key 对。这些是独立的 attention 内核计时;路由、token 置换和设备间通信不在此测量范围内。
尽管 SVG 采用了聪明的方法来选择 attention 掩码,但这种理论上的稀疏性需要转化为真实硬件上的加速。现代 attention 实现并不会将整个注意力矩阵乘法 Q · Kᵀ 实例化;相反,它们将其划分为多个 tile(如 Q[i₁:i₂] · K[j₁:j₂]ᵀ),并通过这些 tile 累积必要的统计量来计算 attention 输出 SoftMax(QKᵀ / √d)V。在访问的 tile 内,被排除的 query-key 对的分数在被 softmax 之前被设置为负无穷,这样它们就不会贡献注意力权重。这强制执行了掩码,但并没有首先避免计算这些分数。
掩码将 tile 分为三类。全量 tile(Full tiles)只包含保留的 query-key 对,不需要 tile 内的掩码。边界 tile(Boundary tiles)同时包含保留和排除的对,因此它们的分数需要逐元素掩码。空 tile(图中有 Skipped tiles 标记)不包含任何保留的对,可以完全跳过。逻辑稀疏性统计的是被排除的对,但硬件节省取决于我们能跳过哪些 tile 以及在我们访问的 tile 内还剩多少工作。
Splash attention 已经支持稀疏掩码并能跳过空 tile。然而,仅这一点并不能保证加速:内核如何处理它访问的 tile 也很重要。在我们最初的稀疏原型(B1:naive block traversal)中,掩码保留了 39% 的 query-key 对,即约 61% 的逻辑稀疏性。但内核访问了 44% 的外层 tile,因为边界 tile 也包含了将被掩码排除的对。它跳过了剩余的 tile,但仍然在每个被访问的 tile 内精确评估掩码。该实现耗时 96.37 ms,而全量 Splash 只需 78.70 ms:尽管跳过了一半以上的 tile,延迟反而高出 22%。
精确评估掩码需要确定哪些 query-key 坐标是有效的,并将该谓词应用到注意力分数上。在 TPU 上,对每个被访问的 tile 执行这种逐元素掩码评估会成为向量处理单元(VPU)的瓶颈,并使矩阵乘法单元(MXU)停滞——即使 tile 中的每一对都是有效的。我们可以通过在执行 attention 之前识别全量和边界 tile,并为它们分别提供处理路径来避免这种不必要的 work。
为解决这个问题,我们的第二次迭代(B2:full/boundary tile specialization)给全量 tile 提供了一条完全无掩码的快速路径,只在边界 tile 上执行精确坐标掩码。两条路径贡献相同的 attention 输出,它们的部分结果使用相应的 softmax 归一化统计量进行合并。保留的 query-key 对保持不变。延迟从 96.37 ms 降至 54.12 ms——比 naive 稀疏遍历(B1)减少 44%,比全量 Splash 快 31%。这一比较表明,在保留精确稀疏掩码的同时,专门化 tile 执行是有效的。
较小的 tile 能更紧密地跟随掩码边界,减少需要 tile 内掩码的 tile 比例。然而,tile 大小也会改变内核对计算和数据移动的分组方式,因此最小化边界 work 并不一定最小化延迟。
为说明这种权衡,我们在 key/value tile 大小(BKV)固定为 1536 的同时改变 query tile 大小(BQ)。每种配置都已经使用了独立的 full 和 boundary 路径。当 BQ = 1024 时,只有 16.95% 的执行 tile 需要内部掩码,但延迟为 66.88 ms。将 BQ 增加到 3328 会使该比例上升到 27.55%,同时将延迟降低到 54.12 ms。进一步将 BQ 增加到 4864 会再次提高延迟至 58.62 ms。图示显示了这种权衡:掩码 work 最少的配置并非最快的配置。这些百分比描述的是需要内部掩码的 work/tile 比例,而不是评估掩码所花费的时间。
我们可以通过将 SVG 掩码与内核执行的 tile 对齐来进一步减少边界掩码。将边界向外舍入会包含一些先前被排除的对;向内舍入则会移除一些先前保留的对。平衡这些选择会轻微修改掩码,同时大致保留 query-key 对的预算。与之前的优化不同,这会改变哪些交互对输出有贡献。
舍入消除了在每个被访问的 tile 内强制执行精确 SVG 边界的需要。由于 75,600 个 token 不是 3328 × 1536 tile 维度的整数倍,延伸到实际序列长度之外的边缘 tile 仍然包含填充,而填充位置必须贡献恰好为零的注意力权重。因此我们复用 full/boundary 区分,现在只需要对包含填充的 tile 进行掩码。
在我们最终的配置中(B3:tile-aligned sparse traversal),这将需要掩码的 tile 比例从 27.55% 降低到仅 2.22%。将掩码严格限制在边界填充上使延迟降至 32.76 ms——实现了比全量 Splash attention 快 2.40 倍的加速。
综合来看,这些结果表明,稀疏 attention 需要同时具备合适的掩码和高效的执行策略:跳过空 tile,避免对全量 tile 进行掩码,并将整体掩码与硬件计算的 tile 对齐。
我们现在有了更快的稀疏 attention 内核,但将其集成到分布式推理中需要处理动态 token 布局和不同 head 的不同掩码。
路由步骤返回每个 head 的布尔标志,以选择 spatial 或 temporal attention。这些标志在保持张量维度不变的情况下在原始和置换后的 token 布局之间进行选择。与动态将 head 拆分为可变大小的 spatial 和 temporal 子集不同,在运行时更改 head 的分配不会改变张量形状或触发昂贵的 XLA 重编译。相同的标志在 attention 之后将 temporal head 恢复到原始 token 顺序,因此下游层接收到它们期望的顺序。Spatial head 保持其原始顺序。
将 token 从帧主序(F, H, W)置换到时序顺序(H, W, F),可以将跨帧处于相同空间位置的 token 放在一起——将时序掩码(之前在图 3 右侧看到的那种分散的对角条纹)转换为单个连续带区,使 tiled 稀疏内核能够高效遍历。我们尝试过避免这种置换而在内核内部处理时序访问,但在我们对 temporal head 的基准测试中,延迟从 44.40 增加到 74.14 ms,尽管更快的那条路径包含了置换和恢复的额外开销。这与内核内部的额外索引和数据重排是一致的,尽管我们没有单独测量这些成本。
我们执行这种置换的位置在分布式推理中很重要。当 40 个 head 分布在四个设备上时,每个设备最初为序列的四分之一持有全部 40 个 head。全到全(all-to-all)通信将这重新分配为每个设备 10 个 head,每个 head 具有完整序列。我们称之为 head-local 布局。
在交换之前置换 token 可能需要额外的通信,因为一个 head 的序列仍然分散在各个设备上。交换之后,重排序设备本地 head 所需的每个 token 都已在那里。因此,我们在输入 head 交换之后应用置换,执行稀疏 attention,并在返回交换之前恢复原始 token 顺序。路由在交换之前计算,每个 head 携带其路由标志。
在一对一捕获的步骤和层上的匹配 TPU v6e 实验中,将 token 重排序移入 head-local 区域将总 attention 延迟从 75.64 ms 降至 46.81 ms,减少了 38%。这包括路由、置换、通信和稀疏内核。这表明布局转换的位置与稀疏内核本身同样重要。
在优化了稀疏内核及其周边开销之后,我们现在来衡量这些改进如何转化为更快的去噪。我们评估了 81 帧 720p 视频生成,40 个去噪步骤,八个 TPU v6e 芯片。该表比较了三种逐步减少注意力对的保留比例的调度,使用相同的 prompt 和 seed,并为每种配置提供本地全量对照。去噪时间为三次预热运行的中位数,加速比计算为全量时间除以 SVG 时间。FFmpeg PSNR 衡量编码后 SVG 和全量视频之间的相似度;值越高表示输出越接近。
激进的调度将去噪时间从 153.50 s 降至 119.86 s,实现了 1.28 倍加速,延迟降低 22%。这些结果反映了结合了分别调优的全量和稀疏配置的组合实现,包括它们的分布式布局。收益小于独立的内核加速,因为去噪还包括稀疏 attention 之外的工作。
由于 self-attention 在 transformer block 延迟中所占的比例随序列长度增长而增加——从 720p 的 55.5% 上升到 1080p 的 72.5%,再到 1440p(2K)的 88.2%——稀疏 attention 的端到端加速随分辨率大幅扩展。在不同分辨率上评估激进调度可以显示这些节省如何在更长的序列上累积:
在 1080p(171K tokens)下,激进调度将去噪时间从 683 s 降至 457 s(1.49 倍加速,PSNR 24.66 dB)。在 1440p(302K tokens,是 720p 序列长度的四倍)下,去噪时间从 2,471 s 降至 1,461 s——实现了 1.69 倍端到端加速,每次生成视频节省超过 16 分钟,同时保持 24.05 dB PSNR。
Sorry, your browser doesn't support playback for this video
关键要点:TPU 上高效的稀疏 attention 取决于硬件能避免多少 work,而不仅仅是掩码去除了多少注意力对。跳过空 tile、限制 tile 内掩码、以及为高效访问安排 token,帮助我们将算法稀疏性转化为实际的端到端生成加速。