Summary (Overview)

  • This paper reverse-engineers the internal algorithm used by Mamba models to perform Associative Recall (AR) and Multi-Query Associative Recall (MQAR), identifying that Mamba implicitly learns linear hash functions via a similarity-preserving mechanism.
  • The authors develop a theoretical framework called "Recall Scaling Laws" that predicts the embedding dimension DD and state dimension NN required for perfect recall given vocabulary size VV and number of facts NfN_f.
  • The key theoretical result shows that high-probability recall requires model dimensions satisfying ND=Ω(Nflog⁡V)ND = \Omega(N_f \log V) — the state memory scales linearly in the number of facts and logarithmically in vocabulary size.
  • The framework extends to multi-layer models (where depth Λ\Lambda multiplies effective state size) and multi-head SSM patterns (where MHA improves recall but MQA/MKA degrade it).
  • Extensive empirical validation confirms the theoretical predictions across linear and full nonlinear Mamba models, with accuracy collapsing onto predicted one-dimensional scaling curves.

Introduction and Theoretical Foundation

Background

Transformers suffer from linear memory growth during decoding due to key-value caching. Recent architectures like Mamba [1], RWKV [2], and other linear RNNs use a fixed-size recurrent state, enabling constant memory complexity. However, this fixed-size state imposes inherent information compression limitations.

The paper addresses a core question: How does the fixed-size state of Mamba limit its recall capabilities, and what is the exact scaling relationship?

Theoretical Foundations

The analysis leverages two key theoretical tools:

  1. Johnson–Lindenstrauss (JL) Lemma: States that points in a high-dimensional space can be embedded into a lower-dimensional space while approximately preserving pairwise distances. Formally, for a set of VV points, an embedding into Rk\mathbb{R}^k with k=O(log⁡V/ε2)k = O(\log V / \varepsilon^2) preserves distances up to factor (1±ε)(1 \pm \varepsilon) with high probability.

  2. Mechanistic Interpretability: The approach of reverse-engineering neural networks by identifying specific "circuits" — compact computational units within the network that implement identifiable algorithms.

The MQAR Task

The Multi-Query Associative Recall task partitions vocabulary V\mathcal{V} into key vocabulary Vk\mathcal{V}_k and value vocabulary Vv\mathcal{V}_v, each of size V/2V/2. A prompt contains:

  • Context: NfN_f non-repeating key-value pairs (ki,vi)(k_i, v_i)
  • Query section: NfN_f query tokens (duplicates of context keys) plus padding

The model must retrieve the value v∗v^* corresponding to each query key k∗k^*.

Mamba Architecture

The Mamba block operates as (simplified):

x^=SiLU(Conv1D(Linear(xe))),z^=SiLU(Linear(xe)),y^=SSM(x^)⊙z^\hat{x} = \text{SiLU}(\text{Conv1D}(\text{Linear}(x^e))), \quad \hat{z} = \text{SiLU}(\text{Linear}(x^e)), \quad \hat{y} = \text{SSM}(\hat{x}) \odot \hat{z}

with the SSM update:

htd=Aˉtd⊙ht−1d+x^tdBˉtd,y^td=htd⋅Cth_t^d = \bar{A}_t^d \odot h_{t-1}^d + \hat{x}_t^d \bar{B}_t^d, \quad \hat{y}_t^d = h_t^d \cdot C_t

Methodology

Minimal Model and Simplified SSM

The authors construct a simplified single-layer Mamba with gating, discretization, nonlinearities, biases, and normalization removed. Without discretization and with A=IA = I, the SSM becomes:

ht=ht−1+x^tBt⊤,y^t=htCth_t = h_{t-1} + \hat{x}_t B_t^\top, \quad \hat{y}_t = h_t C_t

The Ladder of Theoretical Guarantees

The paper establishes four levels of recall guarantees:

LevelGuaranteeDimension RequirementType
1Exact recall (worst-case)N=D=VN = D = VNon-compressive (Thm. 3.1)
2Exact recall (worst-case)ND=O(Nflog⁡V)\sqrt{ND} = O(N_f \log V)Hash-based (Lem. 4.1)
3High-probability recallND=O(Nflog⁡V)ND = O(N_f \log V)Mean-case (Thm. 4.2)
4Lower bound (necessity)ND=Ω(Nflog⁡V)ND = \Omega(N_f \log V)Information-theoretic (Lem. 4.5)

Mechanistic Interpretability Validation

The authors define invariant operators that are robust to orthogonal transformations of weights:

E^in=(diag(Wconv0)PinE∣diag(Wconv1)PinE),Πv,out=E⊤Pout\hat{E}_{\text{in}} = \left(\text{diag}(W_{\text{conv}}^0) P_{\text{in}} E \mid \text{diag}(W_{\text{conv}}^1) P_{\text{in}} E\right), \quad \Pi_{v,\text{out}} = E^\top P_{\text{out}} Πk,in=SBE^in,Πq,in=SCE^in,Gvv=Πv,outΠv,in,Gkq=Πk,in⊤Πq,in\Pi_{k,\text{in}} = S_B \hat{E}_{\text{in}}, \quad \Pi_{q,\text{in}} = S_C \hat{E}_{\text{in}}, \quad G_{vv} = \Pi_{v,\text{out}} \Pi_{v,\text{in}}, \quad G_{kq} = \Pi_{k,\text{in}}^\top \Pi_{q,\text{in}}

The model output becomes:

yt=∑τ=0tGvvξτξτ⊤Gkqξty_t = \sum_{\tau=0}^{t} G_{vv} \xi_\tau \xi_\tau^\top G_{kq} \xi_t

where ξt=(xt−1xt)\xi_t = \binom{x_{t-1}}{x_t} are input token pairs.

Hidden State as Hash Table

The ideal non-compressive circuit stores facts as an outer product:

Ht≡Poutht=∑n=1Nfvnkn⊤,yt=HtqtH_t \equiv P_{\text{out}} h_t = \sum_{n=1}^{N_f} v_n k_n^\top, \quad y_t = H_t q_t

This is a V×VV \times V table where entries are 1 where facts exist. The compressive version uses:

Ht′=∑n=1Nfvn′kn′⊤,yt=E⊤∑n=1Nfvn′kn′⊤qt′H_t' = \sum_{n=1}^{N_f} v_n' k_n'^\top, \quad y_t = E^\top \sum_{n=1}^{N_f} v_n' k_n'^\top q_t'

with compressed tokens vn′=Evnv_n' = Ev_n, kn′=E~knk_n' = \tilde{E}k_n, qt′=E~qtq_t' = \tilde{E}q_t.


Empirical Validation / Results

Key Theoretical Results

Theorem 3.1 (Perfect non-compressive recall): A single-layer simplified Mamba with D=N=VD = N = V, expand=2\text{expand} = 2, Dconv=2D_{\text{conv}} = 2 perfectly solves MQAR (recall probability = 1).

Theorem 3.2 (Efficient compressive recall): With ND=O(Nflog⁡V)ND = O(N_f \log V), expand=2\text{expand} = 2, Dconv=2D_{\text{conv}} = 2, a single-layer Mamba solves the task with high probability.

Theorem 4.2 (Trained model recall scaling laws):

pAR≈Φ(NDNf−2log⁡V)p_{\text{AR}} \approx \Phi\left(\sqrt{\frac{ND}{N_f}} - \sqrt{2\log V}\right) pMQAR≈1L−2Nf∑t=2NfLΦ(ND12Nf+14t−2log⁡V)p_{\text{MQAR}} \approx \frac{1}{L - 2N_f} \sum_{t=2N_f}^{L} \Phi\left(\sqrt{\frac{ND}{\frac{1}{2}N_f + \frac{1}{4}t}} - \sqrt{2\log V}\right)

Remark 4.3 (Unified form):

precall≈Φ(NDaNf+N−2log⁡V)p_{\text{recall}} \approx \Phi\left(\sqrt{\frac{ND}{aN_f + N}} - \sqrt{2\log V}\right)

where a=1a = 1 for AR and a≈5/4a \approx 5/4 for MQAR.

Theorem 4.6 (Multi-layer):

precall≈Φ(ΛNDaNf+ΛN−2log⁡V)p_{\text{recall}} \approx \Phi\left(\sqrt{\frac{\Lambda ND}{aN_f + \Lambda N}} - \sqrt{2\log V}\right)

Theorem 4.7 (Multi-head): Given fixed D,ND, N:

pMQA=pMKA≤psingle=pMVA≤pMHAp_{\text{MQA}} = p_{\text{MKA}} \leq p_{\text{single}} = p_{\text{MVA}} \leq p_{\text{MHA}}

Empirical Findings

Figure 1 results — Accuracy grids show:

  • Theoretical predictions (column b) closely align with trained linear models (column c)
  • Full nonlinear models (column d) match linear models, confirming simplification preserves core recall behavior
  • All columns exhibit the predicted inverse DD–NN tradeoff

Figure 4 results — Scaling curves:

  • Accuracy collapses onto one-dimensional curves of the form precall(x)≈Φ(x−b)p_{\text{recall}}(x) \approx \Phi(x - b)
  • The scaling variable is x=1/σx = 1/\sigma where σ2=aNfND+1D\sigma^2 = \frac{aN_f}{ND} + \frac{1}{D}
  • b≈2log⁡Vb \approx \sqrt{2\log V}
  • Full nonlinear Mamba fits with afull≈12alineara_{\text{full}} \approx \frac{1}{2}a_{\text{linear}}, indicating improved performance under the same scaling law

Multi-layer results (Figure 5):

  • Recall accuracy depends only on effective state size Neff=ΛNN_{\text{eff}} = \Lambda N
  • Accuracy collapses onto a single curve independent of depth Λ\Lambda

Multi-head results (Figure 6):

  • MHA improves recall by increasing effective state size
  • MQA and MKA are weaker due to shared-value compression across heads
  • MVA matches the single-head baseline

Mechanistic Validation

  • Hidden state inversion (Figure 3): Projecting the hidden state back to vocabulary space reveals Ht′′≈HtH_t'' \approx H_t, confirming a compression-decompression scheme
  • Conv1D as copy-shift: Gvvξτ≈xτG_{vv}\xi_\tau \approx x_\tau and Gkq⊤ξτ≈ξτ−1G_{kq}^\top\xi_\tau \approx \xi_{\tau-1}, verified empirically
  • Ablation: Removing Conv1D causes complete failure of recall

Theoretical and Practical Implications

Theoretical Implications

  1. Optimality of scaling: The upper bound ND=O(Nflog⁡V)ND = O(N_f \log V) is matched by an information-theoretic lower bound ND=Ω(Nflog⁡V)ND = \Omega(N_f \log V), establishing that Mamba's recall capacity is essentially optimal for fixed-state recurrent models.

  2. Hash-table interpretation: The paper provides strong evidence that Mamba learns similarity-preserving linear hash functions, connecting neural network behavior to classical data structures.

  3. Multi-head design guidance: The result that MQA/MKA degrade recall while MHA improves it provides theoretical justification for architectural choices in Mamba-2.

Practical Implications

  • Model sizing: Given a target vocabulary size and number of facts, practitioners can predict required dimensions for reliable recall
  • Architecture selection: MVA is validated as the stronger multi-head design when controlling for parameter count, consistent with real-world perplexity results
  • Diagnostic tool: The scaling laws provide a principled way to evaluate whether a model is memory-limited or has other bottlenecks

Connection to Language Modeling

The trends in Thm. 4.7 are consistent with perplexity ablations for full Mamba-2 (Dao & Gu, 2024): MQA and MKA both yield worse perplexity than MVA, matching the finding that shared-value compression degrades recall. This suggests the analysis captures tradeoffs observed in real-world NLP tasks.


Conclusion

Main Takeaways

  1. Mamba performs associative recall by implicitly learning linear hash functions, verified through mechanistic interpretability
  2. The required state memory scales as ND=Θ(Nflog⁡V)ND = \Theta(N_f \log V) — linear in facts, logarithmic in vocabulary
  3. The theoretical framework accurately predicts recall probability across model configurations
  4. Multi-layer models benefit through effective state size ΛN\Lambda N; multi-head patterns have distinct recall characteristics

Future Directions

  • Extend analysis to other architectures: xLSTM, RWKV, DeltaNet
  • Investigate how architectural modifications impact recall capabilities
  • Characterize the theoretical role of the gating branch in recall (currently unclear)
  • Detailed investigation of why full models achieve afull≈12alineara_{\text{full}} \approx \frac{1}{2}a_{\text{linear}} (improved performance under same scaling law)

Limitations

  • Analysis relies on simplified models; the full contribution of each Mamba component (e.g., gating) is not fully characterized
  • The heuristic determination of afulla_{\text{full}} for nonlinear models requires further theoretical justification

Related papers