"""Decode-step latency: CUDA kernel against the eager PyTorch step.""" import math import sys from pathlib import Path import torch import torch.nn.functional as F ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(ROOT)) sys.path.insert(0, str(ROOT / "ref")) import load_local as mamba3 # noqa: E402 from mamba3_ref import rotate_pairs # noqa: E402 def eager_step(q, k, v, z, adt, dt, trap, qb, kb, ang, mv, mo, mz, D, st): angle, S_st, kprev, vprev = st B, R, G, N = q.shape H, P = v.shape[1], v.shape[2] if G != H: q = q.repeat_interleave(H // G, dim=2) k = k.repeat_interleave(H // G, dim=2) qt = (q + qb.permute(1, 0, 2)[None]).permute(0, 2, 1, 3) kt = (k + kb.permute(1, 0, 2)[None]).permute(0, 2, 1, 3) vt = v.unsqueeze(2).permute(0, 3, 2, 1)[..., 0].unsqueeze(2) if False else \ v.unsqueeze(1).permute(0, 2, 1, 3) * mv.permute(1, 0, 2)[None].permute(0, 2, 1, 3) zt = z.unsqueeze(1).permute(0, 2, 1, 3) * mz.permute(1, 0, 2)[None].permute(0, 2, 1, 3) angle = angle + torch.tanh(ang) * dt.unsqueeze(-1) * math.pi cos, sin = torch.cos(angle).unsqueeze(2), torch.sin(angle).unsqueeze(2) q_rot, k_rot = rotate_pairs(qt, cos, sin), rotate_pairs(kt, cos, sin) lam, alpha = torch.sigmoid(trap), torch.exp(adt) beta, gamma = (1 - lam) * dt * alpha, lam * dt prev_kv = torch.einsum("bhrd,bhrp->bhpd", kprev, vprev) curr_kv = torch.einsum("bhrd,bhrp->bhpd", k_rot, vt) S_new = (alpha[..., None, None] * S_st + beta[..., None, None] * prev_kv + gamma[..., None, None] * curr_kv) out = torch.einsum("bhpd,bhrd->bhrp", S_new, q_rot) + D[None, :, None, None] * vt return torch.einsum("bhrp,hrp->bhp", out * F.silu(zt), mo), (angle, S_new, k_rot, vt) def timeit(fn, iters=200, warmup=30): for _ in range(warmup): fn() torch.cuda.synchronize() s, e = torch.cuda.Event(True), torch.cuda.Event(True) ts = [] for _ in range(iters): s.record(); fn(); e.record(); torch.cuda.synchronize() ts.append(s.elapsed_time(e)) ts.sort() return ts[len(ts) // 2] def main(): dev = "cuda" H, P, N, R, Na, G = 32, 64, 128, 4, 32, 1 print(f"{torch.cuda.get_device_name(0)} H={H} P={P} N={N} R={R} (Mamba-3 default geometry)") print(f"{'batch':>6} {'kernel':>10} {'eager':>10} {'speedup':>8} {'state GB/s':>11}") for B in (1, 4, 16, 64, 256): torch.manual_seed(0) f = lambda *s: torch.randn(*s, device=dev) q, k = f(B, R, G, N), f(B, R, G, N) v, z = f(B, H, P), f(B, H, P) dt = F.softplus(-3.0 + f(B, H)) adt = -F.softplus(f(B, H)).clamp(max=-1e-4) * dt trap = torch.rand(B, H, device=dev) * 0.5 qb, kb = f(H, R, N), f(H, R, N) mv, mo, mz = (torch.rand(H, R, P, device=dev) / R for _ in range(3)) D, ang = f(H), torch.rand(B, H, Na, device=dev) state = mamba3.DecodeState(B, H, P, N, R, Na, device=dev) est = (torch.zeros(B, H, Na, device=dev), torch.zeros(B, H, P, N, device=dev), torch.zeros(B, H, R, N, device=dev), torch.zeros(B, H, R, P, device=dev)) kt = timeit(lambda: state.step(q, k, v, qb, kb, mv, mo, ang, adt, dt, trap, z=z, mimo_z=mz, D=D)) et = timeit(lambda: eager_step(q, k, v, z, adt, dt, trap, qb, kb, ang, mv, mo, mz, D, est)) gb = B * H * P * N * 4 * 2 / (kt * 1e-3) / 1e9 print(f"{B:>6} {kt:>9.4f}ms {et:>9.4f}ms {et/kt:>7.2f}x {gb:>10.1f}") if __name__ == "__main__": main()