Accelerating Spatio-Temporal Attention for Video Diffusion on TPUs
Google Developers Blog

Accelerating Spatio-Temporal Attention for Video Diffusion on TPUs

Accelerating Spatio-Temporal Attention for Video Diffusion on TPUs

Problem Overview

Self-attention in high-resolution video diffusion models suffers from a quadratic latency bottleneck. For large video sequence lengths, self-attention becomes a dominant contributor to per-step latency. Scaling from 720p (HD) to 1440p (2K) can quadruple the sequence length, causing full attention's share of per-layer latency to grow from 55.5% to 88.2%. Since attention is the primary driver of latency within a single step, any aggressive inference optimization must target attention first.

Video diffusion models face two main challenges: a large number of denoising steps and high cost per individual denoising step. At 81 frames of 720p video, sequence lengths range from 50K to 400K. Scaling from 720p to 1440p quadruples the sequence length. Due to the quadratic nature of attention, its contribution to per-step latency increases dramatically. Sparse attention can help by retaining only the most important interactions while discarding less significant ones.

Sparse VideoGen (SVG) Approach

Sparse VideoGen (SVG) recognizes and exploits a recurring pattern in video diffusion attention: heads naturally categorize into spatial heads or temporal heads. Each category exhibits highly regular geometric structure:

  • Spatial heads - patches attend mostly to other patches within the same or close-by frames.
  • Temporal heads - patches attend to a small spatial region across a large number of frames.

SVG dynamically profiles attention heads at inference time and routes each head to either a spatial or temporal mask. It samples a small number of queries, computes their attention outputs under dense, spatial, and temporal attention, and selects the sparse mask whose output deviates least from the dense baseline. This preserves higher quality at a given sparsity level.

Both masks retain full attention to the first frame (F00) as an "attention sink" for anchoring global scene appearance, alongside a local band for frame-to-frame interactions.

Hardware Optimizations on TPUs

Translating algorithmic sparsity into physical hardware speedups on TPUs required optimizing the Splash Attention kernel through several strategies:

  • Bypassing empty memory tiles - skip tiles that contain no retained query-key pairs entirely.
  • Restricting exact coordinate masking strictly to boundary tiles - only apply expensive masking to tiles that actually contain mixed retained/excluded pairs.
  • Permuting token memory layouts into temporal-major order - ensures contiguous access for the tiled sparse kernel to traverse efficiently.

By aligning sparse masks with actual hardware tile execution, these optimizations significantly reduce wasted matrix operations and achieve up to a 1.69x end-to-end inference speedup for 1440p video generation.

Tile-Based Attention Mechanics

Modern attention implementations do not materialize the entire attention matrix multiplication (Q \cdot K^T). Instead, they divide it into tiles such as (Q[i_1:i_2] \cdot K[j_1:j_2]^T) and accumulate statistics across tiles to compute (\text{SoftMax}(QK^T / \sqrt{d})V).

Within a visited tile, scores for excluded query-key pairs are set to negative infinity before softmax so they contribute zero attention weight. The mask divides tiles into three types:

  • Full tiles - contain only retained query-key pairs and need no intra-tile masking.
  • Boundary tiles - contain both retained and excluded pairs, so their scores require element-wise masking.
  • Empty tiles (labeled Skipped tiles) - contain no retained pairs and can be skipped entirely.

Logical sparsity counts excluded pairs, but hardware savings depend on which tiles can be skipped and how much work remains inside visited tiles.

Iteration 1: Naive Block Traversal (B1)

In the initial sparse prototype (B1), the mask retains 39% of query-key pairs (approximately 61% logical sparsity). However, the kernel visits 44% of outer tiles because boundary tiles also contain pairs that get masked out. Despite skipping more than half the tiles, it still evaluates the exact mask inside every visited tile. This implementation took 96.37 ms compared to 78.70 ms for dense Splash - a 22% increase in latency.

The core bottleneck was evaluating the exact mask on every visited tile. On TPUs, this element-wise mask evaluation on visited tiles bottlenecks the Vector Processing Unit (VPU) and stalls the Matrix Multiply Unit (MXU), even when every pair in a tile is valid.

Iteration 2: Full/Boundary Tile Specialization (B2)

The second iteration (B2) 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 partial results combined using corresponding softmax normalization statistics. Retained query-key pairs remain unchanged.

Latency dropped from 96.37 ms to 54.12 ms - a 44% reduction compared to the naive sparse traversal (B1) and 31% faster than dense Splash.

Iteration 3: Tile-Aligned Sparse Traversal (B3)

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. Because 75,600 tokens is not an exact multiple of the 3328 × 1536 tile dimensions, edge tiles extending beyond the real sequence length still contain padding that must contribute exactly zero attention weight.

We therefore reuse the full/boundary distinction, with masking now needed only for tiles containing padding. In the final configuration (B3), 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.

Distributed Inference Considerations

Integrating the optimized sparse kernel into distributed inference requires addressing dynamic token layouts and different masks for different heads. The routing step returns a Boolean flag for each head to select spatial or temporal attention, choosing between original and permuted token layouts while keeping tensor dimensions fixed.

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, ensuring downstream layers receive expected ordering. Spatial heads retain their original ordering.

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 scattered diagonal stripes of the temporal mask into a single contiguous band that the tiled sparse kernel can traverse efficiently.

Moving token reordering after the input head exchange (rather than before) reduced total attention latency from 75.64 ms to 46.81 ms, a 38% reduction. This improvement includes routing, permutation, communication, and the sparse kernel itself.

Benchmark Results

The experiments were conducted on 720p video generation with 81 frames, 40 denoising steps, and eight TPU v6e chips. Three schedules were compared, retaining progressively fewer attention pairs, using the same prompt and seed with a local dense control for each configuration.

Schedule Denoising Time Speedup PSNR
Aggressive 153.50 s 1.28× 22% lower latency
Moderate - - -
Conservative - - -

At 1080p (171K tokens), the aggressive schedule reduced denoising time from 683 s to 457 s (1.49× speedup at 24.66 dB PSNR).

At 1440p (302K tokens, four times the sequence length of 720p), denoising time dropped from 2,471 s to 1,461 s - achieving a 1.69× end-to-end speedup and saving over 16 minutes per generated video while maintaining 24.05 dB PSNR.

Key Takeaway

Efficient sparse attention on TPUs depends on how much work the hardware can avoid, not just how many attention pairs the mask removes. The combination of skipping empty tiles, limiting intra-tile masking, and aligning the overall mask with the tiles the hardware computes yields substantial speedups across resolutions.

Read on Google Developers Blog ↗ ← Back to News

Comments

No comments yet. Start the discussion.