| """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: |
| rep = H // Gqk |
| q = q.repeat_interleave(rep, dim=3) |
| k = k.repeat_interleave(rep, dim=3) |
|
|
| |
| v_r = v.unsqueeze(2) * mimo_v.permute(1, 0, 2)[None, None] |
| 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] |
| 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) |
| kt = (k[:, t] + kb).permute(0, 2, 1, 3) |
| vt = v_r[:, t].permute(0, 2, 1, 3) |
| zt = None if z_r is None else z_r[:, t].permute(0, 2, 1, 3) |
|
|
| dt_t = dt[:, :, t] |
| 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) |
|
|