#!/usr/bin/env python3 """Kernel-focused SDPA benchmark using MiniMax-H3's traced tensor shapes.""" import argparse import contextlib import json import math import os import statistics import time from pathlib import Path import torch import torch.nn.functional as F from torch.nn.attention import SDPBackend, sdpa_kernel BACKENDS = { "flash": SDPBackend.FLASH_ATTENTION, "cudnn": SDPBackend.CUDNN_ATTENTION, "efficient": SDPBackend.EFFICIENT_ATTENTION, } def stats(values): values = sorted(values) n = len(values) return { "samples_ms": values, "min_ms": min(values), "median_ms": statistics.median(values), "mean_ms": statistics.mean(values), "p90_ms": values[min(n - 1, math.ceil(0.90 * n) - 1)], "max_ms": max(values), "stdev_ms": statistics.stdev(values) if n > 1 else 0.0, "cv": statistics.stdev(values) / statistics.mean(values) if n > 1 else 0.0, } def make_qkv(layout, batch, heads, sequence, head_dim, dtype, device): torch.manual_seed(1101) if layout == "model_strided": fused = torch.empty( (batch, sequence, 3, heads, head_dim), dtype=dtype, device=device ).normal_(mean=0.0, std=0.02) q = fused[:, :, 0].permute(0, 2, 1, 3) k = fused[:, :, 1].permute(0, 2, 1, 3) v = fused[:, :, 2].permute(0, 2, 1, 3) owner = fused elif layout == "contiguous": q = torch.empty( (batch, heads, sequence, head_dim), dtype=dtype, device=device ).normal_(mean=0.0, std=0.02) k = torch.empty_like(q).normal_(mean=0.0, std=0.02) v = torch.empty_like(q).normal_(mean=0.0, std=0.02) owner = (q, k, v) else: raise ValueError(layout) return q, k, v, owner def backend_context(name): if name == "auto": return contextlib.nullcontext() return sdpa_kernel(BACKENDS[name]) def call_attention(q, k, v, backend): with backend_context(backend): return F.scaled_dot_product_attention( q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False ) def event_time_ms(fn): start = torch.cuda.Event(enable_timing=True) end = torch.cuda.Event(enable_timing=True) start.record() value = fn() end.record() end.synchronize() return start.elapsed_time(end), value def benchmark_layout(args, layout, dtype): q, k, v, owner = make_qkv( layout, args.batch, args.heads, args.sequence, args.head_dim, dtype, "cuda", ) torch.cuda.synchronize() for _ in range(args.warmup): output = call_attention(q, k, v, args.backend) torch.cuda.synchronize() warm = [] for _ in range(args.iterations): elapsed, output = event_time_ms( lambda: call_attention(q, k, v, args.backend) ) warm.append(elapsed) # Touch a buffer larger than L2 before each measured launch. The attention # inputs themselves are much larger than L2 for both target shapes. thrash = torch.empty(args.cache_thrash_mib * 1024 * 1024 // 4, device="cuda") cold = [] for _ in range(args.cold_iterations): thrash.add_(1.0) torch.cuda.synchronize() elapsed, output = event_time_ms( lambda: call_attention(q, k, v, args.backend) ) cold.append(elapsed) graph_result = {"supported": False} try: graph = torch.cuda.CUDAGraph() torch.cuda.synchronize() with torch.cuda.graph(graph): graph_output = call_attention(q, k, v, args.backend) graph.replay() torch.cuda.synchronize() graph_times = [] for _ in range(args.iterations): elapsed, _ = event_time_ms(graph.replay) graph_times.append(elapsed) graph_result = {"supported": True, **stats(graph_times)} del graph_output, graph except Exception as exc: graph_result = {"supported": False, "error": repr(exc)} # One-call trace supplies the exact selected kernel and launch timeline. trace_path = Path(args.output).with_name( f"{Path(args.output).stem}-{layout}.trace.json" ) try: with torch.profiler.profile( activities=[ torch.profiler.ProfilerActivity.CPU, torch.profiler.ProfilerActivity.CUDA, ], record_shapes=True, ) as profiler: output = call_attention(q, k, v, args.backend) torch.cuda.synchronize() profiler.export_chrome_trace(str(trace_path)) trace_error = None except Exception as exc: trace_error = repr(exc) warm_stats = stats(warm) cold_stats = stats(cold) flops = 4 * args.batch * args.heads * args.sequence**2 * args.head_dim output_checksum = float(output.float().mean().item()) result = { "layout": layout, "dtype": str(dtype), "q_shape": list(q.shape), "q_stride": list(q.stride()), "k_stride": list(k.stride()), "v_stride": list(v.stride()), "q_contiguous": q.is_contiguous(), "warm": warm_stats, "cold": cold_stats, "cold_over_warm": cold_stats["median_ms"] / warm_stats["median_ms"], "effective_tflops": flops / (warm_stats["median_ms"] * 1e-3) / 1e12, "io_lower_bound_gbps": ( 4 * args.batch * args.heads * args.sequence * args.head_dim * dtype.itemsize / (warm_stats["median_ms"] * 1e-3) / 1e9 ), "cuda_graph": graph_result, "output_checksum": output_checksum, "trace": str(trace_path), "trace_error": trace_error, } del q, k, v, owner, output, thrash torch.cuda.empty_cache() return result def main(): parser = argparse.ArgumentParser() parser.add_argument("--scenario", required=True, choices=["F3", "RVA"]) parser.add_argument( "--backend", required=True, choices=["flash", "cudnn", "efficient", "auto"] ) parser.add_argument("--sequence", required=True, type=int) parser.add_argument("--output", required=True) parser.add_argument("--batch", type=int, default=1) parser.add_argument("--heads", type=int, default=28) parser.add_argument("--head-dim", type=int, default=128) parser.add_argument("--warmup", type=int, default=3) parser.add_argument("--iterations", type=int, default=12) parser.add_argument("--cold-iterations", type=int, default=5) parser.add_argument("--cache-thrash-mib", type=int, default=512) args = parser.parse_args() Path(args.output).parent.mkdir(parents=True, exist_ok=True) metadata = { "scenario": args.scenario, "backend": args.backend, "sequence": args.sequence, "batch": args.batch, "heads": args.heads, "head_dim": args.head_dim, "pid": os.getpid(), "torch": torch.__version__, "cuda": torch.version.cuda, "cudnn": torch.backends.cudnn.version(), "gpu": torch.cuda.get_device_name(), "capability": list(torch.cuda.get_device_capability()), "started_at": time.time(), "results": [], } try: for layout in ("model_strided", "contiguous"): metadata["results"].append( benchmark_layout(args, layout, torch.bfloat16) ) metadata["status"] = "success" except Exception as exc: metadata["status"] = "failed" metadata["error"] = repr(exc) raise finally: metadata["finished_at"] = time.time() Path(args.output).write_text( json.dumps(metadata, indent=2, ensure_ascii=False) + "\n" ) if __name__ == "__main__": main()