Transformer Based Linear Attention with Optimized GPU Kernel Implementation
Armin Gerami, Ramani Duraiswami
TL;DR
The paper targets the bottleneck of Softmax-based attention in Transformers, which scales as $O(N^2D)$. It introduces a Transformer-based Linear Attention (LA) with a forward-backward formulation and a normalization scheme, paired with a highly optimized CUDA kernel that achieves $O(ND^2)$ time and $O(ND)$ memory. Analytical gradients are derived for the LA head to enable efficient training, and extensive CUDA-level optimizations minimize data movement. Empirically, the approach yields $3.3x$ speedup and $3.6x$ memory reduction over state-of-the-art LA on standalone tests and enables end-to-end training of a 1.4B-parameter LLM with comparable expressivity to Regular Attention, demonstrating practical viability for long-context Transformers on edge devices.
Abstract
The original softmax-based attention mechanism (regular attention) in the extremely successful Transformer architecture computes attention between $N$ tokens, each embedded in a $D$-dimensional head, with a time complexity of $O(N^2D)$. Given the success of Transformers, improving their runtime during both training and inference is a popular research area. One such approach is the introduction of the linear attention (LA) mechanisms, which offers a linear time complexity of $O(ND^2)$ and have demonstrated comparable accuracy to regular attention. However, LA in practice lags behind its theoretical efficiency. We propose a novel method for LA's forward and backward passes, along with a highly-optimized CUDA implementation. Our approach outperforms the state-of-the-art by 3.3 times in speed and reduces memory consumption by 3.6 times. We validate these improvements in both single-layer and end-to-end settings by training a 1.4 billion parameter language model, which demonstrates similar expressivity to regular attention on major reasoning benchmarks.
