From softmax attention to retention
Standard attention computes Attention(X) = softmax(QK^T / √d_k) V. The softmax does two jobs: it makes the weights non-negative and it normalizes each row to sum to 1. It is also what blocks a cheap recurrent form — because exp does not factor across positions, you cannot accumulate a running state, so decoding needs the full growing KV cache and costs O(N) per token.
Retention throws the softmax out. In its place it uses a deterministic decay by relative position: a token m steps in the past is weighted by γ^m for a fixed scalar 0 < γ < 1. The core operation becomes a plain bilinear form Q_n K_m^T scaled by γ^(n-m) and summed over the past. Because that decay does factor — γ^(n-m) = γ^n / γ^m — the same computation can be rolled into a constant-size recurrent state. Positional information rides along via an xPos/rotary-style complex rotation on Q and K, so retention encodes relative position in both a phase (the rotation) and a magnitude (the γ decay).
The recurrent form: a constant-size state
The recurrent form is where retention earns its inference story. Instead of keeping every past key and value, it keeps a single state matrix S_n of fixed shape [d_k, d_v] that summarizes the entire history:
State: S_n = γ · S_(n-1) + K_n^T V_n S_n : [d_k, d_v], S_0 = 0
Output: o_n = Q_n · S_n o_n : [1, d_v]
Q_n, K_n, V_n : [1, d] (row vectors for position n)
K_n^T V_n : [d_k, d_v] rank-1 outer-product updateEach step does one outer product to update the state and one vector–matrix product to read it out. The old state is simply faded by γ before the new token is added. Crucially, S_n never grows: its size is d_k × d_v regardless of how many tokens have gone by. So generation costs O(1) time and O(1) memory per token — no KV cache that swells with context. This is the RNN-style inference profile that Transformers cannot match.
The parallel form: a decay mask instead of softmax
The recurrent form is sequential, which is death for training throughput. So retention has an equivalent parallel form that processes the whole sequence at once, just like attention — but with the softmax replaced by an elementwise decay mask D:
Retention(X) = (Q K^T ⊙ D) V
D_(nm) = γ^(n-m) if n ≥ m (causal + exponential decay)
= 0 if n < m (no peeking at the future)
Q, K, V = X W_Q, X W_K, X W_V Q,K,V : [N, d]
QK^T : [N, N] ⊙ D (elementwise) → × V → [N, d_v]D is a single lower-triangular matrix that folds two things together: causal masking (zeros above the diagonal) and the decay (γ^(n-m) below it). There is no row-wise softmax — the weights are fixed by position, not learned per query, though a GroupNorm on the output plus a swish gate restores stable scale. This form is fully parallel across positions and maps onto the same dense matmul hardware attention already uses, so training runs at Transformer speed.
Why the two forms are the same computation
The magic is that the parallel and recurrent forms are not approximations of each other — they are algebraically identical. Unroll the recurrence from S_0 = 0:
S_n = Σ_(m=1..n) γ^(n-m) K_m^T V_m
o_n = Q_n S_n = Σ_(m=1..n) γ^(n-m) (Q_n K_m^T) V_m
= Σ_(m=1..n) (Q_n K_m^T · D_(nm)) V_mThat last line is precisely row n of (QK^T ⊙ D)V. The factorization γ^(n-m) = γ^n · γ^(-m) is what lets the double sum collapse into a single running state: the γ^n pulls out as the fade applied to the whole accumulated state each step. This is the same identity that powers linear attention; retention just adds the decay. The upshot is a rare luxury — you train with the parallel form and, with the same learned weights, serve with the recurrent form. No distillation, no re-training, no mismatch. One model, two schedules.
The chunkwise-recurrent form: parallel within, recurrent across
The parallel form is O(N^2) — fine for training on moderate context, painful for very long sequences. The pure recurrent form is O(N) but sequential, wasting the parallel hardware. The chunkwise-recurrent form is the hybrid: split the sequence into chunks of size B, run the parallel form inside each chunk, and carry a recurrent state between chunks.
For chunk i (inner positions j = 1..B):
Inner (parallel, in-chunk): I_i = (Q_i K_i^T ⊙ D) V_i
Cross (recurrent carry): C_i = (Q_i R_(i-1)) ⊙ ξ, ξ_j = γ^j
State update: R_i = γ^B R_(i-1) + K_i^T (V_i ⊙ ζ), ζ_j = γ^(B-j)
Retention(X_i) = I_i + C_iInside a chunk you pay the quadratic cost, but only O(B^2); across chunks you pay a single state carry. Total cost is O(NBd + Nd^2) — linear in sequence length for a fixed chunk size, while still doing most work as big parallel matmuls. This is the form you reach for to train or prefill on 100K-token context without the N^2 memory wall.
The decay factor &amp;amp;amp;gamma; and multi-scale retention
γ is the whole personality of a retention head. It sets an effective memory horizon: the weight on a token k steps back is γ^k, so influence decays geometrically. Work it numerically for γ = 0.9:
distance k : 0 1 5 10 20 50
γ^k (0.9) : 1.00 0.90 0.59 0.35 0.12 0.005A token 50 back contributes essentially nothing. A useful rule of thumb: the effective window is about 1/(1-γ) — roughly 10 tokens for γ=0.9, 100 for γ=0.99. A single fixed decay would be a straitjacket, so RetNet uses multi-scale retention (MSR): each head h gets its own γ_h, spread geometrically (roughly γ_h = 1 - 2^(-5-h)). Short-γ heads capture local detail; long-γ heads (near 1) hold long-range context. The decays are fixed, not learned — a deliberate simplification that keeps the three forms clean while covering many timescales at once.
Complexity and the impossible triangle
Retention’s pitch is that it grabs all three corners of a triangle usually considered to allow only two: parallel training, cheap O(1) inference, and strong performance. Transformers own the first and third but pay O(N) per decode step; classic RNNs own the second but cannot parallelize training; older linear-attention models got the first two but lagged on quality.
| Form | Time | Per-step decode | Used for |
|---|---|---|---|
| Parallel | O(N^2 d) | — | Training |
| Recurrent | O(N d^2) | O(1) time & memory | Inference / generation |
| Chunkwise | O(NBd + N d^2) | — | Long-sequence train/prefill |
| Transformer | O(N^2 d) | O(N) (KV cache grows) | (for contrast) |
The constant-memory decode is the standout: whether you are 1K or 100K tokens deep, the retention state is the same d_k × d_v matrix, so throughput and memory stay flat with context length — the opposite of a KV cache that grows without bound.