← All Posts
Deep Learning · Transformers· Recurrent and hybrid models

Kimi Delta Attention: Channel-Wise Memory Control

KDA lets different key-feature channels forget at different rates. It applies a diagonal decay to the memory, then uses the delta rule to fit the current association. The result is still a fixed-size recurrent state, with a more expressive transition than one scalar decay per head.

Change one piece of Gated DeltaNet

Keep $S_t\in\mathbb R^{d_k\times d_v}$, column keys and queries of width $d_k$, and values of width $d_v$. Replace the scalar $\alpha_t$ with a vector $\boldsymbol\alpha_t\in(0,1)^{d_k}$ and define $D_t=\operatorname{diag}(\boldsymbol\alpha_t)$.

$$\overline S_t=D_tS_{t-1},\qquad S_t=\overline S_t+\beta_tk_t(v_t-\overline S_t^\top k_t)^\top,\qquad o_t=S_t^\top q_t.$$

This is the recurrence introduced in Kimi Linear, expressed with the paper's key-by-value state orientation. The model architecture combines KDA with other layers; KDA is the token-mixing mechanism, not a synonym for the entire Kimi model family.

What a channel-wise decay actually scales

Left-multiplication by $D_t$ scales the rows of $S$, because those rows index key features. It does not select one previous token to erase, and it does not directly scale the value-coordinate columns. Several historical associations may overlap in the affected key-feature directions.

For a two-dimensional state $S=I$ and retention factors $(1,0.25)$, the decayed state becomes $\operatorname{diag}(1,0.25)$. A query along the first coordinate reads the old first value unchanged, while a query along the second sees one quarter of its old value. Scalar decay could not make those choices independently within this head.

Try the mechanism Change forgetting in two independent directions

The two operations generally do not commute. Read the error from the decayed state before writing it back. Enable JavaScript to change the example inputs; the complete calculation remains in the article.

Where diagonal-plus-low-rank structure appears

Expand the recurrence:

$$S_t=A_tS_{t-1}+B_t,\quad A_t=(I-\beta_tk_tk_t^\top)D_t=D_t-\beta_tk_t(k_t^\top D_t),\quad B_t=\beta_tk_tv_t^\top.$$

$A_t$ is a diagonal matrix plus a rank-one correction. This is the specialized diagonal-plus-low-rank structure that a chunkwise algorithm can exploit. It avoids treating every per-token transition as an arbitrary dense $d_k\times d_k$ matrix.

The phrase describes an algebraic structure, not a claim that every product of such transitions remains diagonal plus rank one. Products can acquire more complicated structure. Efficient implementations organize chunk computations around the actual recurrence rather than assuming the rank never grows.

The two transformations usually do not commute

With scalar decay, multiplying by $\alpha I$ commutes with the overwrite operator. With different channel rates, $(I-\beta kk^\top)D$ generally differs from $D(I-\beta kk^\top)$.

Choose $k=(1,1)/\sqrt2$, $\beta=1$, and $D=\operatorname{diag}(1,0.5)$. Then

$$ (I-kk^\top)D=\begin{pmatrix}0.5&-0.25\\-0.5&0.25\end{pmatrix},\qquad D(I-kk^\top)=\begin{pmatrix}0.5&-0.5\\-0.25&0.25\end{pmatrix}.$$

Swapping the order changes the state transition. In code, computing the prediction error from the decayed state makes the intended ordering clear. The lab checks the equivalence of the expanded and update forms, and checks that this swapped example differs.

The core update in a few lines

def kda_step(state, key, value, retention, beta, query):
    # state: [dk, dv], retention: [dk]
    decayed = retention[:, None] * state
    error = value - decayed.T @ key
    updated = decayed + beta * np.outer(key, error)
    return updated, updated.T @ query

For equal retention factors, this reduces to the scalar-gated update. With all factors one, it reduces to DeltaNet. Those limits are useful correctness tests because they compare independent descriptions of the same operation.

A complete layer contains more than this state equation

The surrounding layer needs projections that generate $q,k,v$, retention factors, and the write rate; normalization conventions; a way to incorporate short-range order; and an output transformation. The Kimi Linear design uses a hardware-oriented chunkwise algorithm and a layerwise hybrid with MLA. These implementation and architecture choices are part of what was evaluated in its report.

Later KDA-based models can modify gate ranges, output gates, kernels, and surrounding MLPs. The Kimi K3 case study names its specific changes rather than assuming the Kimi Linear configuration applies unchanged.

Expressive forgetting does not create a token list

The persistent matrix still has $d_kd_v$ entries per head at fixed dimensions. Channel-wise retention improves how that limited state is managed; it does not make it equivalent to retaining every historical key and value. Arbitrary retrieval quality must be measured, and a hybrid model's full-attention layers still have growing history storage.

Try it: Which multiplication would decay value-coordinate columns instead of key-feature rows?

Right-multiplying a key-by-value state by a diagonal matrix of value width. That is a different gating design. Here the retention vector has key width and multiplies rows on the left.

See where it belongs in a stack

Hybrid attention explains why recurrent KDA and token-addressable MLA can complement each other. Their memory costs must be added, not averaged away.