From Fast Weight Programmers to KDA
Introduction
The goal of this blog is to better understand the progression that led to KDA and, by building some intuition for each step along the way, think about where this family of recurrent attention might go next.
Linear Transformers are Secretly Fast Weight Programmers
The crux of this paper comes down to asking whether we can use "slow weights" (i.e. your actual, trainable model weights) to program "fast weights" (i.e. per-layer accumulated state derived from those slow weights and your inputs) while also side-stepping the softmax bottleneck (i.e. what linear attention does).
The paper derives its fast weight programmer in this way:
Now, when you replace the softmax with a dot-product, via the associativity property, you can write
where denotes an outer product and
Now, applying this to softmax attention, the goal is to replace the softmax kernel with another kernel , where
Now, we can rewrite softmax attention as
and, rearranging a bit
And when we introduce for the numerator and for the denominator (as the paper does), we arrive at our fast weight programmer:
So, this is the general fast weight programmer (FWP) idea. The paper then goes a step further. They note that once the sequence length exceeds the key feature dimension , the model can end up in an over-capacity regime. This is not to say that every token necessarily consumes one independent slot, but that bounds the number of mutually orthogonal keys that the memory can hold. And this makes sense in that information is both stored in, and retrieved from, matrices, and exact interference-free retrieval requires the keys to be orthogonal. Otherwise, a query can retrieve a mixture of multiple values. And once you're in this regime, the ideal memory should be able to interact with its contents and selectively determine which associations to remember and which to forget.
To do this, they introduce a delta rule. The FWP first accesses the current state of memory and retrieves the value currently paired with the key . The model then combines this retrieved value with the new input value using a learned weight . With these enhancements, we now have:
The delta (from the name DeltaNet) comes from updating your fast weights based on the difference (the delta) between the value you currently retrieve from memory, , and the target value, . In the formulation above, controls how strongly the model edits its memory. When , the model leaves the existing association alone, whereas when , it applies the full correction. With a properly normalized key and no interference from other associations, this exactly replaces the value associated with the current key, and more generally, it pushes the retrieved value toward the target. Also something of note in the paper, in order to push out the memory-capacity frontier, they want to project into a higher-dimensional space while also remaining nonnegative so that the resulting similarities can still act as valid attention weights.
Parallelizing Linear Transformers with the Delta Rule
The question that this paper asks is can the DeltaNet state update be more efficient? The straightforward implemenation of DeltaNet requires you to process tokens sequentially as each next state relies on the previous state. The python code (for a single user) would resemble:
def delta_net(x):
# x: [sequence_length, model_dimension]
q = W_q(x)
k = W_k(x)
v = W_v(x)
beta = torch.sigmoid(W_beta(x)).squeeze(-1)
S = torch.zeros(
v.shape[-1],
k.shape[-1],
dtype=x.dtype,
device=x.device,
)
outputs = []
for i in range(x.shape[0]):
v_old = S @ k[i]
correction = beta[i] * (v[i] - v_old)
S = S + torch.outer(correction, k[i])
outputs.append(S @ q[i])
return torch.stack(outputs)
This paper will reparameterize the state update rule, which will allow processing of tokens within chunks to be done largely in parallel using matrix multiplies. State will still propogate between chunks, but there are far fewer sequential steps.
This paper first simplifies the notation a bit, using for the fast-weight state. In this notation, the DeltaNet update from above becomes
They then rewrite this same update as
The authors note that the term in parentheses can be viewed as a generalized Householder transformation. A standard Householder matrix takes the form
which is the identity minus a rank-one outer product. Our transition matrix has this same general structure:
Why identifying this as a generalized Householder transformation is useful becomes clearer when we unroll the recurrence a few times. To keep things readable, (the authors and I) let
so that our update is simply
Expanding the first three steps gives us
Now, if we were to compute our output state using this formula, we'd have to store a lot of intermediate results and have many matrix multiplies, which does not seem ideal. And so this is where reframing our transition matrix becomes important. Since each has the generalized Householder form (an identity matrix minus a rank-one matrix) we can use the compact WY representation for products of Householder matrices to write:
is simply our chunk of keys, while contains an adjusted key for each of them. As for what these adjusted keys are:
then, define
then we can write this as
Now lets add a second transition:
So, is not merely , but also adjusted by , which tells us how much the second key overlaps with the first. This adjustment is how keeps track of the fact that the transitions are necessarily ordered.
We can repeat the same process to give the general recurrence
and therefore
Before collecting these recurrences into a matrix, we can rearrange the recurrence above as
Now, the recurrence for the vectors can be collected into a triangular linear system. First let
and
Left-multiplying by rescales each row of by that row's corresponding , while trilkeeps only interactions with earlier keys. With those coefficients collected in , all the individual recurrences can be written together as (notice that the form of this expression maps directly to our rewritten recurrence formula above):
Solving gives
The paper calls
so that
With that, the formulas above map fairly directly to PyTorch. The following leaves out batch and head dimensions and assumes that the sequence length is divisible by the chunk size:
def delta_rule(Q, K, V, beta, chunk_size):
"""
Q, K: [sequence_length, key_dimension]
V: [sequence_length, value_dimension]
beta: [sequence_length]
"""
sequence_length, key_dimension = K.shape
value_dimension = V.shape[-1]
num_chunks = sequence_length // chunk_size
Q = Q.reshape(num_chunks, chunk_size, key_dimension)
K = K.reshape(num_chunks, chunk_size, key_dimension)
V = V.reshape(num_chunks, chunk_size, value_dimension)
beta = beta.reshape(num_chunks, chunk_size)
# Fold D = diag(beta) into K when constructing L = tril(D K K^T, -1).
K_beta = K * beta.unsqueeze(-1)
# Compute (I + L)^-1 with the paper's vectorized forward substitution.
T_inverse = -torch.tril(
K_beta @ K.transpose(-1, -2),
diagonal=-1,
)
for row in range(1, chunk_size):
T_inverse[..., row, :row] = T_inverse[..., row, :row] + (
T_inverse[..., row, :, None] * T_inverse[..., :, :row]
).sum(dim=-2)
I = torch.eye(chunk_size, dtype=K.dtype, device=K.device)
T_inverse = T_inverse + I
# T = (I + L)^-1 D. Right-multiplication by diagonal D
# is equivalent to scaling column c by beta[c].
T = T_inverse * beta.unsqueeze(-2)
W = T @ K
U = T @ V
S = torch.zeros(
value_dimension,
key_dimension,
dtype=V.dtype,
device=V.device,
)
outputs = []
for chunk in range(num_chunks):
# U - W S^T
corrections = U[chunk] - W[chunk] @ S.transpose(-1, -2)
# Q S^T -- reads state produced by all previous chunks
inter_chunk = Q[chunk] @ S.transpose(-1, -2)
# tril(Q K^T) (U - W S^T)
attention = torch.tril(Q[chunk] @ K[chunk].transpose(-1, -2))
intra_chunk = attention @ corrections
outputs.append(inter_chunk + intra_chunk)
# S_next = S + (U - W S^T)^T K
S = S + corrections.transpose(-1, -2) @ K[chunk]
return torch.stack(outputs).reshape(sequence_length, value_dimension)
Gated Delta Networks
Gated DeltaNet, simply, adds a data-dependent decay factor . Our DeltaNet recurrence
becomes
This follow-up paper, again led by Songlin Yang, notes that DeltaNet has no good way to rapidly clear outdated information. The delta rule is quite good at making a precise edit, which is to say that, for a given key, it can retrieve the value currently associated with it and replace that value with a better one. However, it can only forget an association when it has a particular key to target and a replacement to write.
Mamba-2 provides the inspirational behavior, whereby it (Mamba-2) multiplies the old state by a data-dependent decay factor, which allows the model to clear its memory before adding the new write at full strength. Notably though, this decay is not precise in that all of the existing associations are weakened together, which could be find when you want to make a bunch of cache irrelevant at once.
Gated DeltaNet folds this idea into DeltaNet. (the Mamba-2 contribution) controls how much of the cache generally survives, while and the delta rule control the precise edit at the current key. When , we get regular DeltaNet, and as approaches zero, the old state is cleared while the new key-value association is still written.
Kimi Delta Attention
KDA then makes this forgetting more fine-grained. In Gated DeltaNet, is a single scalar per head at each time step, and so every feature dimension of the old state decays by the same amount. KDA instead predicts a head-dimension-sized vector , giving each feature dimension (i.e. a channel) its own decay rate.
We can now make the same change to our simplified delta_rule implementation above. The main difference is that alpha has shape [sequence_length, key_dimension], so all of the cumulative decay terms now carry that additional key-feature dimension:
def kda(Q, K, V, alpha, beta, chunk_size):
"""
Q, K: [sequence_length, key_dimension]
V: [sequence_length, value_dimension]
alpha: [sequence_length, key_dimension]
beta: [sequence_length]
"""
sequence_length, key_dimension = K.shape
value_dimension = V.shape[-1]
num_chunks = sequence_length // chunk_size
Q = Q.reshape(num_chunks, chunk_size, key_dimension)
K = K.reshape(num_chunks, chunk_size, key_dimension)
V = V.reshape(num_chunks, chunk_size, value_dimension)
alpha = alpha.reshape(num_chunks, chunk_size, key_dimension)
beta = beta.reshape(num_chunks, chunk_size)
# gamma_r: cumulative per-channel decay from the start of the chunk to r
gamma = alpha.cumprod(dim=-2)
# Gamma[r, s, d] = gamma[r, d] / gamma[s, d]
# Each pair of positions now has one decay value per key channel.
Gamma = gamma[:, :, None, :] / gamma[:, None, :, :]
K_beta = K * beta.unsqueeze(-1)
V_beta = V * beta.unsqueeze(-1)
# The same triangular system as DeltaNet, now with the key-key
# interactions scaled by their per-channel decay between s and r.
T = -torch.tril(
torch.einsum("nrd,nsd,nrsd->nrs", K_beta, K, Gamma),
diagonal=-1,
)
for row in range(1, chunk_size):
T[..., row, :row] = T[..., row, :row] + (
T[..., row, :, None] * T[..., :, :row]
).sum(dim=-2)
I = torch.eye(chunk_size, dtype=K.dtype, device=K.device)
T = T + I
# W absorbs the cumulative channel-wise decay; U is unchanged.
W = T @ (K_beta * gamma)
U = T @ V_beta
S = torch.zeros(
value_dimension,
key_dimension,
dtype=V.dtype,
device=V.device,
)
outputs = []
for chunk in range(num_chunks):
gamma_chunk = gamma[chunk]
gamma_end = gamma_chunk[-1]
corrections = U[chunk] - W[chunk] @ S.transpose(-1, -2)
# Previous-chunk state, decayed independently along each key channel.
inter_chunk = (Q[chunk] * gamma_chunk) @ S.transpose(-1, -2)
# Causal interactions within the chunk carry the same channel-wise decay.
attention = torch.tril(
torch.einsum(
"rd,sd,rsd->rs",
Q[chunk],
K[chunk],
Gamma[chunk],
)
)
intra_chunk = attention @ corrections
outputs.append(inter_chunk + intra_chunk)
# Decay the old state to the chunk boundary, then add the corrected writes.
decay_to_end = gamma_end / gamma_chunk
S = S * gamma_end.unsqueeze(0)
S = S + corrections.transpose(-1, -2) @ (K[chunk] * decay_to_end)
return torch.stack(outputs).reshape(sequence_length, value_dimension)
Gated DeltaNet-2
Ok, so KDA refines Gated DeltaNet by making the decay channel-wise. GDN-2 says that's all well and good, but is still performing two roles. On the key side, it determines how much of the current read should be erased, and on the value side, it determines how much of the new value should be written. There is no real reason that these two decisions need to be the same, and so GDN-2 separates them, and so replaces with a channel-wise erase gate and a channel-wise write gate .
And so, letting , their new update is
Another way to see the same update is to first decay the state and retrieve the old value through the gated key:
and then write the residual
And GDN-2 splitting things in this way actually seems quite intuitive. The erase gate uses one pattern of key features to decide what old content to remove, while the write gate uses a separate pattern of value features to decide what new content to store. The erase gate therefore changes which key channels are used to retrieve the old value, but the resulting correction is still stored at the original key . In this sense, GDN-2 changes what is erased without changing where the update is written. Earlier models tie both decisions to the same scalar , meaning that erasing a lot necessarily means writing a lot. GDN-2 can instead perform, for example, a strong erase with a weak write, or vice versa, while varying both decisions across their respective channels.
If we set and , this collapses back to KDA. If we additionally collapse to one scalar per head, we get regular Gated DeltaNet. Importantly as the paper notes, the erase term is still rank one, so the same general WY-style chunkwise formulation remains available even though erasing and writing are no longer tied to the same scalar.
Conclusion
The (ongoing) history of this family of recurrent attention seems to follow a pretty clear pattern. Each iteration takes some decision that was previously tied together and gives the model more granular control over it. I would expect this to continue. The update is still limited to a single rank-one edit at each token, the erase direction is still derived from the current key, and the diagonal decay can scale channels independently but cannot mix information between them. Future versions might therefore introduce multiple erase directions, higher-rank updates, or some structured form of cross-channel interaction. The challenge is that this additional expressivity cannot come at the expense of the whole reason these models are interesting and useful, i.e. whatever gets added still needs to be packed into hardware-friendly matrix multiplications.