| """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 |
| from mamba3_ref import rotate_pairs |
|
|
|
|
| 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() |
|
|