Jonah's

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:

y(i)=V(i)softmax((K(i))⊤q(i)).

Now, when you replace the softmax with a dot-product, via the associativity property, you can write

y(i)=V(i)((K(i))⊤q(i))y(i)=(V(i)(K(i))⊤)q(i)y(i)=(∑j=1iv(j)⊗k(j))q(i)y(i)=W(i)q(i)

where ⊗ denotes an outer product and

W(i)=W(i−1)+v(i)⊗k(i).

Now, applying this to softmax attention, the goal is to replace the softmax kernel κ with another kernel κ′, where

κ′(k,q)=ϕ(k)⊤ϕ(q).

Now, we can rewrite softmax attention as

y(i)=∑j=1iv(j)ϕ(k(j))⊤ϕ(q(i))∑j′=1iϕ(k(j′))·ϕ(q(i)),

and, rearranging a bit

y(i)=(∑j=1iv(j)ϕ(k(j))⊤)ϕ(q(i))(∑j′=1iϕ(k(j′)))·ϕ(q(i)).

And when we introduce W(i) for the numerator and z(i) for the denominator (as the paper does), we arrive at our fast weight programmer:

W(i)=W(i−1)+v(i)⊗ϕ(k(i))z(i)=z(i−1)+ϕ(k(i))y(i)=W(i)ϕ(q(i))z(i)·ϕ(q(i)).

So, this is the general fast weight programmer (FWP) idea. The paper then goes a step further. They note that once the sequence length L exceeds the key feature dimension ddot, 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 ddot 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 W(i−1) and retrieves the value v¯(i) currently paired with the key k(i). The model then combines this retrieved value with the new input value v(i) using a learned weight β(i). With these enhancements, we now have:

k(i),v(i),q(i)=Wkx(i),Wvx(i),Wqx(i)v¯(i)=W(i−1)ϕ(k(i))β(i)=σ(Wβx(i))vnew(i)=β(i)v(i)+(1−β(i))v¯(i)W(i)=W(i−1)+vnew(i)⊗ϕ(k(i))⏟write−v¯(i)⊗ϕ(k(i))⏟removeW(i)=W(i−1)+β(i)(v(i)−v¯(i))⊗ϕ(k(i))y(i)=W(i)ϕ(q(i)).

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, v¯(i), and the target value, v(i). In the formulation above, β(i) controls how strongly the model edits its memory. When β(i)=0, the model leaves the existing association alone, whereas when β(i)=1, 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 St for the fast-weight state. In this notation, the DeltaNet update from above becomes

St=St−1−βt(St−1kt−vt)kt⊤.

They then rewrite this same update as

St=St−1−vtoldkt⊤+vtnewkt⊤St=St−1−βt(St−1kt)kt⊤+βtvtkt⊤St=St−1(I−βtktkt⊤)+βtvtkt⊤.

The authors note that the term in parentheses can be viewed as a generalized Householder transformation. A standard Householder matrix takes the form

H=I−2uu⊤u⊤u,

which is the identity minus a rank-one outer product. Our transition matrix has this same general structure:

Pt=I−βtktkt⊤.

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

Ht=βtvtkt⊤,

so that our update is simply

St=St−1Pt+Ht.

Expanding the first three steps gives us

S1=S0P1+H1S2=(S0P1+H1)P2+H2S2=S0P1P2+H1P2+H2S3=(S0P1P2+H1P2+H2)P3+H3S3=S0P1P2P3+H1P2P3+H2P3+H3.

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 Pi 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:

P1P2⋯PC=I−W⊤K.

K is simply our chunk of keys, while W contains an adjusted key wi for each of them. As for what these adjusted keys are:

P1=I−β1k1k1⊤.

then, define

w1=β1k1,

then we can write this as

P1=I−w1k1⊤.

Now lets add a second transition:

P1P2=(I−w1k1⊤)(I−β2k2k2⊤)P1P2=I−w1k1⊤−β2k2k2⊤+β2w1(k1⊤k2)k2⊤P1P2=I−w1k1⊤−β2(k2−w1(k1⊤k2))⏟w2k2⊤P1P2=I−w1k1⊤−w2k2⊤.

So, w2 is not merely β2k2, but also adjusted by k1⊤k2, which tells us how much the second key overlaps with the first. This adjustment is how W keeps track of the fact that the transitions are necessarily ordered.

We can repeat the same process to give the general recurrence

wt=βt(kt−∑i=1t−1wi(ki⊤kt)),

and therefore

P1P2⋯Pt=I−∑i=1twiki⊤=I−W⊤K.

Before collecting these recurrences into a matrix, we can rearrange the recurrence above as

wt⊤=βtkt⊤−∑i<tβt(kt⊤ki)wi⊤wt⊤+∑i<tβt(kt⊤ki)wi⊤=βtkt⊤.

Now, the recurrence for the wt vectors can be collected into a triangular linear system. First let

D=diag(β)

and

L=tril(DKK⊤,−1).

Left-multiplying by D rescales each row of KK⊤ by that row's corresponding βt, while trilkeeps only interactions with earlier keys. With those coefficients collected in L, all the individual recurrences can be written together as (notice that the form of this expression maps directly to our rewritten recurrence formula above):

(I+L)W=DK.

Solving gives

W=(I+L)−1DK.

The paper calls

T=(I+L)−1D,

so that

W=TK.

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 αt∈(0,1). Our DeltaNet recurrence

St=St−1(I−βtktkt⊤)+βtvtkt⊤

becomes

St=St−1(αt(I−βtktkt⊤))+βtvtkt⊤.

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. αt (the Mamba-2 contribution) controls how much of the cache generally survives, while βt and the delta rule control the precise edit at the current key. When αt=1, we get regular DeltaNet, and as αt 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, αt 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 αt, giving each feature dimension (i.e. a channel) its own decay rate.

St=St−1Diag(αt)(I−βtktkt⊤)+βtvtkt⊤.

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 βt 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 βt with a channel-wise erase gate bt and a channel-wise write gate wt.
And so, letting Dt=Diag(αt), their new update is

St=St−1Dt(I−(bt⊙kt)kt⊤)+(wt⊙vt)kt⊤.

Another way to see the same update is to first decay the state and retrieve the old value through the gated key:

S―t=St−1Dt,v^t=S―t(bt⊙kt),

and then write the residual

St=S―t+(wt⊙vt−v^t)kt⊤.

And GDN-2 splitting things in this way actually seems quite intuitive. The erase gate bt uses one pattern of key features to decide what old content to remove, while the write gate wt 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 kt. 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 βt, 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 bt=βt1dk and wt=βt1dv, this collapses back to KDA. If we additionally collapse αt 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.