244 lines
7.6 KiB
Python
Executable File
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()
|