Architectures · Beyond attention

Mamba

Mamba-130m · a selective state-space model

compared
Mamba-130m vs GPT-2 small
layers
24 GPT-2 12
dmodel
768 GPT-2 768
mixer
scan · state 16 GPT-2 attention · 12 heads
memory
1.3 MiB, any length GPT-2 36 KiB per token
params
130M GPT-2 124M

All steps

  1. GPT-2 vs Mamba

    Mamba-130m is GPT-2 small’s size, but it has no attention. Each of its 24 blocks is one mixer: a widening projection, a short causal convolution, a selective state-space scan and a gate. Tokens exchange information only through a state carried from left to right. Click a label to jump.

    24 × 768 · 130M · no attention

  2. A running state instead of a cache

    To predict the next token, attention compares it with every earlier token and keeps all their keys and values, so the KV cache grows with every token. Mamba carries a fixed-size state instead: each new token costs the same, and past about 38 tokens its state is smaller than GPT-2’s cache.

    KV cache n × 36 KiB · state 1.3 MiB

  3. The scan

    The state is a small vector per channel (16 numbers in Mamba). At each token it is decayed by Ā and the input is written in through B̄; C reads the output. Here one toy channel with a 4-number state and the same step size Δ for every token, as in the earlier S4 models.

    ht = Ā ht−1 + B̄ xt · yt = C ht

  4. Selectivity

    Mamba makes the step size Δ, and B and C, depend on the current token. A large Δ writes the token in strongly and forgets more of the past; a small one lets it pass and keeps the memory. Right: real Δ from Mamba-130m; in its middle layers the names get the largest steps.

    Δt = softplus(W xt) · Ā_t = exp(Δt A)

  5. Same size, same job

    Real next-token guesses for the same prompt from GPT-2 small, with attention, and Mamba-130m, without it: two models of about the same size, trained on different web text. Both reach for places a cat might sit.

    same prompt, two architectures

Code

# the Mamba block: one mixer where GPT-2 has attention and an MLP
x, z = in_proj(rms_norm(u)).chunk(2, dim=-1)       # 768 → 2 × 1,536
x = F.silu(causal_conv1d(x))                       # depthwise, width 4
dt, B, C = x_proj(x).split([48, 16, 16], dim=-1)   # all three depend on the token
dt = F.softplus(dt_proj(dt))                       # Δ: a step size per channel
A = -torch.exp(A_log)                              # (1536, 16) learned decay rates
for t in range(T):                                 # the selective scan
    h = torch.exp(dt[t, :, None] * A) * h + dt[t, :, None] * B[t] * x[t, :, None]
    y[t] = (h * C[t]).sum(-1) + D * x[t]
u = u + out_proj(y * F.silu(z))

Go deeper