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