28. Linear attention and Gated DeltaNet
In this chapter
- Why softmax attention's cost grows with context, and how removing the softmax turns attention into a recurrence with a fixed-size state.
- The state as an associative memory, why additive writes interfere, and how decay and the delta rule fix it.
- Gated DeltaNet: the recurrent form for decode and the chunked form, all matrix multiplications, for prefill.
- The full Qwen3.5 / Flash-Next linear-attention layer: short convolution, gates, and the state each request carries.
You will build
engine/gdn.py: the recurrent and chunked Gated DeltaNet, the causal convolution with carried state, and the complete layer, matching Hugging Face's Flash-Next layer.
Time: 6-8 hours. GPU: not needed.
The cost of remembering everything
Softmax attention keeps every past key and value. At decode, each new token reads the whole KV cache: time and memory per token grow linearly with context, and prefill grows quadratically. Chapter 16’s formula put Qwen3-0.6B’s cache at 3.5 GiB for a 32k context, three times the model. At the million-token contexts that agents and long documents want, the cache, not the weights, is the model.
Could a layer instead keep a fixed-size summary of the past, updated once per token, like an RNN? That’s linear attention. Flash-Next uses it in 36 of its 48 layers.
Removing the softmax
Attention for query $t$ is $o_t = \sum_{i \le t} \frac{\exp(q_t \cdot k_i)}{Z}, v_i$. The exponential couples $q_t$ and $k_i$ inside each term, so nothing can be precomputed. Replace $\exp(q \cdot k)$ by a plain dot product (or by $\phi(q) \cdot \phi(k)$ for some feature map $\phi$), and the sum factorizes (Katharopoulos et al., 2020, Transformers are RNNs):
$$ o_t = \sum_{i \le t} (q_t^\top k_i), v_i = q_t^\top \underbrace{\Big(\sum_{i \le t} k_i v_i^\top\Big)}_{S_t ,\in, \mathbb{R}^{d_k \times d_v}} . $$
$S_t$ is a $d_k \times d_v$ matrix, the state, and it updates with one outer product per token:
$$ S_t = S_{t-1} + k_t v_t^\top, \qquad o_t = S_t^\top q_t . $$
Per-token cost and memory are now constant, whatever the context. Prefill is linear in $T$ instead of quadratic.
The state is an associative memory
Think of $S$ as a memory that stores value $v$ under key $k$. Reading with a query equal to a stored key returns $S^\top k_j = \sum_i (k_i \cdot k_j), v_i$: the right value, if keys are orthonormal, plus interference from every other stored value in proportion to how similar its key is. A $d_k$-dimensional state can hold at most $d_k$ orthogonal keys. Writing more than that, and real sequences are thousands of tokens, the memories blur together.
python run.py linear
The first part stores $n$ random key-value pairs in a 64 × 64 state and reads each back with its key (unit-norm keys, relative error of the retrieved value):
{"pairs": 16, "key_dim": 64, "additive_error_all": 0.45, "delta_error_all": 0.303, "delta_error_last_8": 0.224}
{"pairs": 64, "key_dim": 64, "additive_error_all": 0.967, "delta_error_all": 0.714, "delta_error_last_8": 0.23}
{"pairs": 256, "key_dim": 64, "additive_error_all": 2.013, "delta_error_all": 1.199, "delta_error_last_8": 0.282}
Plain additive writes degrade steadily, and past 64 pairs the “retrieved” value is mostly noise. Two ideas fix this, and Gated DeltaNet uses both.
Forget: decay
Multiply the state by a decay $\alpha_t \in (0, 1)$ before each write: $S_t = \alpha_t S_{t-1} + k_t v_t^\top$. Old memories fade, making room for new ones. If $\alpha_t$ is computed from the input (a gate), the model can decide per token how much to forget: keep everything inside a sentence, wipe the slate at a document boundary. This is the idea behind RetNet, GLA and Mamba-2.
Correct: the delta rule
Instead of adding $v_t$ blindly, first ask what the memory currently returns for $k_t$, and write only the error:
$$ S_t = S_{t-1} + \beta_t, k_t \big(v_t - S_{t-1}^\top k_t\big)^\top . $$
This is the delta rule (Widrow-Hoff, 1960; DeltaNet, Schlag et al., 2021): one step of gradient descent on $\tfrac12 \lVert S^\top k_t - v_t \rVert^2$ with learning rate $\beta_t$. With $\beta = 1$ and a unit key, reading $k_t$ afterwards returns exactly $v_t$: the old association along $k_t$ is replaced, not piled on. That’s the “delta” columns above: the most recent pairs are retrieved well even at 256 pairs, and overall error is much lower.
Worked example (the milestone test). $k = q = [1, 0]$, $v = [2, 4]$, empty state. With $\beta = 0.5$: the prediction is $[0, 0]$, the error is $[2, 4]$, and $S$ gains $0.5 \cdot [1, 0]^\top [2, 4]$, so its first row is $[1, 2]$ and the output is $[1, 2]$. Repeat with $\beta = 1$: the prediction is now $[1, 2]$, the error $[1, 2]$, and the first row becomes $[2, 4]$, exactly $v$.
Gated DeltaNet
Combine both (Yang, Kautz and Hatamizadeh, 2024):
$$ S_t = \alpha_t S_{t-1} + \beta_t, k_t \big(v_t - \alpha_t S_{t-1}^\top k_t\big)^\top, \qquad o_t = S_t^\top q_t , $$
with $\alpha_t = e^{g_t}$ for a learned log-decay $g_t \le 0$, and $q, k$ L2-normalized (so that $\beta \le 1$ keeps updates stable).
The recurrent form
Decode processes one token at a time, so the recurrence is exactly what it needs:
def recurrent_gated_delta_rule(q, k, v, g, beta, state=None):
"""One token at a time. (Your engine: Chapter 28)
q, k [B, H, T, Dk] (already L2-normalized; q also scaled by Dk^-0.5); v [B, H, T, Dv];
g, beta [B, H, T]; state [B, H, Dk, Dv] or None. Returns (o [B, H, T, Dv], final state).
"""
b, h, t, dk = k.shape
dv = v.shape[-1]
S = torch.zeros(b, h, dk, dv, device=q.device, dtype=torch.float32) if state is None else state.float()
q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta))
out = torch.empty(b, h, t, dv, device=q.device, dtype=torch.float32)
for i in range(t):
S = S * g[:, :, i, None, None].exp()
prediction = torch.einsum("bhk,bhkv->bhv", k[:, :, i], S)
delta = (v[:, :, i] - prediction) * beta[:, :, i, None]
S = S + torch.einsum("bhk,bhv->bhkv", k[:, :, i], delta)
out[:, :, i] = torch.einsum("bhk,bhkv->bhv", q[:, :, i], S)
return out, S
// S is [dk][dv] row-major: decay, predict v from k, write beta * error along k, read with q.
inline std::vector<float> delta_step(std::vector<float>& S, const std::vector<float>& q, const std::vector<float>& k,
const std::vector<float>& v, float decay, float beta) {
size_t dv = v.size();
for (float& s : S) s *= decay;
std::vector<float> pred(dv, 0.f), out(dv, 0.f);
for (size_t i = 0; i < k.size(); ++i) for (size_t j = 0; j < dv; ++j) pred[j] += k[i] * S[i * dv + j];
for (size_t i = 0; i < k.size(); ++i) for (size_t j = 0; j < dv; ++j) S[i * dv + j] += beta * k[i] * (v[j] - pred[j]);
for (size_t i = 0; i < q.size(); ++i) for (size_t j = 0; j < dv; ++j) out[j] += q[i] * S[i * dv + j];
return out;
}
#![allow(unused)]
fn main() {
/// S is [dk][dv] stored row-major. Decay, predict v from k, write the scaled error, read with q.
pub fn delta_step(s: &mut [f32], q: &[f32], k: &[f32], v: &[f32], decay: f32, beta: f32) -> Vec<f32> {
let dv = v.len();
s.iter_mut().for_each(|x| *x *= decay);
let mut prediction = vec![0.0; dv];
for (i, ki) in k.iter().enumerate() {
for j in 0..dv {
prediction[j] += ki * s[i * dv + j];
}
}
for (i, ki) in k.iter().enumerate() {
for j in 0..dv {
s[i * dv + j] += beta * ki * (v[j] - prediction[j]);
}
}
let mut out = vec![0.0; dv];
for (i, qi) in q.iter().enumerate() {
for j in 0..dv {
out[j] += qi * s[i * dv + j];
}
}
out
}
}
Each step is a few $d_k \times d_v$ elementwise operations and mat-vecs per head: memory-bound on the state, but the state is small (16 KiB per head in FP32 for $d_k = d_v = 64$) and doesn’t grow. The state stays in FP32 even when the model runs in BF16, because it accumulates over thousands of steps (Chapter 13).
The chunked form: recurrence as matrix multiplication
For prefill, a loop of $T$ sequential steps wastes a GPU. The chunked algorithm processes the sequence in chunks of $C$ tokens (64 is typical). Within a chunk, it unrolls the recurrence algebraically into matrix products; across chunks, it carries the state.
Let $G_i = \sum_{j \le i} g_j$ be the cumulative log-decay within the chunk, and $D_{ij} = e^{G_i - G_j}$ for $j \le i$ (the decay from position $j$ to $i$). Unrolling shows that the values actually written, $U$ (each $v_i$ minus what the memory predicted for $k_i$, scaled by $\beta_i$), satisfy a unit lower-triangular linear system: each position’s write depends on the writes before it in the chunk. Solving that system (the “UT transform”) turns the sequential dependency into one triangular solve per chunk. Then the outputs and the next chunk’s state are matrix products:
$$ \begin{aligned} \big(I + \operatorname{strict_tril}(\beta_i, k_i \cdot k_j, D_{ij})\big), U &= \beta \odot V - \beta \odot e^{G} \odot (K S_0) \ O &= e^{G} \odot (Q S_0) + \operatorname{tril}!\big(QK^\top \odot D\big), U \ S_\text{end} &= e^{G_C} S_0 + \textstyle\sum_j e^{G_C - G_j}, k_j U_j^\top \end{aligned} $$
def chunk_gated_delta_rule(q, k, v, g, beta, state=None, chunk=64):
"""Same contract and result as the recurrent form, computed chunk by chunk. (Your engine: Chapter 28)
Inside a chunk, let G_i be the cumulative log-decay up to position i and
D_ij = exp(G_i - G_j) for j <= i (0 above the diagonal). Unrolling the recurrence shows the
values actually written, U, solve a unit-lower-triangular system:
(I + strict_tril(beta_i k_i.k_j D_ij)) U = beta * v - beta * exp(G) * (K S_0)
(the "UT transform"); the chunk's outputs and final state then follow from matmuls:
O = exp(G) * (Q S_0) + (tril(Q K^T * D)) U
S_end = exp(G_last) S_0 + sum_j exp(G_last - G_j) k_j U_j^T
"""
b, h, t, dk = k.shape
dv = v.shape[-1]
q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta))
pad = (-t) % chunk
if pad:
q, k, v = (F.pad(x, (0, 0, 0, pad)) for x in (q, k, v))
g, beta = (F.pad(x, (0, pad)) for x in (g, beta))
n = q.shape[2] // chunk
q, k, v = (x.view(b, h, n, chunk, -1) for x in (q, k, v))
g, beta = g.view(b, h, n, chunk), beta.view(b, h, n, chunk)
G = g.cumsum(-1) # [b,h,n,C]
lower = torch.ones(chunk, chunk, dtype=torch.bool, device=q.device).tril()
decay = (G[..., :, None] - G[..., None, :]).masked_fill(~lower, float("-inf")).exp() # D_ij
k_beta, v_beta = k * beta[..., None], v * beta[..., None]
system = (k_beta @ k.transpose(-1, -2)) * decay # beta_i k_i.k_j D_ij
system = system.tril(-1) + torch.eye(chunk, device=q.device) # I + strictly-lower part
# Solve once for both right-hand sides: the part that depends on S_0 and the part that does not.
w = torch.linalg.solve_triangular(system, k_beta * G[..., None].exp(), upper=False, unitriangular=True)
u = torch.linalg.solve_triangular(system, v_beta, upper=False, unitriangular=True)
attn = (q @ k.transpose(-1, -2)) * decay # causal, decayed Q K^T
S = torch.zeros(b, h, dk, dv, device=q.device) if state is None else state.float()
out = torch.empty_like(v)
for c in range(n):
U = u[:, :, c] - w[:, :, c] @ S # values actually written
out[:, :, c] = (q[:, :, c] * G[:, :, c, :, None].exp()) @ S + attn[:, :, c] @ U
tail = (G[:, :, c, -1:] - G[:, :, c]).exp() # exp(G_last - G_j)
S = S * G[:, :, c, -1, None, None].exp() + (k[:, :, c] * tail[..., None]).transpose(-1, -2) @ U
return out.view(b, h, n * chunk, dv)[:, :, :t], S
The milestone test runs chunk sizes 1, 5, 8 and 64 against the recurrent form on a sequence of 37 tokens (so the last chunk is partial) with a non-zero initial state. They all agree to $10^{-5}$.
Note
The first version of this function kept the wrong triangle of the system matrix, a transposition that’s easy to make when deriving it on paper. Outputs looked plausible, and the chunk-size-1 test even passed, because a 1 × 1 chunk has no off-diagonal part. Testing several chunk sizes, including ones that don’t divide the length, caught it. Test a fast algorithm against the slow one over every boundary case you can construct.
On a laptop CPU, for 4 heads and 1,024 tokens:
{"form": "recurrent", "tokens": 1024, "ms": 132.0}
{"form": "chunked", "tokens": 1024, "ms": 12.2}
{"max_difference": 1.043081283569336e-07, "state_bytes_per_head": 16384, "kv_bytes_per_head_at_1024_tokens": 262144}
Eleven times faster, numerically identical, and the state is 16× smaller than one head’s KV cache at only 1,024 tokens (in BF16). The production kernels (the flash-linear-attention library’s Triton kernels, which Qwen’s models use) fuse these steps per chunk and run at near-matmul speed.
The full layer
The Gated DeltaNet layer in Qwen3.5 and Flash-Next wraps the rule with projections, gates and a short convolution. Parameter names match the checkpoints:
class GatedDeltaNet(nn.Module):
"""A Qwen3.5 / Flash-Next linear-attention layer. (Your engine: Chapter 28)"""
def __init__(self, hidden, key_heads, value_heads, key_dim, value_dim, conv_kernel=4, eps=1e-6,
gate_activation="silu"):
super().__init__()
if value_heads % key_heads:
raise ValueError("value heads must be a multiple of key heads")
self.kh, self.vh, self.dk, self.dv = key_heads, value_heads, key_dim, value_dim
self.key_size, self.value_size = key_heads * key_dim, value_heads * value_dim
conv_dim = 2 * self.key_size + self.value_size
self.in_proj_qkv = nn.Linear(hidden, conv_dim, bias=False)
self.in_proj_z = nn.Linear(hidden, self.value_size, bias=False)
self.in_proj_b = nn.Linear(hidden, value_heads, bias=False)
self.in_proj_a = nn.Linear(hidden, value_heads, bias=False)
self.conv1d = nn.Conv1d(conv_dim, conv_dim, conv_kernel, groups=conv_dim, bias=False)
self.dt_bias = nn.Parameter(torch.ones(value_heads))
self.A_log = nn.Parameter(torch.log(torch.empty(value_heads).uniform_(0.01, 16)))
self.norm = RMSNormGated(value_dim, eps, gate_activation)
self.out_proj = nn.Linear(self.value_size, hidden, bias=False)
def forward(self, x, state=None, chunk=64):
"""x [B, T, D] -> (y [B, T, D], new LinearAttentionState). (Your engine: Chapter 28)"""
b, t, _ = x.shape
state = state or LinearAttentionState()
mixed, conv_state = causal_conv1d(self.in_proj_qkv(x).transpose(1, 2), self.conv1d.weight[:, 0], state.conv)
q, k, v = mixed.transpose(1, 2).split([self.key_size, self.key_size, self.value_size], dim=-1)
q = q.reshape(b, t, self.kh, self.dk).transpose(1, 2)
k = k.reshape(b, t, self.kh, self.dk).transpose(1, 2)
v = v.reshape(b, t, self.vh, self.dv).transpose(1, 2)
repeat = self.vh // self.kh # several value heads share a key head
q, k = q.repeat_interleave(repeat, 1), k.repeat_interleave(repeat, 1)
q = l2norm(q.float()) * self.dk ** -0.5
k = l2norm(k.float())
beta = torch.sigmoid(self.in_proj_b(x)).transpose(1, 2) # write strength in (0,1)
g = (-self.A_log.float().exp() * F.softplus(self.in_proj_a(x).float() + self.dt_bias)).transpose(1, 2)
rule = recurrent_gated_delta_rule if t == 1 else chunk_gated_delta_rule
kwargs = {} if t == 1 else {"chunk": chunk}
o, recurrent = rule(q, k, v, g, beta, state.recurrent, **kwargs)
z = self.in_proj_z(x).reshape(b, t, self.vh, self.dv)
o = self.norm(o.transpose(1, 2).to(x.dtype), z)
return self.out_proj(o.reshape(b, t, -1)), LinearAttentionState(conv_state, recurrent)
Step by step:
- Project $x$ to $q, k, v$ in one matrix (
in_proj_qkv), plus an output gate $z$ (in_proj_z), a write strength $b$ and a decay input $a$ (one scalar per value head each). - Short causal convolution over time on $q, k, v$ (kernel 4, one filter per channel, then SiLU). It lets each token mix in its three predecessors before the state sees it: cheap local context. Its own state is the last 3 inputs.
- Heads: there are more value heads than key heads (Flash-Next: 16 key heads, 48 value heads), and each key head is shared by several value heads, like grouped-query attention in reverse.
- Normalize: $q$ and $k$ L2-normalized, $q$ scaled by $d_k^{-1/2}$.
- Gates: $\beta = \sigma(b)$, and the log-decay $g = -e^{A_\text{log}} \cdot \operatorname{softplus}(a + \text{dt_bias})$, the same parameterization as Mamba’s $\Delta$: always negative, so $\alpha = e^{g} \in (0, 1)$.
- The rule: recurrent for a single token, chunked otherwise.
- Gated RMSNorm: normalize each head’s output, multiply by $\operatorname{SiLU}(z)$, then project back to the hidden size.
def causal_conv1d(x, weight, state=None, activation=True):
"""Depthwise causal convolution over time with carried history. (Your engine: Chapter 28)
x [B, C, T]; weight [C, K]; state [B, C, K-1] = the previous K-1 inputs (zeros at start).
Returns (y [B, C, T], new_state). Output t sees inputs t-K+1 .. t only.
"""
kernel = weight.shape[-1]
if state is None:
state = x.new_zeros(x.shape[0], x.shape[1], kernel - 1)
joined = torch.cat((state.to(x.dtype), x), dim=-1)
y = F.conv1d(joined, weight[:, None, :].to(x.dtype), groups=x.shape[1])
return (F.silu(y) if activation else y), joined[..., -(kernel - 1):]
class RMSNormGated(nn.Module):
"""RMSNorm(x) * weight * silu(z): normalize the read-out, then let a gate decide how much passes."""
def __init__(self, width, eps=1e-6, activation="silu"):
super().__init__()
self.weight = nn.Parameter(torch.ones(width))
self.eps, self.activation = eps, activation
def forward(self, x, z):
x32 = x.float()
normed = x32 * torch.rsqrt(x32.square().mean(-1, keepdim=True) + self.eps)
gate = F.silu(z.float()) if self.activation == "silu" else torch.sigmoid(z.float())
return (self.weight * normed.to(x.dtype) * gate).to(x.dtype)
The state a request carries
A Gated DeltaNet layer’s per-request state is the convolution history plus the recurrent matrices, fixed size whatever the context. For Flash-Next: 48 value heads × 128 × 128 FP32 values per layer, 3 MiB, times 36 linear layers, 108 MiB per sequence. Its 12 attention layers’ KV cache at 32k tokens is 0.75 GiB; at 262k tokens, 6 GiB. An all-attention model of the same shape would need four times as much.
The fixed state changes the engine in three ways:
- No truncation. A KV cache can be rolled back by forgetting positions (Chapter 26’s speculative decoding). A recurrent state can’t: the rejected tokens have been mixed into it. Engines snapshot the state before speculating and restore it on rejection (Chapter 30’s
HybridState). - Prefix caching needs snapshots too, stored at block boundaries: the state after a shared prefix, not a list of per-position entries.
- Batching is simple: every request’s state has the same shape.
Why hybrid?
Linear attention’s fixed memory is also its weakness. Tasks that need exact recall of an arbitrary earlier token (copy this ID, find the needle in the haystack) are hard for any fixed-size state once the context is long, as the recall demo showed. A few full-attention layers restore precise lookup, and models with a 3:1 ratio of linear to attention layers (Qwen3-Next, Qwen3.5, Flash-Next) match or beat pure attention models on long-context benchmarks at a fraction of the cache. Flash-Next goes one step further: its attention layers are themselves sparse (Chapter 29).
Build it
Engine milestone 28: Gated DeltaNet. Implement recurrent_gated_delta_rule, chunk_gated_delta_rule, causal_conv1d and GatedDeltaNet.forward in engine/gdn.py (the gated norm, the state class and the constructor are provided).
pytest tests/test_ch28_gdn.py
python run.py linear --impl engine
The tests check the worked example, chunked against recurrent for four chunk sizes with an initial state, that the convolution’s carried state makes split calls equal one call, that a layer prefilling 13 tokens and then decoding 8 matches one full forward, and parity with Hugging Face’s Flash-Next Qwen4ExpTextGatedDeltaNet layer with the same weights.
Stretch exercises
- ★ Extend the recall demo with decay: at $\alpha = 0.95$, how does error depend on how long ago a pair was written? Where:
experiments/ch28.py(create it), adapting the recall example inrun.py’scmd_linear. - ★★ Measure the chunked form’s time for chunk sizes 16, 32, 64, 128 and $T$ = 4,096. Explain the optimum in terms of the triangular solve’s cost ($O(C^2)$ per position) against the number of sequential chunk steps ($T/C$). Where:
experiments/ch28.py(create it), callingengine.gdn.chunk_gated_delta_rule. - ★★ Write a Triton kernel for the recurrent step at decode: one program per (sequence, head), the state tile in registers, all of decay, predict, correct and read fused. Where: new
engine/kernels/triton_gdn.py, selected byGatedDeltaNet.forwardinengine/gdn.py. - ★★★ Implement Mamba-2’s selective state-space recurrence ($S_t = \alpha_t S_{t-1} + k_t v_t^\top$ with scalar $\alpha_t$ per head) in the same chunked style, and show it’s Gated DeltaNet with $\beta$’s correction term removed. Where: add a selective-state-space helper beside
chunk_gated_delta_ruleinengine/gdn.py.
Check your understanding
- Why does replacing $\exp(q \cdot k)$ by $q \cdot k$ make attention a recurrence?
- Why do additive writes interfere, and what limits how many associations a $d_k \times d_v$ state can hold?
- What does the delta rule write instead of $v_t$, and why does $\beta = 1$ replace an association exactly?
- Why does the chunked form need a triangular solve?
- Why can’t a recurrent state be truncated like a KV cache, and what do engines do instead?
Going deeper
- Katharopoulos et al., Transformers are RNNs (2020); Schlag, Irie and Schmidhuber, Linear Transformers Are Secretly Fast Weight Programmers (2021); Yang et al., Gated Linear Attention (2023), Parallelizing Linear Transformers with the Delta Rule over Sequence Length (2024) and Gated Delta Networks (2024).
- Gu and Dao, Mamba (2023) and Dao and Gu, Transformers are SSMs (Mamba-2, 2024), for the state-space view of the same family.
- Songlin Yang’s blog series DeltaNet Explained (three parts), the clearest derivation of the chunked algorithm; the
flash-linear-attentionrepository for production kernels. - GPU Mode L20, L21 and L24 (Scan), the parallel-prefix algorithm behind Mamba’s selective scan; PMPP Chapter 11 (Scan: Kogge-Stone and Brent-Kung parallel prefix sums).
- The Qwen3-Next and Qwen3.5 model cards for the hybrid 3:1 layout.