🤖 AI Summary
This work addresses the high computational cost of self-attention in high-resolution video generation and the inefficiency of existing training-agnostic sparse attention methods under multi-GPU sequence parallelism, which suffer from inter-head load imbalance. The authors propose the first sparse attention system supporting runtime load balancing, integrating Top-p/Top-k hybrid routing, video-aware block organization, peer-to-peer attention head migration, and slack-aware sparsity enhancement, along with compute-communication overlap to optimize throughput. Evaluated on the Wan2.2 I2V model, the approach reduces average load imbalance from 1.34 to 1.08, achieves a 4.41× speedup over FlashAttention in attention computation, and accelerates DiT inference by 2.02–2.11× while preserving generation quality.
📝 Abstract
Video Diffusion Transformers process long spatio-temporal sequences, making self-attention the main bottleneck in high-resolution video generation. Training-free sparse attention reduces this cost, but adaptive Top-$p$ routing creates uneven per-head workloads under multi-GPU sequence parallelism. The resulting workload heterogeneity turns sparse attention into a rank-level straggler problem. We present \method{}, a training-free sparse-attention system that improves the distributed execution efficiency of adaptive sparse attention under multi-GPU sequence parallelism. \method{} uses Top-$p$ routing, a Top-$k$ safety floor, and video-aware block organization as the sparse-routing frontend, then repairs the materialized mask at runtime. Runtime Load Balancing migrates a small number of heavy heads via P2P communication to shorten the current critical path. Slack-Aware Sparse Augmentation fills residual non-critical-rank slack with additional high-value blocks, while overlap hides scheduling and migration overhead behind existing computation. On step-distilled Wan2.2 I2V, \method{} reduces average load imbalance from 1.34 to 1.08 and delivers a $4.41\times$ attention speedup over FlashAttention, while achieving a $2.02$--$2.11\times$ DiT inference speedup with competitive video quality.