← 返回 2026-10-05 简报

在TPU上加速视频扩散模型的时空注意力机制

Accelerating Spatio-Temporal Attention for Video Diffusion on TPUs

语音播报
摘要
事件:Google团队在TPU v6e上优化视频扩散模型的时空注意力机制,将理论稀疏性转化为端到端加速,解决长序列生成延迟问题。 要点:利用空间与时间头的结构化稀疏性动态路由掩码,通过自定义JAX内核实现硬件级物理稀疏,显著降低计算开销。 影响:大幅降低高分辨率视频生成的推理成本,为开发者提供高效的TPU部署方案,推动高质量视频生成模型的实用化进程。

一项关于在 TPU 上将理论稀疏性转化为实际端到端加速效果的案例研究。

为什么视频扩散模型速度缓慢?

视频扩散模型在生成视频时通常较慢,原因有二:一是去噪步骤数量庞大,二是单个去噪步骤的成本高昂。在较长的视频序列长度下,自注意力机制可能成为单个去噪步骤延迟的主要贡献者之一。对于当今流行的开源视频生成模型而言,仅 81 帧 720p 视频的序列长度范围就在 50K 到 400K 之间。

表 1:不同分辨率下(81 帧视频在八个 TPU v6e 芯片上)自注意力的延迟及其占单层 Transformer 块延迟的比例。

例如,从 720p(HD)扩展到 1440p(2K)会使序列长度增加四倍。由于全注意力机制的复杂度随序列长度呈二次方增长,其占每层延迟的比例可以从 55.5% 增长到 88.2%。

加速注意力机制

由于注意力是单个步骤内延迟的最大驱动因素,任何激进的推理优化都必须首先针对注意力进行改进。幸运的是,视频扩散中的注意力具有高度结构化特征:许多查询-键(query-key)交互携带的注意力质量很少,因此通常可以跳过大量成对计算。密集注意力会计算每一次交互,无论其重要性如何;而稀疏注意力则使用掩码仅保留最重要的交互并丢弃其余部分。下面展示了一个来自视频扩散模型的注意力矩阵示例,以说明注意力质量可以多么集中。

图 1:空间头(Head 12)的注意力质量。在帧优先排序(左侧)中,注意力紧密集中在局部帧上。在时间排序(右侧)中,相同的注意力质量分散在各个 token 之间——这说明了为什么 token 排序必须与头类型相匹配。

然而,注意力的模式在不同头、不同层甚至不同步骤之间都有变化。稀疏视频生成器(Sparse VideoGen, SVG)是一项开创性工作,它识别并利用了一种注意力模式:不同的头通常可以被归类为空间头或时间头。对于实现而言,重要的是,这些并非任意稀疏模式,两者都具有高度规则的几何结构。在空间头中,一个补丁主要关注同一帧或邻近帧中的其他补丁;而在时间头中,一个补丁会关注跨大量帧的小空间区域。在下图中,展示了空间头(第一行)和时间头(第二行)中对应于单个查询的注意力质量。在空间头中,查询广泛地关注其自身帧及相邻帧中的所有 token。在时间头中,其注意力局限于更窄的空间区域,但跨越更多的帧。

图 2:注意力头的专业化(第 6 步,第 32 层)。青色方框标记了查询 token 在帧 F06 上的空间 (y,x) 位置。在空间头(顶部),94.8% 的注意力质量广泛分布在用于 2D 合成的查询帧上。在时间头(底部),注意力集中在先前帧(F04–F05)中相同的空间 (y,x) 管状区域,以追踪局部运动。

基于这种结构,SVG 在推理时动态分析注意力头,并将每个头路由到空间掩码或时间掩码。它通过采样少量查询,计算它们在密集、空间和注意力下的输出,并选择与密集基线偏差最小的稀疏掩码来实现这一目标。这使得 SVG 能够在给定的稀疏度水平下保持更高的质量。有关分析和路由算法的详细信息,请参阅 SVG 论文。如下文图 3 所示,两种掩码也都保留了对第一帧(F00)的全注意力——作为锚定全局场景外观的注意力汇(attention sink),同时辅以用于帧间交互的局部带。

图3:按帧主序排列的稀疏注意力掩码。左侧:空间掩码(连续的局部带 + F00处的首帧汇聚点)。右侧:时间掩码,在帧主序下呈现为细粒度的对角条纹——这激发了我们后续对token进行置换的动机。

在TPU上实现稀疏注意力

下表总结了我们要比较的自定义JAX和Pallas Splash Attention内核实现。所有计时均使用单个TPU v6e设备上的合成BF16输入,包含75.6K个token、10个头以及128的头维度。参考查询和键/值瓦片大小分别为3328和1536。稀疏变体保留了约38.87%的查询-键对。这些是孤立的注意力内核计时;路由、token置换和设备间的通信不在本测量范围内。

表2:单个TPU v6e芯片上稀疏Splash Attention内核优化的进展(75,600个token,10个头,头维度128)。

图4:对应于表2的瓦片遍历策略(从左到右:B0 密集、B1 精确稀疏、B2 分类瓦片、B3 平衡瓦片舍入)。

逻辑稀疏并非物理稀疏

图5:用于稀疏局部窗口注意力的硬件瓦片分类。瓦片被划分为全瓦片(通过无掩码的快速路径执行)、边界瓦片(需要逐元素坐标掩码)和跳过瓦片(完全省略以节省计算量和带宽)。

虽然SVG采用了一种巧妙的注意力掩码选择方法,但这种理论上的稀疏性需要转化为真实硬件上的加速效果。现代的注意力实现并不会物化整个注意力矩阵乘法 Q · Kᵀ;相反,它们将其划分为瓦片,如 Q[i₁:i₂] · K[j₁:j₂]ᵀ,并在这些瓦片上累积必要的统计信息以计算注意力输出 SoftMax(QKᵀ / √d)V。在访问的瓦片内部,被排除的查询-键对的分数会被设置为负无穷大,以便它们在softmax后贡献零注意力权重。这实现了掩码效果,但并未避免首先计算这些分数。

掩码将瓦片分为三种类型。全瓦片仅包含保留的查询-键对,无需片内掩码。边界瓦片同时包含保留和排除的对,因此其分数需要逐元素掩码。空瓦片(图中标记为跳过瓦片)不包含任何保留的对,可以完全跳过。逻辑稀疏计算的是排除的对数,但硬件节省取决于我们可以跳过哪些瓦片,以及在我们访问的瓦片中剩余多少工作量。

Splash注意力已经支持稀疏掩码并可以跳过空瓦片。然而,这本身并不能保证加速:内核如何处理它访问的瓦片同样重要。在我们的初始稀疏原型(B1:朴素块遍历)中,掩码保留了39%的查询-键对,即约61%的逻辑稀疏度。但是,由于边界瓦片也包含将被掩码排除的对,内核访问了44%的外部瓦片。它跳过了剩余的瓦片,但在每个访问的瓦片内仍会评估精确掩码。与密集Splash的78.70毫秒相比,该实现耗时96.37毫秒:尽管跳过了超过一半的瓦片,延迟却高出22%。

当没有需要掩码的内容时,避免片内掩码

评估精确掩码需要确定哪些查询-键坐标是有效的,并将该谓词应用于注意力分数。在TPU上,在每个访问的瓦片上进行这种逐元素掩码评估会瓶颈化向量处理单元(VPU)并阻塞矩阵乘法单元(MXU),即使瓦片中的每一对都是有效的。我们可以通过在执行注意力之前识别全瓦片和边界瓦片,并为它们提供独立的路径来避免这种不必要的工作。

为了解决这一问题,我们的第二次迭代(B2:完整/边界图块专业化)为完整图块提供了完全无需掩码的快速路径,仅在边界图块上执行精确的坐标掩码。两条路径共同贡献于相同的注意力输出,其部分结果通过相应的 softmax 归一化统计信息进行合并。保留的查询-键对保持不变。延迟从 96.37 毫秒降低至 54.12 毫秒——与朴素稀疏遍历(B1)相比减少了 44%,比密集 Splash 快 31%。这一比较展示了在保持精确稀疏掩码的同时专业化图块执行的好处。

图块大小调整可能带来权衡

较小的图块可以更紧密地跟随掩码边界,从而减少需要内部图块掩码的图块比例。然而,图块大小也会改变内核如何分组计算和数据移动,因此最小化边界工作并不一定能最小化延迟。

图 6:大图块与小图块的边界对比。深色线条描绘了 SVG 掩码边界(包括左侧的垂直第一帧汇聚点)。较小的图块更紧密地跟随边界,从而减少了黄色边界图块的数量,但以牺牲计算图块效率为代价。

为了说明这种权衡,我们在保持键/值图块大小(BKV)固定为 1536 的情况下,变化查询图块大小(BQ)。每种配置都已经使用了独立的完整和边界路径。在 BQ = 1024 时,仅有 16.95% 的执行图块需要内部掩码,但延迟为 66.88 毫秒。将 BQ 增加到 3328 使该比例上升至 27.55%,同时将延迟降低至 54.12 毫秒。进一步将 BQ 增加到 4864 又导致延迟回升至 58.62 毫秒。图表显示了这种权衡:掩码工作量最少的配置并不是最快的。这些百分比描述的是需要内部掩码的工作/图块比例,而非评估掩码所花费的时间。

图 7:折线图展示了在 TPU v6e 上,内核延迟(毫秒)与四种查询图块大小下需要内部图块掩码的工作百分比之间的关系,证明了在 BQ 等于 3328 时存在延迟最小值。

将稀疏掩码与最终执行对齐

我们可以通过将 SVG 掩码与内核执行的图块对齐来进一步减少边界掩码。向外舍入边界会包含一些之前被排除的对;向内舍入则会移除一些之前保留的对。平衡这些选择会略微修改掩码,同时大致保持查询-键对的数量预算。与之前的优化不同,这会改变哪些交互对输出有贡献。

舍去操作消除了在每个访问图块内强制执行精确 SVG 边界的需求。由于 75,600 个令牌不是 3328 × 1536 图块尺寸的精确倍数,超出实际序列长度的边缘图块仍然包含填充,且填充位置必须恰好贡献零注意力权重。因此,我们重新使用完整/边界的区分,现在仅需对包含填充的图块进行掩码。

在我们的最终配置(B3:与图块对齐的稀疏遍历)中,需要掩码的图块比例从 27.55% 降低至仅 2.22%。将掩码严格限制在边界填充上,将延迟降低至 32.76 毫秒——实现了比密集 Splash 注意力快 2.40 倍的速度提升。

这些结果共同表明,为什么稀疏注意力既需要合适的掩码,也需要高效的执行策略:跳过空图块、避免对完整图块进行掩码,并将整体掩码与硬件计算使用的图块对齐。

在 TPU 上支持动态掩码和令牌布局

我们现在拥有了一个更快的稀疏注意力内核,但将其集成到分布式推理中需要处理动态的令牌布局和不同头部的不同掩码。

图 8:动态时空注意力的端到端执行流水线。动态路由评估采样的查询以生成每个头部的布尔标志,这些标志在分派到稀疏 Splash 注意力内核之前选择静态形状张量上的令牌布局,并恢复输出顺序。

支持带有固定张量形状的动态路由

路由步骤为每个头返回一个布尔标志,以选择空间注意力或时间注意力。这些标志在保持张量维度固定的同时,在原始令牌布局和置换后的令牌布局之间进行选择。与将头动态拆分为可变大小的空间和 temporal 子集不同,在运行时更改头的分配不会改变张量形状,也不会触发昂贵的 XLA 重新编译。相同的标志会在注意力处理后将时间头恢复到原始的令牌顺序,以便下游层接收到它们期望的顺序。空间头保留其原始顺序。

安排令牌以实现高效的时间访问

图 9:令牌内存布局转换。帧主序 (F,H,W) 连续存储第 1 帧的空间令牌。将其置换为时间主序 (H,W,F) 会将来自连续帧中相同空间位置的令牌(B1–B4)分组到连续的内存中,以便高效地遍历时间内核。

将令牌从帧主序 (F, H, W) 置换为时间顺序 (H, W, F) 会使跨帧处于相同空间位置的令牌彼此相邻——这将之前图 3(右侧)中看到的时间掩码的分散对角条纹转换为单个连续带,使得分块稀疏内核可以高效地遍历。我们曾尝试避免这种置换,而是在内核内部处理时间访问,但在我们的时间头基准测试中,延迟从 44.40 毫秒增加到 74.14 毫秒,尽管更快的路径包含了置换和恢复的额外开销。这与内核内部的额外索引和数据重排相符,尽管我们没有单独测量这些成本。

在局部头上执行布局转换

我们在何处执行这种置换在分布式推理中至关重要。当 40 个头分布在四个设备上时,每个设备最初持有所有 40 个头,但仅针对序列的四分之一部分。All-to-all 交换将此重新分配为每个设备 10 个头,每个头都包含完整的序列。我们称之为头局部布局(head-local layout)。

在该交换之前置换令牌可能需要额外的通信,因为头的序列仍然分散在设备上。交换后,重排设备局部头所需的所有令牌都已可用。因此,我们在输入头交换之后应用置换,执行稀疏注意力,并在返回交换之前恢复原始令牌顺序。路由在交换之前计算,每个头都携带其路由标志。

在捕获的单个步骤和层的匹配 TPU v6e 实验中,将令牌重排移动到头局部区域将总注意力延迟从 75.64 毫秒降低到 46.81 毫秒,减少了 38%。这包括路由、置换、通信和稀疏内核。这表明布局转换的位置与稀疏内核本身同样重要。

端到端结果与要点

在优化了稀疏内核及其周围开销之后,我们现在衡量这些改进如何转化为更快的去噪过程。我们评估使用 81 帧、40 个去噪步骤和八个 TPU v6e 芯片的 720p 视频生成。该表格比较了三种保留逐渐减少的注意力对的调度方案,每种配置使用相同的提示和种子以及本地密集控制。去噪时间是三次预热运行的中位数,加速比计算为密集时间除以 SVG 时间。FFmpeg PSNR 衡量编码后的 SVG 与密集视频之间的相似度;数值越高表示输出越接近。

表 3:SVG 稀疏调度下的端到端 720p 去噪延迟(在八个 TPU v6e 芯片上处理 81 帧、40 步)及输出质量。

激进的调度计划将去噪时间从153.50秒缩短至119.86秒,实现了1.28倍的加速,并将延迟降低了22%。这些结果反映了密集和稀疏配置分别调优后的联合实现,包括它们的分布式布局。由于去噪过程还包含稀疏注意力之外的其他工作,因此整体增益小于孤立内核的加速幅度。

随着序列长度的增加,自注意力在Transformer块延迟中所占的比例越来越大——从720p时的55.5%上升到1080p时的72.5%,以及在1440p(2K)时的88.2%——因此,稀疏注意力带来的端到端加速效果随

本条评分 8.4 score-v1
  • 来源权威 8
    注册表 priority=8(Google Developers Blog)
  • 时效 0.4
    没有发布日期(nodate 源),给地板分 0.4(证明不了它新,不按新的加分)
  • 多源印证 0
    只有 1 家在报(无旁证)
  • 社区信号 0
    无社区数据(本管线走 RSS,HN 的 hn_fetcher 未接入)
历史
刊期得分排名结果
2026-10-05 8.4 19 入选
2026-10-04 8.4 32 未入选
2026-10-03 8.4 42 未入选
2026-10-02 8.4 45 未入选
原文链接:https://developers.googleblog.com/accelerating-spatio-temporal-attention-for-video-diffusion-on-tpus/
来源:Google Developers Blog
以上内容由 AI 自动翻译,仅供参考。
← 返回简报