mamba3 / ref /mamba3_ref.py
phanerozoic's picture
Mamba-3 MIMO chunked forward and recurrent decode, native CUDA
d5a6d53 verified
Raw
History Blame
4.45 kB
"""PyTorch reference for the Mamba-3 MIMO decode recurrence.
Carried state per head: angle (Na,), S (P, N), kprev (R, N), vprev (R, P).
angle <- angle + tanh(theta_t) dt_t pi
q_t, k_t <- rotate_pairs(q_t + q_bias, angle), rotate_pairs(k_t + k_bias, angle)
alpha <- exp(a_t dt_t), gamma <- sigmoid(trap_t) dt_t, beta <- (1 - sigmoid(trap_t)) dt_t alpha
S <- alpha S + beta (kprev^T vprev) + gamma (k_t^T v_t)
y_t <- fold_r(gate(S q_t^T))
The three-term update is the trapezoidal discretization; it reaches one step
back, which is why kprev and vprev are state.
"""
import math
import torch
import torch.nn.functional as F
def rotate_pairs(x, cos, sin):
"""Rotate adjacent pairs of x (..., N) by cos/sin (..., Na); pairs past Na are identity."""
pairs = x.reshape(*x.shape[:-1], -1, 2)
x0, x1 = pairs[..., 0], pairs[..., 1]
npair = x0.shape[-1]
if cos.shape[-1] < npair:
pad = npair - cos.shape[-1]
cos = F.pad(cos, (0, pad), value=1.0)
sin = F.pad(sin, (0, pad), value=0.0)
r0 = x0 * cos - x1 * sin
r1 = x0 * sin + x1 * cos
return torch.stack([r0, r1], dim=-1).reshape_as(x)
def mimo_step_ref(
q, k, v, adt, dt, trap, q_bias, k_bias, angles, mimo_v, mimo_o,
D=None, z=None, mimo_z=None, state=None,
fused_norm=False, outproj_norm_weight=None, outproj_norm_eps=1e-5,
):
"""Sequential Mamba-3 MIMO recurrence; returns (y (B,S,H,P), state).
q, k (B,S,R,Gqk,N); v, z (B,S,H,P); adt, dt, trap (B,H,S); q_bias, k_bias
(H,R,N); angles (B,S,H,Na) pre-tanh; mimo_* (H,R,P) elementwise; D (H,).
"""
B, S, R, Gqk, N = q.shape
H, P = v.shape[2], v.shape[3]
Na = angles.shape[-1]
dev = q.device
if Gqk != H: # grouped-query: replicate q/k across heads
rep = H // Gqk
q = q.repeat_interleave(rep, dim=3)
k = k.repeat_interleave(rep, dim=3)
# Rank expansion is elementwise in P, so it is a broadcast multiply, not a projection.
v_r = v.unsqueeze(2) * mimo_v.permute(1, 0, 2)[None, None] # (B,S,R,H,P)
z_r = None if z is None else z.unsqueeze(2) * mimo_z.permute(1, 0, 2)[None, None]
qb = q_bias.permute(1, 0, 2)[None] # (1,R,H,N)
kb = k_bias.permute(1, 0, 2)[None]
if state is None:
angle = torch.zeros((B, H, Na), dtype=torch.float32, device=dev)
S_st = torch.zeros((B, H, P, N), dtype=torch.float32, device=dev)
kprev = torch.zeros((B, H, R, N), dtype=q.dtype, device=dev)
vprev = torch.zeros((B, H, R, P), dtype=v.dtype, device=dev)
else:
angle, S_st, kprev, vprev = (t.clone() for t in state)
S_st = S_st.float()
ys = []
for t in range(S):
qt = (q[:, t] + qb).permute(0, 2, 1, 3) # (B,H,R,N)
kt = (k[:, t] + kb).permute(0, 2, 1, 3)
vt = v_r[:, t].permute(0, 2, 1, 3) # (B,H,R,P)
zt = None if z_r is None else z_r[:, t].permute(0, 2, 1, 3)
dt_t = dt[:, :, t] # (B,H)
angle = angle + torch.tanh(angles[:, t].float()) * dt_t.unsqueeze(-1) * math.pi
cos, sin = torch.cos(angle).unsqueeze(2), torch.sin(angle).unsqueeze(2)
q_rot = rotate_pairs(qt, cos, sin)
k_rot = rotate_pairs(kt, cos, sin)
lam = torch.sigmoid(trap[:, :, t].float())
alpha = torch.exp(adt[:, :, t].float())
beta = (1.0 - lam) * dt_t * alpha
gamma = lam * dt_t
prev_kv = torch.einsum("bhrd,bhrp->bhpd", kprev.float(), vprev.float())
curr_kv = torch.einsum("bhrd,bhrp->bhpd", k_rot.float(), vt.float())
S_st = (alpha[..., None, None] * S_st
+ beta[..., None, None] * prev_kv
+ gamma[..., None, None] * curr_kv)
out = torch.einsum("bhpd,bhrd->bhrp", S_st, q_rot.float())
if D is not None:
out = out + D[None, :, None, None].float() * vt.float()
if fused_norm:
out = out * torch.rsqrt(out.square().mean(-1, keepdim=True) + outproj_norm_eps)
if outproj_norm_weight is not None:
out = out * outproj_norm_weight[None, :, None, :].float()
out = out * F.silu(zt.float())
elif zt is not None:
out = out * F.silu(zt.float())
ys.append(torch.einsum("bhrp,hrp->bhp", out, mimo_o.float()))
kprev, vprev = k_rot, vt
return torch.stack(ys, dim=1), (angle, S_st, kprev, vprev)