DART: Decoded Attention over Recurrent States

Modern large language models
Modern language models largely build on three families: Transformers, recurrent models, and hybrids.
Transformers represent history as token-level keys and values. At each step, a query compares against earlier keys and retrieves a weighted combination of their values. This gives the model explicit, content-dependent access to previous positions. The cost grows with context: full attention requires quadratic training compute, and autoregressive inference maintains a KV cache whose size grows linearly with sequence length.
Recurrent models, including state space models (SSMs) and linear attention, compress history into a compact state that is updated as new tokens arrive. Mamba and Mamba-2 use input-dependent state dynamics to control what is written, retained, and read. Linear attention also admits a recurrent formulation, accumulating key–value associations in a state. These formulations enable linear-time sequence processing and a fixed-size recurrent memory during inference. Compression, however, can make precise retrieval difficult, especially when the query is unknown at the time information is written. Zoology studies this recall gap between attention and efficient recurrent alternatives.
Hybrid models combine recurrent efficiency with attention-based retrieval. Architectures such as Jamba and Zamba interleave recurrent computation with attention layers, improving the balance between efficiency and model quality. These designs still maintain two forms of memory: token-level KV caches for attention and compact states for recurrence.
This motivates the central question of DART:
Can a unified memory representation support both recurrent compression and attention-style retrieval?
A unified memory representation
Mamba-2 provides a natural starting point through state space duality (SSD). SSD shows that a selective state space layer can also be written as structured causal attention: its coefficients \(A,B,C,X\) correspond to the attention roles of decay, keys, queries, and values. This connects recurrent computation and attention through the same underlying state representation.
The recurrent update and readout take the form
\[H_t = a_t H_{t-1} + B_t^\top X_t, \quad Y_t^{\mathrm{SSM}} = C_tH_t.\]The state \(H_t\in\mathbb{R}^{N\times P}\) can therefore be viewed as a compressed KV cache. The outer product \(B_t^\top X_t\) writes a token-dependent association into it, while \(a_t\) controls how the previous state carries forward. The readout \(C_tH_t\) uses the current token’s read vector to extract a value from this accumulated memory.
The key observation behind DART is that Mamba-2 reads values from the compressed KV cache, but does not decode keys from that cache for attention. Decoding keys as well would make it possible to compare a query against compressed recurrent memories and retrieve the relevant information from them.
SSD also provides an efficient chunkwise algorithm. Alongside the recurrent outputs, it computes each chunk’s contribution to the state. DART retains these contributions as chunk state memories. Each stored matrix represents what a particular chunk writes, giving retrieval access to individual compressed portions of the history.
This leads to state-memory attention (SMA). For the current token, \(C_t\) decodes a value and \(E_t\) decodes a key from each completed historical chunk memory. The query \(Q_t\) scores those keys, and attention combines the corresponding values. DART adds the retrieved readout to the native Mamba-2 output through a gated residual connection.
The two branches share the chunk state memories: recurrence compresses tokens into states, and attention retrieves from those states. The following sections explain how these memories are formed and how their keys and values are decoded.
From chunk contributions to stored memories
Mamba-2’s chunkwise computation divides a sequence into chunks of \(S\) tokens. Alongside local outputs, it computes each chunk’s contribution to the recurrent state. For chunk \(c\), ending at position \(e(c)\), that contribution is
\[\Delta H_{[c]} = \sum_{s\in[c]} \left(\prod_{k=s+1}^{e(c)} a_k\right) B_s^\top X_s.\]Each matrix \(\Delta H_{[c]}\in\mathbb{R}^{N\times P}\) summarizes the associations written within that chunk.
DART retains the individual contribution matrices as chunk state memories. The native recurrent state continues to evolve, while SMA gains access to a sequence of compressed historical memories.
Decoding keys and values for the current token
For a token at position \(t\), the block input \(U_t\) produces three vectors:
\[Q_t = U_t W_Q, \quad C_t = U_t W_C, \quad E_t = (U_t W_E)^\top.\]The query \(Q_t\in\mathbb{R}^{1\times N}\) determines which chunk memories receive attention. The native Mamba-2 read vector \(C_t\in\mathbb{R}^{1\times N}\) extracts values. The verifier \(E_t\in\mathbb{R}^{P\times 1}\) reads the value axis to construct a key for each memory.
For a historical chunk \(c\), the two decoded representations are
\[V_{t,c} = C_t\Delta H_{[c]} \in\mathbb{R}^{1\times P}, \quad K_{t,c} = \Delta H_{[c]}E_t \in\mathbb{R}^{N\times 1}.\]The stored object is a matrix; its decoded key and value depend on the current token. The same chunk memory can yield different representations at different positions. Both readouts are token-conditioned through \(C_t\) and \(E_t\), with a separate query \(Q_t\) used to score the decoded key.
The attention calculation then has a familiar form:
\[\ell_{t,c} = \frac{Q_t K_{t,c}}{\sqrt{N}}, \quad g_{t,c} = \operatorname{softmax}_{c<c(t)}(\ell_{t,c}), \quad R_t = \sum_{c<c(t)} g_{t,c}V_{t,c}.\]Here \(c(t)=\lceil t/S\rceil\) is the current chunk. For clarity, these equations omit normalization in the score computation. The implementation applies RMS normalization to \(Q_t\) and \(E_t\), and scales the state rows by their inverse root mean square when decoding keys.
The retrieved value enters as a gated correction to the native state readout:
\[G_t = \operatorname{SiLU}(U_t W_G), \quad Y_t = C_tH_t + G_tR_t.\]The gate \(G_t\) is scalar, and \(W_G\) is initialized to zero. The retrieval contribution therefore starts at zero and is learned during training. The surrounding Mamba-2 block processes the combined readout through its usual output operations.
Making retrieval efficient
DART reuses Mamba-2’s chunkwise computation to construct the memories shown in Figure 2. Efficient retrieval requires an additional step: a custom FlashAttention-style kernel for SMA.
The kernel processes a tile of queries while streaming through historical chunk matrices. It decodes their keys and values in on-chip memory, updates an online softmax, and accumulates the retrieved output. Decoded keys, values, logits, and attention weights remain temporary quantities. The forward pass writes the readout and log-sum-exp statistics to device memory; backward computation reconstructs the temporary quantities as needed.
This avoids materializing large token-by-chunk intermediate tensors. It reduces memory traffic while retaining the additional arithmetic required by SMA.
Cache size
Let the sequence length be \(L\) and the number of chunks be \(M=\lceil L/S\rceil\). A matched token-level attention head with key and value dimension \(P\) stores approximately \(2LP\) elements. DART’s historical chunk matrices require approximately \(MNP\) elements. Their leading, length-dependent cache ratio is
\[\frac{MNP}{2LP} \approx \frac{N}{2S}.\]At \(N=128\) and \(S=256\), the ratio is \(1/4\): 75% less length-dependent cache under this matched comparison. For example, at \(L=4096\) and \(P=64\), the chunk matrices contain 131,072 elements per head, compared with 524,288 elements for token keys and values.
This comparison assumes the same element precision and excludes fixed recurrent state, small normalization statistics, model weights, and other runtime allocations. It describes the dominant historical cache term. Total device-memory savings depend on the complete implementation and the attention baseline’s head-sharing configuration.
Compute cost
Each token can visit roughly \(L/S\) historical chunks, with key and value decoding costing \(O(NP)\) per visit. The added SMA computation is
\[O(LMNP)=O\!\left(\frac{L^2NP}{S}\right).\]For fixed \(N\) and \(S\), this is quadratic in sequence length, on top of Mamba-2’s linear-time state computation. Relative to token-level attention’s \(O(L^2P)\) arithmetic, the leading ratio is \(N/S\). Actual latency also depends on kernel utilization and memory movement; this ratio alone gives no wall-clock speedup guarantee.
Smaller chunks preserve more separately addressable memories and increase retrieval work. Larger chunks reduce cache and computation while compressing more information into each memory. Chunk size controls this trade-off.
Experiments
Associative recall
Multi-query associative recall, or MQAR, asks a model to remember several key–value pairs and predict the associated value when a key reappears. The paper compares DART and Mamba-2 at matched model widths and state sizes, alongside a Transformer++ baseline.
At \(N=64\), DART reaches near-perfect recall across all tested widths and sequence lengths. Even with \(N=16\), it exceeds Mamba-2 with \(N=64\) at matched widths in these experiments. The gains suggest that selective access to stored states can make compressed associations more useful.
The chunk size here is \(S=16\). The separate 75% cache-saving example uses \(S=256\) and \(N=128\); that percentage should not be assigned to this MQAR setting.
Language modeling and extraction
The language models are pretrained on the Pile for 100 billion tokens. DART adds fewer than 3% parameters to the corresponding Mamba-2 backbones. The main language-modeling configuration uses \(S=256\) and \(N=128\).
The average covers seven downstream accuracy metrics reported in the paper. General language-modeling quality stays close to the backbone, with modest average downstream gains. Most examples in these downstream benchmarks are shorter than 256 tokens, so many do not activate historical-chunk retrieval at this chunk size. This evaluation primarily checks that training with SMA preserves the recurrent branch’s general capability.
The changes are larger on information extraction. SWDE asks for attribute values from web pages; FDA asks for values associated with requested keys in regulatory documents. Both are scored with contains accuracy, which checks whether the generated response contains the expected answer.
These entries use DART’s main \(S=256\) setting. Synthetic needle-in-a-haystack tests also show substantial gains in several configurations: for example, on NIAH-Single-2 at 2,048 tokens, the 780M/795M comparison improves from 24.6% to 70.0%.
The results vary by task and context length. Many 4,096-token NIAH-Single-2 and Single-3 settings remain difficult. There are regressions as well: on NIAH-Single-1 at 4,096 tokens, DART 133M scores 57.2% versus Mamba-2 130M’s 74.2%. Full-attention Pythia remains stronger on several extraction and question-answering metrics, and Attamba is strong on several synthetic retrieval settings. The complete comparisons above show both the retrieval gains and the settings where other architectures remain stronger.
What the ablations reveal
Removing the SMA branch at evaluation time tests how much the trained model relies on retrieval from historical chunk states. It strongly reduces extraction accuracy and MQAR recall. The NIAH effects depend on the task and context length, so the SMA branch should not be read as a uniform improvement on every retrieval metric. This intervention uses a trained DART checkpoint; it is distinct from training a standalone Mamba-2 baseline.
Sharing the native SSM read vector \(C_t\) also matters. In the memory-constrained MQAR setting below, the shared design achieves high recall, while the otherwise matched variants with an independent SMA read vector remain at zero accuracy. This result supports sharing \(C_t\) for decoding historical memories in this tested setting.