Skip to content
AI360Xpert
Beta

Efficient Transformers

Sub-quadratic and IO-aware attention mechanisms that conquer the quadratic memory and computation bottleneck. They scale transformers to massive context lengths.

Efficient transformer architectures reduce quadratic complexity through low-rank projection, hashing, kernel linearization, and IO-aware SRAM tiling.
Efficient transformer architectures reduce quadratic complexity through low-rank projection, hashing, kernel linearization, and IO-aware SRAM tiling.

Why Does This Exist?

Standard self-attention in vanilla transformers computes compatibility across all pairs of sequence tokens:

Attention(Q,K,V)=softmax(QK⊤dk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{Q K^\top}{\sqrt{d_k}}\right) V

Given an input sequence of length NN and hidden dimension dd, the product QK⊤Q K^\top produces an explicit N×NN \times N attention matrix. Both the computational time to calculate this matrix and the memory storage to materialize it scale quadratically as O(N2)\mathcal{O}(N^2).

For short sequences (N=512N = 512), this quadratic overhead is negligible. However, as language and multimodal models scale to long documents, codebases, and audio streams (N=32,768N = 32{,}768 to N=1,000,000N = 1{,}000{,}000 tokens), O(N2)\mathcal{O}(N^2) becomes completely intractable: an N=64kN = 64\text{k} 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 4,000×4,000=16,000,0004{,}000 \times 4{,}000 = 16{,}000{,}000 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 N×NN \times N attention matrix P=softmax(QK⊤/d)P = \text{softmax}(Q K^\top / \sqrt{d}) has low intrinsic rank: its singular values decay exponentially. Linformer exploits this by projecting the sequence length of Keys and Values from NN down to a fixed constant dimension k≪Nk \ll N using learned linear projection matrices E,F∈Rk×NE, F \in \mathbb{R}^{k \times N}:

K~=EK∈Rk×d,V~=FV∈Rk×d\tilde{K} = E K \in \mathbb{R}^{k \times d}, \quad \tilde{V} = F V \in \mathbb{R}^{k \times d}

The reduced attention operation evaluates to:

LinformerAttention(Q,K,V)=softmax(QK~⊤d)V~\text{LinformerAttention}(Q, K, V) = \text{softmax}\left(\frac{Q \tilde{K}^\top}{\sqrt{d}}\right) \tilde{V}

The attention matrix QK~⊤Q \tilde{K}^\top shrinks from N×NN \times N to N×kN \times k. This achieves O(N⋅k)\mathcal{O}(N \cdot k) linear computational time and memory footprint with respect to sequence length NN.

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:

h(x)=arg⁡max⁡[[xR ; −xR]]h(x) = \arg\max \left[ [x R \,;\, -x R] \right]

where R∈Rd×b/2R \in \mathbb{R}^{d \times b/2} 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 O(Nlog⁡N)\mathcal{O}(N \log N). 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 ϕ:Rd→Rm\phi: \mathbb{R}^d \to \mathbb{R}^m:

softmax(q⊤kd)≈ϕ(q)⊤ϕ(k)\text{softmax}\left(\frac{q^\top k}{\sqrt{d}}\right) \approx \phi(q)^\top \phi(k)

Because the kernel decomposes into an inner product, matrix multiplication associativity allows re-ordering the computation:

Attention(Q,K,V)=(Q′K′⊤)V≡Q′(K′⊤V)\text{Attention}(Q, K, V) = (Q' K'^\top) V \equiv Q' (K'^\top V)

where Q′=ϕ(Q)∈RN×mQ' = \phi(Q) \in \mathbb{R}^{N \times m} and K′=ϕ(K)∈RN×mK' = \phi(K) \in \mathbb{R}^{N \times m}. By computing K′⊤V∈Rm×dK'^\top V \in \mathbb{R}^{m \times d} first (O(Nmd)\mathcal{O}(N m d) operations), followed by multiplication with Q′Q', the full N×NN \times N matrix is never formed, delivering strict O(N)\mathcal{O}(N) 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 N×NN \times N 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 Q,K,VQ, K, V into blocks that fit entirely within SRAM. It uses the online softmax trick to incrementally update running maximums and normalizers across tiles:

mnew=max⁡(mold,xi),dnew=doldemold−mnew+exi−mnewm_{\text{new}} = \max(m_{\text{old}}, x_i), \quad d_{\text{new}} = d_{\text{old}} e^{m_{\text{old}} - m_{\text{new}}} + e^{x_i - m_{\text{new}}}

By fusing softmax scaling and value multiplication into a single GPU CUDA kernel, FlashAttention avoids ever writing the N×NN \times N matrix to HBM, cutting memory accesses from O(N2)\mathcal{O}(N^2) to O(N)\mathcal{O}(N) 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 N=4096N = 4096, feature dimension d=64d = 64, and feature projection dimension m=64m = 64.

1. Vanilla Attention O(N2d)\mathcal{O}(N^2 d)

  • Step A: Compute QK⊤Q K^\top: Multiply Q∈R4096×64Q \in \mathbb{R}^{4096 \times 64} by K⊤∈R64×4096K^\top \in \mathbb{R}^{64 \times 4096}:

    FLOPs1=2×N×N×d=2×4096×4096×64=2,147,483,648≈2.15×109 FLOPs\text{FLOPs}_1 = 2 \times N \times N \times d = 2 \times 4096 \times 4096 \times 64 = 2{,}147{,}483{,}648 \approx 2.15 \times 10^9 \text{ FLOPs}
  • Step B: Multiply AVA V: Multiply Attention matrix A∈R4096×4096A \in \mathbb{R}^{4096 \times 4096} by V∈R4096×64V \in \mathbb{R}^{4096 \times 64}:

    FLOPs2=2×N×N×d=2,147,483,648≈2.15×109 FLOPs\text{FLOPs}_2 = 2 \times N \times N \times d = 2{,}147{,}483{,}648 \approx 2.15 \times 10^9 \text{ FLOPs}
  • Total Vanilla FLOPs:

    Totalvanilla≈4.29×109 FLOPs\text{Total}_{\text{vanilla}} \approx 4.29 \times 10^9 \text{ FLOPs}

2. Linear Kernel Associative Attention O(Nmd)\mathcal{O}(N m d)

With Q′∈R4096×64Q' \in \mathbb{R}^{4096 \times 64}, K′∈R4096×64K' \in \mathbb{R}^{4096 \times 64}, and V∈R4096×64V \in \mathbb{R}^{4096 \times 64}:

  • Step A: Compute M=K′⊤VM = K'^\top V: Multiply K′⊤∈R64×4096K'^\top \in \mathbb{R}^{64 \times 4096} by V∈R4096×64V \in \mathbb{R}^{4096 \times 64}:

    FLOPs1=2×m×N×d=2×64×4096×64=33,554,432≈3.36×107 FLOPs\text{FLOPs}_1 = 2 \times m \times N \times d = 2 \times 64 \times 4096 \times 64 = 33{,}554{,}432 \approx 3.36 \times 10^7 \text{ FLOPs}
  • Step B: Multiply Q′MQ' M: Multiply Q′∈R4096×64Q' \in \mathbb{R}^{4096 \times 64} by M∈R64×64M \in \mathbb{R}^{64 \times 64}:

    FLOPs2=2×N×m×d=2×4096×64×64=33,554,432≈3.36×107 FLOPs\text{FLOPs}_2 = 2 \times N \times m \times d = 2 \times 4096 \times 64 \times 64 = 33{,}554{,}432 \approx 3.36 \times 10^7 \text{ FLOPs}
  • Total Linear Kernel FLOPs:

    Totallinear=3.36×107+3.36×107=6.71×107 FLOPs\text{Total}_{\text{linear}} = 3.36 \times 10^7 + 3.36 \times 10^7 = 6.71 \times 10^7 \text{ FLOPs}

Speedup Ratio:

TotalvanillaTotallinear=4.29×1096.71×107≈64.0\frac{\text{Total}_{\text{vanilla}}}{\text{Total}_{\text{linear}}} = \frac{4.29 \times 10^9}{6.71 \times 10^7} \approx 64.0

By swapping evaluation order via kernel associativity, the operation requires 64 times fewer computations on an N=4096N = 4096 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 (O(N2)\mathcal{O}(N^2)), 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.