h3-blackwell-runtime/tools/smoke_attention_backends.py
2026-08-20 17:39:44 +07:00

58 lines
2 KiB
Python

"""Verify and time request-selectable H3 attention kernels on the active GPU."""
import argparse
import json
import time
import torch
from h3_blackwell_runtime.attention import run_attention
parser = argparse.ArgumentParser()
parser.add_argument("--backends", nargs="+", default=("sage2", "cudnn_sdpa", "ck_int8"))
parser.add_argument("--sequence", type=int, default=512)
parser.add_argument("--heads", type=int, default=56)
parser.add_argument("--head-dim", type=int, default=128)
parser.add_argument("--warmup", type=int, default=2)
parser.add_argument("--iterations", type=int, default=5)
parser.add_argument("--seed", type=int, default=440407)
args = parser.parse_args()
torch.manual_seed(args.seed)
q = torch.randn(1, args.heads, args.sequence, args.head_dim, device="cuda", dtype=torch.bfloat16)
results = {}
reference = None
with torch.inference_mode():
for backend in args.backends:
for _ in range(args.warmup):
output = run_attention(q, q, q, backend=backend, is_causal=False)
torch.cuda.synchronize()
elapsed = []
for _ in range(args.iterations):
started = time.perf_counter()
output = run_attention(q, q, q, backend=backend, is_causal=False)
torch.cuda.synchronize()
elapsed.append(time.perf_counter() - started)
if reference is None:
reference = output
delta = (output.float() - reference.float()).abs()
results[backend] = {
"mean_seconds": sum(elapsed) / len(elapsed),
"min_seconds": min(elapsed),
"finite": bool(torch.isfinite(output).all()),
"shape": list(output.shape),
"dtype": str(output.dtype),
"max_abs_vs_reference": delta.max().item(),
"mean_abs_vs_reference": delta.mean().item(),
}
print(json.dumps({
"gpu": torch.cuda.get_device_name(),
"torch": torch.__version__,
"cuda": torch.version.cuda,
"shape": list(q.shape),
"reference": args.backends[0],
"results": results,
}, indent=2))