"""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)