mamba3 / benchmarks /bench_chunk.py
phanerozoic's picture
Sync v1 sources and card to main
2fe518c verified
Raw
History Blame
1.99 kB
"""Chunk-size sweep. The state workspace scales as B*H*(S/C)*N*P."""
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))
import load_local as mamba3 # noqa: E402
def timeit(fn, iters=20, warmup=5):
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]
dev = "cuda"
B, H, G, P, N, R = 1, 32, 1, 64, 128, 4
Na = N // 2
print(f"{torch.cuda.get_device_name(0)} B={B} H={H} P={P} N={N} R={R}")
for S in (2048, 4096):
torch.manual_seed(0)
f = lambda *s: torch.randn(*s, device=dev)
dt = F.softplus(-3.0 + f(B, H, S))
adt = -F.softplus(f(B, H, S)).clamp(max=-1e-4) * dt
trap = torch.rand(B, H, S, device=dev) * 0.5
raw = torch.rand(B, S, H, Na, device=dev)
q, k, v, z = f(B, S, R, G, N), f(B, S, R, G, N), f(B, S, H, P), f(B, S, H, P)
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 = f(H)
print(f" S={S}")
base = None
for C in (16, 32, 64, 128):
try:
fn = lambda: mamba3.forward(q, k, v, qb, kb, mv, mo, raw, adt, dt, trap,
z=z, mimo_z=mz, D=D, chunk_size=C)
fn()
except RuntimeError as ex:
print(f" C={C:3d} unavailable: {str(ex).splitlines()[0][:70]}")
continue
torch.cuda.reset_peak_memory_stats()
t = timeit(fn)
pk = torch.cuda.max_memory_allocated() / 1e6
ws = B * H * (S // C) * N * P * 4 / 1e6
base = base or t
print(f" C={C:3d} CR={C*R:3d} {t:7.3f} ms {base/t:5.2f}x "
f"workspace {ws:6.0f} MB peak {pk:5.0f} MB {B*S/t:7.1f} ktok/s")