opt: cumsum/chunked attention kernels, memory-flat CE, block checkpointing, v2 trainer (proven equivalent, 46 tests)
e488197 verified Download benchmarks/bench_attention.py from thefinalboss/fractus-cte: direct link, hf CLI and curl.
- Browser
- Download file 4.54 kB
-
https://huggingface.co/thefinalboss/fractus-cte/resolve/1be8a5a2b51e8bc52109e469403f0015115341ee/benchmarks/bench_attention.py
- Command line
-
hf download hf://thefinalboss/fractus-cte@1be8a5a2b51e8bc52109e469403f0015115341ee/benchmarks/bench_attention.py
-
curl -L -o bench_attention.py https://huggingface.co/thefinalboss/fractus-cte/resolve/1be8a5a2b51e8bc52109e469403f0015115341ee/benchmarks/bench_attention.py
4.54 kB
| """Micro-benchmark of the linear-attention core at true 1B shapes. | |
| 1B config per block: n_heads=20, d_head=64, n_levels=2, chunk C=128. | |
| The engine flattens (B, nlev, H) into G = B*nlev*H groups of (C, dH). | |
| Measures forward-only and fwd+bwd wall time + RSS delta, per implementation. | |
| Each (impl, B) cell runs in its own subprocess: the einsum reference | |
| allocates O(G*C*dH^2) and can hit the native commit limit at larger B — | |
| a segfault in one cell must not kill the whole table. | |
| Usage: py benchmarks/bench_attention.py [--impl auto] [--iters 10] | |
| """ | |
| import argparse | |
| import subprocess | |
| import sys | |
| from pathlib import Path | |
| import torch | |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) | |
| sys.path.insert(0, str(Path(__file__).resolve().parent)) | |
| from bench_utils import RssTracker, timed | |
| from fractus.nn.attention import FractalLinearAttention | |
| def make_inputs(G, C, dH, seed=0): | |
| torch.manual_seed(seed) | |
| q = torch.rand(G, C, dH) + 0.5 # elu+1 features are positive; keep magnitudes realistic | |
| k = torch.rand(G, C, dH) + 0.5 | |
| v = torch.randn(G, C, dH) | |
| return q, k, v | |
| def run_impl(attn, fn, G, C, dH, carry, iters): | |
| q, k, v = make_inputs(G, C, dH) | |
| S0 = torch.randn(G, dH, dH) * 0.01 if carry else None | |
| z0 = torch.randn(G, dH) * 0.01 if carry else None | |
| # forward only | |
| def fwd(): | |
| with torch.no_grad(): | |
| fn(attn, q, k, v, carry=(S0, z0) if carry else None) | |
| # fwd + bwd (grad wrt inputs) | |
| def fwdbwd(): | |
| qg = q.clone().requires_grad_(True) | |
| kg = k.clone().requires_grad_(True) | |
| vg = v.clone().requires_grad_(True) | |
| y = fn(attn, qg, kg, vg, carry=(S0, z0) if carry else None) | |
| loss = (y[0] if isinstance(y, tuple) else y).sum() | |
| loss.backward() | |
| t_fwd = timed(fwd, iters=iters) | |
| t_fwdbwd = timed(fwdbwd, iters=iters) | |
| with RssTracker() as rss: | |
| fwdbwd() | |
| rss_delta = rss.delta # computed by __exit__, must be read after the block | |
| return t_fwd, t_fwdbwd, rss_delta | |
| def _impl_table(): | |
| from fractus.nn import attention as A | |
| impls = {} | |
| if hasattr(A.FractalLinearAttention, "_linear_attention_causal_einsum"): | |
| impls["einsum"] = A.FractalLinearAttention._linear_attention_causal_einsum | |
| for name in ("cumsum", "chunked"): | |
| if hasattr(A.FractalLinearAttention, f"_linear_attention_causal_{name}"): | |
| impls[name] = getattr(A.FractalLinearAttention, f"_linear_attention_causal_{name}") | |
| if not impls: | |
| # pristine repo state: only the vectorized body exists -> treat it as einsum ref | |
| impls["einsum"] = A.FractalLinearAttention._linear_attention_causal_vectorized | |
| return impls | |
| def run_row(name, B, iters): | |
| """Single cell, executed inside the spawned subprocess.""" | |
| fn = _impl_table()[name] | |
| attn = FractalLinearAttention(d_model=1280, n_heads=20, d_head=64, n_levels=2) | |
| dH, C = 64, 128 | |
| G = B * 2 * 20 | |
| t_fwd, t_fb, rss = run_impl(attn, fn, G, C, dH, carry=True, iters=iters) | |
| print(f"{name:<10} {B:>2} {G:>9} {t_fwd*1e3:>9.1f} {t_fb*1e3:>11.1f} {rss/1e6:>9.0f}", flush=True) | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--impl", default="auto", choices=["auto", "einsum", "cumsum", "chunked"]) | |
| ap.add_argument("--iters", type=int, default=10) | |
| ap.add_argument("--batch-sizes", default="2,4,8") | |
| ap.add_argument("--row", nargs=2, metavar=("IMPL", "B"), help=argparse.SUPPRESS) | |
| args = ap.parse_args() | |
| if args.row: | |
| run_row(args.row[0], int(args.row[1]), args.iters) | |
| return | |
| impls = _impl_table() | |
| chosen = list(impls.items()) if args.impl == "auto" else [(args.impl, impls[args.impl])] | |
| batch_sizes = [int(b) for b in args.batch_sizes.split(",")] | |
| me = str(Path(__file__).resolve()) | |
| print(f"{'impl':<10} {'B':>2} {'G':>9} {'fwd ms':>9} {'fwd+bwd ms':>11} {'RSSdMB':>9}") | |
| for name, _ in chosen: | |
| for B in batch_sizes: | |
| r = subprocess.run([sys.executable, "-u", me, "--row", name, str(B), | |
| "--iters", str(args.iters)], | |
| capture_output=True, text=True) | |
| if r.returncode == 0 and r.stdout.strip(): | |
| print(r.stdout.strip().splitlines()[-1], flush=True) | |
| else: | |
| why = "segfault/crash" if r.returncode else "no output" | |
| print(f"{name:<10} {B:>2} {B * 2 * 20:>9} {'-':>9} {'-':>11} " | |
| f"{'-':>9} <- {why} (rc={r.returncode})", flush=True) | |
| if __name__ == "__main__": | |
| main() | |