mamba3 / benchmarks /bench_step.py
phanerozoic's picture
Sync v1 sources and card to main
2fe518c verified
Raw
History Blame
3.57 kB
"""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()