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