"""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))