2026-08-31 15:57:13 +08:00

244 lines
7.6 KiB
Python
Executable File

#!/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()