Architectures · Beyond attention
- 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
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 attentionA 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 MiBThe 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 htSelectivity
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)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))