IO-Aware Tiling: A New Principle for Attention Acceleration
Transformers are slow and memory-intensive on long sequences because the time and memory complexity of self-attention grows quadratically with sequence length.
EvidenciaE3 inspeccionableAnálisisEstándar
Transformers are slow and memory-intensive on long sequences because the time and memory complexity of self-attention grows quadratically with sequence length. Approximate attention methods attempt to trade model quality for computational savings, but often fail to deliver actual wall-clock speedups.
The paper proposes a missing principle: making attention algorithms IO-aware, i.e., considering reads and writes between different levels of GPU memory. FlashAttention uses tiling to reduce the number of read/write operations between GPU High Bandwidth Memory (HBM) and on-chip SRAM.
It claims to be an exact attention algorithm, requiring no approximation, with results consistent with standard algorithms. It also extends to an approximate version that is faster than all existing approximate methods.
It achieves 3x speedup on GPT-2 (sequence length 1K) and 2.4x speedup on Long Range Arena (sequence lengths 1K-4K). GPT-2 perplexity improves by 0.7, and long-document classification accuracy increases by 6.4 percentage points. It is the first method to enable Transformers to exceed random baseline performance on Path-X (16K sequence, 61.4%) and Path-256 (64K sequence, 63.1%).
Experiments were primarily conducted on GPUs, using specific models and sequence lengths. Beyond certain lengths, some approximate methods may outperform it. One cannot conclude from this that approximate methods are universally ineffective, nor extrapolate these findings to non-GPU hardware or untested configurations.
Suppose you are training a long-context model and notice low GPU utilization but saturated memory bandwidth—this is a typical manifestation of an IO bottleneck. In such cases, reducing data movement between HBM and SRAM is more effective than increasing compute power.