= Linear Attention & State-Space Models

Softmax attention is powerful, but its cache grows with the whole prefix. Linear attention and
state-space models trade exact softmax weights for recurrent states, so decoding can keep a
fixed-size summary instead of every key and value. That trade is central to long-context LLMs:
most designs keep some full attention, then use linear or state-space layers where constant
per-token state is worth the approximation.

[#sec-feature-map]
== Attention as a feature-map kernel

Causal softmax attention at time stem:[t] is a normalized weighted average over previous
values. Linear attention replaces the exponential score by a positive feature-map dot product:

[latexmath#eq-linear-kernel]
++++
K(\vq_t, \vk_s) = \vphi(\vq_t)^\T \vphi(\vk_s) .
++++

The output is

[latexmath#eq-linear-attention]
++++
\vy_t = \frac{\sum_{s\le t} K(\vq_t,\vk_s)\vv_s}
{\sum_{s\le t} K(\vq_t,\vk_s)} .
++++

Because the kernel factorizes, move the query-dependent term outside the prefix sum:

[latexmath#eq-recurrent-state]
++++
\mS_t = \mS_{t-1} + \vphi(\vk_t)\vv_t^\T,
\qquad \vc_t = \vc_{t-1} + \vphi(\vk_t) .
++++

Then stem:[\vy_t = \vphi(\vq_t)^\T\mS_t / (\vphi(\vq_t)^\T\vc_t)]. The numerator state
stem:[\mS_t] stores key-value outer products; the vector stem:[\vc_t] stores the normalizer.
This is the whole trick from linear Transformers <<katharopoulos2020transformers>>. It changes
the computation from a growing attention matrix to a recurrent update, while preserving a
parallel form for training.

.Causal linear attention: parallel and recurrent forms
[source,python]
----
include::code/linear_ops.py[tag=linear-attention]
----

The chapter tests generate seeded stem:[\vq,\vk,\vv] arrays and assert that the parallel lower
triangle and recurrent state produce the same outputs to float64 precision. That test is the
minimum correctness check: if the order of the outer product or the normalizer is wrong, the two
forms disagree immediately.

The feature map is doing the approximation work. Softmax attention compares every query with
every key and then normalizes those exact scores. Linear attention chooses a representation in
which the score already looks like an inner product of transformed vectors. Positivity matters
because the denominator should behave like a sum of weights, not like a cancellation between
positive and negative terms. Different papers choose different maps, but the implementation
pattern is the same once stem:[\vphi] has been chosen.

[#sec-decoding-state]
== Constant-state decoding

During autoregressive decoding, softmax attention caches every previous stem:[\vk_s] and
stem:[\vv_s]. A linear-attention layer instead caches only stem:[\mS] and stem:[\vc]. For fixed
feature and value widths, each new token updates the same two arrays and emits one output. The
cache size is independent of the prefix length, so decoding is stem:[O(1)] memory per token for
that layer.

This is not free. The state is a compressed summary: once two different prefixes produce the
same stem:[\mS] and stem:[\vc], the layer cannot distinguish them later. Full attention keeps
all past keys and values and can revisit a rare token exactly. Linear layers win when the model
can use a lossy recurrent memory for much of the stack and reserve exact attention for layers
that need retrieval.

This distinction also explains why training can still be parallel. During training, the whole
sequence is known, so a scan or a masked matrix computation can form all prefix states at once.
During decoding, tokens arrive one at a time, and the recurrent form becomes the natural
implementation. A useful mental model is therefore not \"linear attention is an RNN instead of
attention,\" but \"the attention kernel was chosen so that its causal prefix sums are RNN
states.\" The tests exercise both views on the same tensors.

[#sec-delta-rule]
== Delta-rule memories

The delta rule makes the recurrent state an error-correcting associative memory. Let
stem:[\mS] map key features to values. Before writing token stem:[t], predict its value from
its key, compute the error, and write only the error:

[latexmath#eq-delta-rule]
++++
\mS_t = \mS_{t-1} + \beta_t\,\vphi(\vk_t)
(\vv_t - \vphi(\vk_t)^\T\mS_{t-1})^\T .
++++

If the memory already predicts the value, the update is small. If it is wrong, the update
corrects that key. Gated DeltaNet adds a forget gate before the correction:

[latexmath#eq-gated-delta]
++++
\mS_t = \gamma_t\mS_{t-1} + \beta_t\,\vphi(\vk_t)
(\vv_t - \vphi(\vk_t)^\T\gamma_t\mS_{t-1})^\T .
++++

The gate lets the model decide how fast old associations decay, while the delta term still
writes prediction error. Gated Delta Networks use this idea to improve Mamba-style sequence
models <<yang2024gated>>.

The delta rule differs from the plain linear-attention write in one important way. A plain write
adds stem:[\vphi(\vk_t)\vv_t^\T] every time, even if the state already maps that key to the
right value. The delta update first asks what the state would return, then writes the residual.
That makes repeated evidence for the same association converge instead of growing without
bound. The gated version adds controlled forgetting, so new evidence can overwrite stale
associations rather than merely adding to them.

.Delta and gated-delta updates
[source,python]
----
include::code/linear_ops.py[tag=delta]
----

[#sec-ssm-hybrids]
== State-space models and hybrids

A linear state-space model keeps a hidden state stem:[\vh_t]. Discretizing a continuous system
gives a recurrence of the form

[latexmath#eq-ssm]
++++
\vh_t = \bar{\mA}_t\vh_{t-1} + \bar{\mB}_t\vx_t,
\qquad \vy_t = \mC_t\vh_t .
++++

Mamba's selective SSM makes the discretization and projections input-dependent, so the model can
choose what to remember or forget as a function of the current token <<gu2023mamba>>. The tiny
NumPy version below uses a diagonal state: stem:[\bar{\mA}_t] is an exponential decay and
stem:[\bar{\mB}_t,\mC_t] vary with the input.

.A tiny selective state-space recurrence
[source,python]
----
include::code/linear_ops.py[tag=ssm]
----

Hybrid LLMs mix these mechanisms rather than declaring one winner. Jamba is a hybrid
Transformer-Mamba model <<lieber2024jamba>>; MiniMax-01 uses Lightning Attention in a foundation
model stack <<minimax2025minimax01>>; Kimi Linear proposes an efficient linear-attention
architecture <<zhang2025kimi>>. The common recipe is to interleave cheap recurrent layers with
occasional full-attention layers, keeping exact retrieval paths while reducing average cache and
attention cost.

The state-space view is broader than linear attention, but the engineering motivation is
similar: replace an ever-growing table of past activations with a state updated by a recurrence.
The selective parameters are what keep that recurrence from being a fixed filter applied to all
tokens. A token can ask for a slow decay, a fast decay, or a different input projection. That is
why Mamba-like layers are usually discussed as content-dependent sequence models rather than as
ordinary convolutions.

[NOTE,caption=In practice]
====
Use linear attention when long prefixes make exact attention too expensive and the task can
benefit from a compressed recurrent memory. Use full attention when exact copying, retrieval, or
cross-token comparison matters. Modern long-context stacks usually combine them: recurrent or
linear layers carry most tokens cheaply, while full-attention layers refresh global access.
Mamba and DeltaNet-style layers should be read as sequence-memory layers, not as drop-in exact
softmax replacements <<gu2023mamba>><<yang2024gated>>.
====

[.key-equations#key-equations]
.Key equations
****
[latexmath]
++++
K(\vq,\vk) = \vphi(\vq)^\T\vphi(\vk)
++++

[latexmath]
++++
\mS_t = \mS_{t-1} + \vphi(\vk_t)\vv_t^\T, \quad
\vc_t = \vc_{t-1} + \vphi(\vk_t)
++++

[latexmath]
++++
\vy_t = \frac{\vphi(\vq_t)^\T\mS_t}
{\vphi(\vq_t)^\T\vc_t}
++++

[latexmath]
++++
\mS_t = \mS_{t-1} + \beta_t\vphi(\vk_t)
(\vv_t - \vphi(\vk_t)^\T\mS_{t-1})^\T
++++

[latexmath]
++++
\vh_t = \bar{\mA}_t\vh_{t-1} + \bar{\mB}_t\vx_t, \quad
\vy_t = \mC_t\vh_t
++++
****

[.teach]
[#sec-teach]
== Teach it

*The one-sentence version.* Linear attention factorizes the attention score so the whole prefix
can be summarized by a recurrent key-value state.

*An analogy.* Softmax attention keeps every note you ever wrote; linear attention keeps a
running ledger. The ledger is compact and fast to update, but it cannot recover a note that was
summarized away.

*At the board.*

. Replace stem:[e^{\vq^\T\vk}] by stem:[\vphi(\vq)^\T\vphi(\vk)].
. Pull stem:[\vphi(\vq_t)] outside the prefix sum and define stem:[\mS_t] and stem:[\vc_t].
. Show the one-token decoding update: add one outer product, then read with the query.
. Contrast a plain write with the delta rule: write the prediction error, not the whole value.

*Misconceptions to address.* Linear attention is not exact softmax attention. Constant-state
decoding saves memory per layer, not necessarily all model memory. Mamba is a selective SSM,
not just attention with a different kernel.

*Check for understanding.* What information is lost when a layer keeps only stem:[\mS] and
stem:[\vc] instead of all past keys and values?

[#sec-exercises]
== Exercises

[#ex-linear-attention-kernel.exercise]
.★ Kernel replacement
====
State the condition a feature map must satisfy for <<eq-linear-attention>> to be a normalized
weighted average, and explain why positivity matters.
====

[#ex-linear-attention-recurrence.exercise]
.★★ Derive the recurrent form
====
Starting from <<eq-linear-attention>>, derive <<eq-recurrent-state>> and the readout
stem:[\vphi(\vq_t)^\T\mS_t / (\vphi(\vq_t)^\T\vc_t)].
====

[#ex-linear-attention-delta.exercise]
.★★ Error-correcting write
====
For one key feature stem:[\vk = (1,0)] and value stem:[\vv], show what repeated delta-rule
updates with stem:[\beta = 1/2] do to the prediction error.
====

[#ex-linear-attention-implement.exercise]
.★★★ Decode one token
====
Implement a one-token decoder update that receives stem:[\vq_t,\vk_t,\vv_t,\mS,\vc] and returns
the output and updated state. Compare a full prefix processed this way with
`recurrent_linear_attention`.
====

[bibliography]
[#sec-references]
== References

include::../../book/sources.adoc[tags=katharopoulos2020transformers;gu2023mamba;lieber2024jamba;minimax2025minimax01;yang2024gated;zhang2025kimi]
