Efficient Transformers
Sub-quadratic and IO-aware attention mechanisms that conquer the quadratic memory and computation bottleneck. They scale transformers to massive context lengths.
Why Does This Exist?
Standard self-attention in vanilla transformers computes compatibility across all pairs of sequence tokens:
Given an input sequence of length and hidden dimension , the product produces an explicit attention matrix. Both the computational time to calculate this matrix and the memory storage to materialize it scale quadratically as .
For short sequences (), this quadratic overhead is negligible. However, as language and multimodal models scale to long documents, codebases, and audio streams ( to tokens), becomes completely intractable: an sequence would require materializing a 4-billion-element matrix per attention head at every layer, triggering catastrophic GPU Out-Of-Memory (OOM) failures.
Efficient transformer architectures tackle this bottleneck through four distinct mathematical and systems paradigms: low-rank projection (Linformer), sparse hashing (Reformer), kernel feature linearization (Performer), and hardware IO-aware memory tiling (FlashAttention).
Think of It Like This
A town hall debate compared across four crowd-management strategies
Imagine a town hall meeting with 4,000 citizens where every citizen needs to hear and respond to everyone else's ideas:
In the vanilla setup, every citizen conducts a private one-on-one consultation with every other person in the room. This requires separate conversations. The building fills with noise, gridlock ensues, and nobody finishes before midnight.
To make the assembly viable, four distinct committee designs emerge:
First, Linformer (The Representative Delegation): rather than debating 4,000 citizens individually, the crowd elects 128 regional delegates. Every citizen speaks only to the 128 delegates, who aggregate the viewpoints. Total discussions drop from sixteen million to roughly half a million.
Second, Reformer (Interest Circles): citizens are sorted into topical breakout rooms using a quick personality questionnaire (Locality-Sensitive Hashing). Citizens debate only within their specific interest room, ignoring unrelated groups.
Third, Performer (The Summary Scoreboard): instead of pairwise chats, each speaker writes their views as numerical attributes on an open community scoreboard. Listeners read the aggregated column totals directly, bypassing pairwise interaction entirely through associative addition.
Fourth, FlashAttention (The High-Speed Micro-Conference): rather than changing the debate rules or taking approximations, the organizers realize the bottleneck was citizens walking back and forth to an off-site archive. Organizers divide attendees into small room batches (SRAM tiles), let them compute exact consensus in small fast sprints, and immediately record final tallies without ever cluttering the main hallway.
How It Actually Works
Four Paradigms for Breaking Quadratic Scaling
The literature on efficient transformers divides into mathematical approximation methods and exact hardware-aware kernel optimizations.
1. Linformer: Low-Rank Sequence Projection
Wang et al. (2020) demonstrated theoretically and empirically that the attention matrix has low intrinsic rank: its singular values decay exponentially. Linformer exploits this by projecting the sequence length of Keys and Values from down to a fixed constant dimension using learned linear projection matrices :
The reduced attention operation evaluates to:
The attention matrix shrinks from to . This achieves linear computational time and memory footprint with respect to sequence length .
2. Reformer: Locality-Sensitive Hashing (LSH)
Kitaev et al. (2020) observed that in the softmax calculation, tokens with large positive dot products dominate the output, while distant orthogonal tokens receive near-zero weights. Reformer uses Locality-Sensitive Hashing (LSH) via random spherical projections to assign query and key vectors into discrete hash buckets:
where is a random Gaussian projection matrix. Only queries and keys falling into the identical hash bucket are permitted to attend to one another. Sorting by bucket reduces complexity to . In addition, Reformer replaces standard residual layers with reversible RevNet layers, eliminating the need to cache intermediate activations for backpropagation.
3. Performer: Fast Attention via Positive Orthogonal Random Features (FAVOR+)
Choromanski et al. (2020) approximate the softmax attention kernel using positive random feature maps :
Because the kernel decomposes into an inner product, matrix multiplication associativity allows re-ordering the computation:
where and . By computing first ( operations), followed by multiplication with , the full matrix is never formed, delivering strict linear time.
4. FlashAttention: IO-Aware Tiling
Dao et al. (2022, 2023) showed that on modern GPUs, attention runtime is bounded by memory bandwidth (IO) rather than arithmetic compute throughput (FLOPS). Standard attention repeatedly reads and writes the intermediate matrix between slow GPU High-Bandwidth Memory (HBM, 80 GB at ~2 TB/s) and fast on-chip SRAM cache (approx. 20 MB per Streaming Multiprocessor at ~19 TB/s).
FlashAttention computes the exact, unapproximated softmax attention by tiling into blocks that fit entirely within SRAM. It uses the online softmax trick to incrementally update running maximums and normalizers across tiles:
By fusing softmax scaling and value multiplication into a single GPU CUDA kernel, FlashAttention avoids ever writing the matrix to HBM, cutting memory accesses from to and delivering 2-4x wall-clock speedups.
Worked Example
Compare the computational floating-point operations (FLOPs) of Vanilla Attention versus a Linear Associative Kernel (Performer) on a sequence of length , feature dimension , and feature projection dimension .
1. Vanilla Attention
-
Step A: Compute : Multiply by :
-
Step B: Multiply : Multiply Attention matrix by :
-
Total Vanilla FLOPs:
2. Linear Kernel Associative Attention
With , , and :
-
Step A: Compute : Multiply by :
-
Step B: Multiply : Multiply by :
-
Total Linear Kernel FLOPs:
Speedup Ratio:
By swapping evaluation order via kernel associativity, the operation requires 64 times fewer computations on an sequence.
Code
import numpy as np
def linformer_attention( Q: np.ndarray, K: np.ndarray, V: np.ndarray, E_proj: np.ndarray, F_proj: np.ndarray) -> np.ndarray: """Compute Linformer attention using low-rank projection matrices.""" # Q: (N, d), K: (N, d), V: (N, d) # E_proj, F_proj: (k, N) where k << N N, d = Q.shape k = E_proj.shape[0]
# Project Keys and Values from N -> k K_proj = np.dot(E_proj, K) # (k, d) V_proj = np.dot(F_proj, V) # (k, d)
# Compute reduced (N x k) attention scores scale = 1.0 / np.sqrt(d) scores = np.dot(Q, K_proj.T) * scale # (N, k)
# Softmax along projected key dimension k scores_shifted = scores - np.max(scores, axis=-1, keepdims=True) exp_scores = np.exp(scores_shifted) attn_weights = exp_scores / np.sum(exp_scores, axis=-1, keepdims=True)
# Output: (N, k) x (k, d) -> (N, d) output = np.dot(attn_weights, V_proj) return np.round(output, 4)
def linear_kernel_attention( Q_prime: np.ndarray, K_prime: np.ndarray, V: np.ndarray) -> np.ndarray: """Compute associative linear attention: Q' * (K'^T * V).""" # Q_prime: (N, m), K_prime: (N, m), V: (N, d) # Associative contraction: (m, N) x (N, d) -> (m, d) context_kv = np.dot(K_prime.T, V)
# Final projection: (N, m) x (m, d) -> (N, d) output = np.dot(Q_prime, context_kv)
# Normalization by row sums normalizer = np.dot(Q_prime, np.sum(K_prime, axis=0, keepdims=True).T) output = output / (normalizer + 1e-6) return np.round(output, 4)
# Demonstration: N = 4, d = 2, k = 2Q_demo = np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0], [0.5, 0.5]])K_demo = np.array([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0], [0.5, 0.5]])V_demo = np.array([[2.0, 1.0], [1.0, 2.0], [3.0, 3.0], [1.5, 1.5]])
# Down-projection matrix from N=4 to k=2E_demo = np.array([[0.5, 0.5, 0.0, 0.0], [0.0, 0.0, 0.5, 0.5]])F_demo = E_demo.copy()
out_linformer = linformer_attention(Q_demo, K_demo, V_demo, E_demo, F_demo)print(f"Linformer output shape: {out_linformer.shape}")# -> Linformer output shape: (4, 2)
print(f"First row output: {out_linformer[0].tolist()}")# -> First row output: [1.8385, 1.8385]Watch Out For
Assuming kernel or low-rank approximations preserve exact needle-in-a-haystack retrieval
A common failure mode is deploying low-rank (Linformer) or kernel-based (Performer) models for tasks requiring precise token retrieval across long documents (such as searching for a specific variable definition in a 50,000-token repository).
Because low-rank projections and random kernel maps smooth out high-frequency dot-product peaks, they blur distinct sharp retrieval spikes. Consequently, while they excel at language modeling perplexity and document classification, they perform poorly on fine-grained retrieval benchmarks. For applications requiring exact associative recall, exact IO-aware methods like FlashAttention maintain zero mathematical loss while delivering high throughput.
The Quick Version
- Vanilla self-attention scales quadratically (), causing compute and memory bottlenecks on long sequences.
- Linformer and Performer achieve linear complexity through low-rank sequence projection and associative random feature kernels.
- FlashAttention delivers exact mathematical attention with 2-4x wall-clock speedups by tiling operations inside fast GPU SRAM.