一项关于在 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%——因此,稀疏注意力带来的端到端加速效果随
A case study in turning theoretical sparsity into an actual end-to-end speedup on TPUs.
Why is video diffusion slow?
Video diffusion models are typically slow at generating videos for two reasons: a large number of denoising steps, and high cost for each individual denoising step. At large video sequence lengths, self-attention can become one of the dominant contributors to the latency of an individual denoising step. The sequence length for just 81 frames of 720p video ranges from 50K to 400K for popular open source video generation models today.
Table 1: Self-attention latency and its share of single-layer transformer block latency across resolutions (81-frame video on eight TPU v6e chips).
As an example, scaling from 720p (HD) to 1440p (2K) can quadruple the sequence length. Because full attention scales quadratically with sequence length, its share of per-layer latency can grow from 55.5% to 88.2%.
Speeding up attention
Since attention is the single greatest driver of latency within a single step, any aggressive inference optimization would need to target attention first. Fortunately, attention in video diffusion is highly structured: many query-key interactions carry little attention mass, so a large fraction of the pairwise computation can often be skipped. Dense attention computes every interaction regardless of its importance, while sparse attention uses a mask to retain only the most important interactions and discard the rest. An example of an attention matrix from a video diffusion model is shown below to illustrate how tightly concentrated attention mass can be.
Figure 1: Attention mass for a spatial head (Head 12). In frame-major ordering (left), attention is tightly concentrated along local frames. In temporal ordering (right), that same mass appears dispersed across tokens—illustrating why token ordering must match head type.
However, the pattern of attention varies across heads and layers, and even across steps. Sparse VideoGen ( SVG ) is a formative work that recognizes and exploits one pattern of attention: different heads are often characterizable as either spatial heads or temporal heads . Importantly for implementation, these are not arbitrary sparse patterns, and both have highly regular geometric structure. Within a spatial head, a patch attends mostly to other patches within the same or close-by frames; within a temporal head, a patch attends to a small spatial region across a large number of frames. In the image below, the attention mass corresponding to a single query is shown for a spatial head (first row) and a temporal head (second row). In the spatial head, the query diffusely attends to all the tokens in its own frame and to adjacent frames. In the temporal head, its attention is limited to a narrower spatial region but across a larger number of frames.
Figure 2: Attention head specialization (Step 6, Layer 32). The cyan box marks the query token’s spatial (y,x) position on frame F06. In the spatial head (top), 94.8% of the attention mass spreads broadly across the query frame for 2D synthesis. In the temporal head (bottom), attention concentrates along that same spatial (y,x) tube across prior frames (F04–F05) to track localized motion.
Based on this structure, SVG dynamically profiles attention heads at inference time and routes each head to either a spatial or temporal mask. It does so by sampling a small number of queries, computing their attention outputs under dense, spatial, and temporal attention, and selecting the sparse mask whose output deviates least from the dense baseline. This allows SVG to preserve higher quality at a given sparsity level. We refer readers to the SVG paper for details of the profiling and routing algorithm. As shown in Figure 3 below, both masks also retain full attention to the first frame (F00)—acting as an attention sink to anchor global scene appearance—alongside a local band for frame-to-frame interactions.
Figure 3: Sparse attention masks in frame-major order. Left: Spatial mask (contiguous local band + first-frame sink at F00). Right: Temporal mask, which appears as fine-grained diagonal stripes in frame-major order—motivating our token permutation later on.
Implementing sparse attention on TPUs
The table below summarizes the custom JAX and Pallas Splash Attention kernel implementations we compare. All timings use synthetic BF16 inputs on a single TPU v6e device, with 75.6K tokens, 10 heads, and a head dimension of 128. The reference query and key/value tile sizes are 3328 and 1536. Sparse variants retain approximately 38.87% of query-key pairs. These are isolated attention-kernel timings; routing, token permutation, and communication between devices are outside this measurement.
Table 2: Progression of sparse Splash Attention kernel optimizations on a single TPU v6e chip (75,600 tokens, 10 heads, head dimension 128).
Figure 4: Tile traversal strategies corresponding to Table 2 (from left to right: B0 Dense, B1 Exact sparse, B2 Classified tiles, and B3 Balanced tile rounding).
Logical sparsity is not physical sparsity
Figure 5: Hardware tile classification for sparse local-window attention. Tiles are partitioned into full tiles (executed via a mask-free fast path), boundary tiles (requiring elementwise coordinate masking), and skipped tiles (omitted entirely to save compute and bandwidth).
Although SVG has a clever approach to select an attention mask, this theoretical sparsity needs to be converted into a speedup on real hardware. Modern attention implementations don’t materialize the entire attention matrix multiplication Q · Kᵀ; rather, they divide it into tiles such as Q[i₁:i₂] · K[j₁:j₂]ᵀ and accumulate necessary statistics across these tiles to compute the attention output SoftMax(QKᵀ / √d)V. Within a visited tile, scores for excluded query-key pairs are set to negative infinity before the softmax so that they contribute zero attention weight. This enforces the mask, but does not avoid computing those scores in the first place.
The mask divides tiles into three types. Full tiles contain only retained query-key pairs and need no intra-tile mask. Boundary tiles contain both retained and excluded pairs, so their scores need elementwise masking. Empty tiles (labeled Skipped tiles in the diagram) contain no retained pairs and can be skipped entirely. Logical sparsity counts excluded pairs, but the hardware savings depend on which tiles we can skip and how much work remains inside those we visit.
Splash attention already supports sparse masks and can skip empty tiles. However, this alone does not guarantee a speedup: how the kernel handles the tiles it visits also matters. In our initial sparse prototype (B1: naive block traversal), the mask retains 39% of query-key pairs, or about 61% logical sparsity. But the kernel visits 44% of outer tiles because boundary tiles also contain pairs that will be masked out. It skips the remaining tiles, but still evaluates the exact mask inside every visited tile. This implementation takes 96.37 ms compared to 78.70 ms for dense Splash: 22% higher latency despite skipping more than half the tiles.
Avoid intra-tile masking when there is nothing to mask
Evaluating the exact mask requires determining which query-key coordinates are valid and applying that predicate to the attention scores. On TPUs, performing this elementwise mask evaluation on every visited tile bottlenecks the Vector Processing Unit (VPU) and stalls the Matrix Multiply Unit (MXU), even when every pair in a tile is valid. We can avoid that unnecessary work by identifying full and boundary tiles before executing attention and giving them separate paths.
To fix this, our second iteration (B2: full/boundary tile specialization) gives full tiles a completely mask-free fast path, executing exact coordinate masking only on boundary tiles. Both paths contribute to the same attention output, with their partial results combined using the corresponding softmax normalization statistics. The retained query-key pairs are unchanged. Latency drops from 96.37 ms down to 54.12 ms—a 44% reduction compared to the naive sparse traversal (B1) and 31% faster than dense Splash. This comparison shows the benefit of specializing tile execution while preserving the exact sparse mask.
Tile-size tuning can impose a tradeoff
Smaller tiles can follow the mask boundary more closely, reducing the fraction of tiles requiring intra-tile masking. However, tile size also changes how the kernel groups computation and data movement, so minimizing boundary work does not necessarily minimize latency.
Figure 6: Large vs. small tile boundaries. The dark line traces the SVG mask boundary (including the vertical first-frame sink on the left). Smaller tiles follow the boundary more closely, reducing yellow boundary tiles at the cost of smaller compute tile efficiency.
To illustrate this tradeoff, we vary the query tile size (BQ) while keeping the key/value tile size (BKV) fixed at 1536. Every configuration already uses separate full and boundary paths. At BQ = 1024, only 16.95% of executed tiles require internal masking, but latency is 66.88 ms. Increasing BQ to 3328 raises that fraction to 27.55% while reducing latency to 54.12 ms. Increasing BQ further to 4864 raises latency again, to 58.62 ms. The plot shows this tradeoff: the configuration with the least masking work is not the fastest. These percentages describe work/tiles requiring internal masking, not time spent evaluating the mask.
Figure 7: Line plot showing kernel latency in milliseconds versus percentage of work requiring intra-tile masking across four query tile sizes on TPU v6e, demonstrating the latency minimum at BQ equals 3328.
Align the sparse mask with final execution
We can reduce boundary masking further by aligning the SVG mask with the tiles the kernel executes. Rounding a boundary outward includes some previously excluded pairs; rounding it inward removes some previously retained pairs. Balancing these choices slightly modifies the mask while approximately preserving the query-key pair budget. Unlike the previous optimization, this changes which interactions contribute to the output.
Rounding removes the need to enforce the exact SVG boundary within each visited tile. Because 75,600 tokens is not an exact multiple of the 3328 × 1536 tile dimensions, the edge tiles extending beyond the real sequence length still contain padding, and padded positions must contribute exactly zero attention weight. We therefore reuse the full/boundary distinction, with masking now needed only for tiles containing padding.
In our final configuration (B3: tile-aligned sparse traversal), this reduces the fraction of tiles requiring masking from 27.55% down to just 2.22%. Restricting masking strictly to boundary padding brings latency down to 32.76 ms—achieving a 2.40× speedup over dense Splash attention.
Together, these results show why sparse attention needs both a suitable mask and an efficient execution strategy: skip empty tiles, avoid masking full tiles, and align the overall mask with the tiles the hardware computes.
Supporting dynamic masks and token layouts on TPUs
We now have a faster sparse attention kernel, but integrating it into distributed inference requires dealing with dynamic token layouts and different masks for different heads.
Figure 8: End-to-end execution pipeline for dynamic spatio-temporal attention. Dynamic routing evaluates sampled queries to produce per-head boolean flags, which select token layouts on static-shape tensors before dispatching to the sparse Splash Attention kernel and restoring output order.
Support dynamic routing with fixed tensor shapes
The routing step returns a Boolean flag for each head to select spatial or temporal attention. These flags choose between the original and permuted token layouts while keeping the tensor dimensions fixed. Unlike dynamically splitting heads into variable-sized spatial and temporal subsets, changing a head’s assignment at runtime does not alter tensor shapes or trigger costly XLA recompilations. The same flags restore temporal heads to the original token order after attention, so downstream layers receive the ordering that they expect. Spatial heads retain their original ordering.
Arrange tokens for efficient temporal access
Figure 9: Token memory layout transformation. Frame-major ordering (F,H,W) stores spatial tokens of frame 1 consecutively. Permuting to temporal-major order (H,W,F) groups tokens from the same spatial position across successive frames (B1–B4) into contiguous memory for efficient temporal kernel traversal.
Permuting tokens from frame-major (F, H, W) to temporal order (H, W, F) places tokens at the same spatial position across frames next to each other—transforming the scattered diagonal stripes of the temporal mask (seen earlier in Figure 3, right) into a single contiguous band that the tiled sparse kernel can traverse efficiently. We tried avoiding this permutation and handling temporal access inside the kernel instead, but latency increased from 44.40 to 74.14 ms in our temporal-head benchmark, even though the faster path included additional overheads of permutation and restoration. This is in line with extra indexing and data rearrangement inside the kernel, although we did not measure those costs separately.
Perform layout transformations on local heads
Where we perform this permutation matters in distributed inference. With 40 heads distributed across four devices, each device initially holds all 40 heads for one-quarter of the sequence. An all-to-all exchange redistributes this into 10 heads per device, each with the complete sequence. We call this the head-local layout.
Permuting tokens before that exchange can require additional communication because a head’s sequence is still spread across devices. After the exchange, every token needed to reorder a device’s local heads is already available there. We therefore apply the permutation after the input head exchange, execute sparse attention, and recover the original token order before the return exchange. Routing is computed before the exchange, and each head carries its routing flag with it.
In a matched TPU v6e experiment on one captured step and layer, moving token reordering into the head-local region reduced total attention latency from 75.64 ms to 46.81 ms, a 38% reduction. This includes routing, permutation, communication, and the sparse kernel. It shows why the placement of layout transformations matters alongside the sparse kernel itself.
End-to-end results and takeaways
Having optimized both the sparse kernel and its surrounding overheads, we now measure how these improvements translate into faster denoising. We evaluate 720p video generation with 81 frames, 40 denoising steps, and eight TPU v6e chips. The table compares three schedules that retain progressively fewer attention pairs, using the same prompt and seed and a local dense control for each configuration. Denoising times are medians of three warm runs, with speedup computed as dense time divided by SVG time. FFmpeg PSNR measures similarity between the encoded SVG and dense videos; higher values indicate closer outputs.
Table 3: End-to-end 720p denoising latency (81 frames, 40 steps on eight TPU v6e chips) and output quality across SVG sparsity schedules.
The aggressive schedule reduces denoising time from 153.50 s to 119.86 s, a 1.28× speedup and 22% lower latency. These results reflect the combined implementation with separately tuned dense and sparse configurations, including their distributed layouts. Gains are smaller than the isolated kernel speedups because denoising also includes work outside sparse attention.
Because self-attention accounts for an increasingly large share of transformer-block latency as sequence length grows—rising from 55.5% at 720p to 72.5% at 1080p and 88.2% at 1440p (2K)—end-to-end speedups from sparse attention scale su
| 刊期 | 得分 | 排名 | 结果 |
|---|---|---|---|
| 2026-10-05 | 8.4 | 19 | 入选 |
| 2026-10-04 | 8.4 | 32 | 未入选 |
| 2026-10-03 | 8.4 | 42 | 未入选 |
| 2026-10-02 | 8.4 | 45 | 未入选 |