← All Posts
From Scratch · C++ · Transformers · Attention Variants · Part 7 of 13

Linear Attention: Kernelized Efficient Attention

Introduction

All attention variants so far compute scores for every (query, key) pair, creating a $(T \times T)$ matrix. Even with sparsity patterns, this is expensive for very long sequences. Linear attention takes a fundamentally different approach: replace the softmax similarity with a positive kernel function and reorder operations using associativity of matrix multiplication. The result: computation reduces from $O(T^2 \cdot d)$ to $O(T \cdot d^2)$.

This post implements linear attention using a feature map $\phi = \text{elu}(x) + 1$ (positive kernel), and shows how reordering transforms the problem.

Mathematical Foundation

Standard attention (softmax):

output_i = sum_j( softmax(Q_i @ K_j^T / sqrt(d)) * V_j ) = sum_j( exp(Q_i @ K_j^T) / Z_i * V_j ) Complexity: O(T²d) — must compute T² scores.

Linear attention (with positive kernel $\phi$):

output_i = sum_j( phi(Q_i) · phi(K_j)^T * V_j ) = phi(Q_i) @ (sum_j( phi(K_j) * V_j )) = phi(Q_i) @ (phi(K)^T @ V) Reorder: compute (phi(K)^T @ V) once for all queries [d x d matrix] Then: Q-stream just multiplies by this summary [O(Td²)] Complexity: O(T·d²) — much better for d << T.

Feature Map: ELU + 1

A feature map $\phi: \mathbb{R}^d \to \mathbb{R}^m$ must be positive (non-negative everywhere) for the reordering to work:

phi(x) = elu(x) + 1 elu(x) = { x, if x > 0 { alpha*(e^x-1), if x <= 0 (alpha=1 by default) phi(x) >= 1 for all x, so it's always positive. This is from "Transformers are RNNs" (Katharopoulos et al., 2020).

We need a positive function so that the reordered matrix multiplication doesn't introduce sign ambiguities. ELU is smooth (differentiable) and simpler than exponential, making it tractable. Adding 1 ensures it's always at least 1, avoiding numerical issues.

LibTorch Implementation

struct LinearAttentionImpl : public torch::nn::Module { torch::nn::Linear wq{nullptr}, wk{nullptr}, wv{nullptr}; int64_t d_model, d_k; LinearAttentionImpl(int64_t d_model_) : d_model(d_model_), d_k(d_model_) { wq = register_module("wq", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wk = register_module("wk", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); wv = register_module("wv", torch::nn::Linear( torch::nn::LinearOptions(d_model, d_k).bias(false) ) ); } torch::Tensor phi(torch::Tensor x) { // Feature map: elu(x) + 1 (always positive) return torch::elu(x) + 1.0; } torch::Tensor forward(torch::Tensor x) { auto Q = wq(x); // (batch, T, d_k) auto K = wk(x); // (batch, T, d_k) auto V = wv(x); // (batch, T, d_k) // Apply positive feature map auto phi_Q = phi(Q); // (batch, T, d_k) auto phi_K = phi(K); // (batch, T, d_k) // Step 1: Compute phi(K)^T @ V [d_k x d_k matrix per batch] auto kv_summary = torch::matmul( phi_K.transpose(-2, -1), // (batch, d_k, T) V // (batch, T, d_k) ); // (batch, d_k, d_k) // Step 2: Compute phi(Q) @ (phi(K)^T @ V) auto output = torch::matmul( phi_Q, // (batch, T, d_k) kv_summary // (batch, d_k, d_k) ); // (batch, T, d_k) // Optional: normalize by phi(Q) @ phi(K)^T @ 1 (per-token norm) // For simplicity, we skip this here. Full impl includes normalization. return output; } }; TORCH_MODULE(LinearAttention);

The Key Reordering Trick

Naive (quadratic): output_i = sum_j( phi(Q_i) ⊙ phi(K_j) * V_j ) Compute all (i, j) pairs first. [O(T² d)] Efficient (linear): Let S = phi(K)^T @ V [shape (d_k, d_k)] output_i = phi(Q_i) @ S [shape (d_k,)] Compute S once: O(T·d²) Apply to all Q: O(T·d²) Total: O(T·d²) instead of O(T²·d) Reordering uses the associativity: (A @ B) @ C = A @ (B @ C)

Interactive Dimension Comparison

▶ Dimension Comparison: $(T \times T)$ vs $(d \times d)$

See why reordering saves computation: when $d \ll T$, $(d \times d) \ll (T \times T)$.

Complete Shape Trace

StepVariableShapeNotes
Inputx(batch, T, d_model)Tokens
ProjectsQ, K, V(batch, T, d_k)Linear projections
Feature mapphi_Q, phi_K(batch, T, d_k)elu(·) + 1
K transposephi_K^T(batch, d_k, T)For matmul
Summarykv_summary(batch, d_k, d_k)phi(K)^T @ V
Outputoutput(batch, T, d_k)phi(Q) @ kv_summary

Key Insights

Positive kernels enable reordering. Softmax is non-linear and cannot be reordered. By replacing it with a positive feature map, we unlock the associative property of matrix multiplication, turning $O(T^2)$ into $O(T)$ (in terms of sequence length).
Trade-off: expressiveness vs. efficiency. Linear attention trades some representational power (kernelized similarity instead of learned softmax) for massive speedup. In practice, on very long sequences, this is a worthwhile trade.
The summary matrix is the key. Thinking of the computation as building a $(d \times d)$ summary of the key-value stream and then querying it simplifies both implementation and understanding.

Other Linear Attention Variants

Linear attention is an active research area. Other popular kernels include:

  • ELU+1 (used here): smooth, stable, from Katharopoulos et al. 2020
  • Exp kernel: $\phi(x) = \exp(x/\sqrt{d})$ (approximates softmax)
  • Polynomial: $\phi(x) = (1 + x/\sqrt{d})^p$ (positive, efficient)

What Comes Next

We've now covered seven fundamental attention patterns: self-attention, softmax building blocks, causal masking, cross-attention, sliding-window, global, and linear attention. Next, we explore multi-head attention (MHA), which runs multiple heads in parallel, and specialized variants like multi-query (MQA) and grouped-query (GQA) attention.

Previous Global Attention