Files
AstrAI/csrc/bench/benchmark_attention.py
T
ViperEkura 6709534d64 refactor: remove bf16 gemm and swiglu kernels and rebuild csrc benchmarks
- delete csrc/kernels/gemm.cu and swiglu.cu and drop their CMake and setup.py registration
- remove the ops wrappers plus backend/linear.py and backend/swiglu.py so Linear and MLP call F.linear directly
- drop the four gemm and swiglu kernel test files and prune the stale cuda_kernels.md sections
- add csrc/bench benchmarks for the remaining kernels: attention decode prefill paged decode paged prefill versus single-launch SDPA references, rotary versus the torch fallback, fp8 quantize and mm_fp8 versus torch baselines
- attention, rotary_emb, and fp8_ops kernels are unchanged
2026-09-05 01:38:10 +08:00

650 lines
20 KiB
Python

"""Benchmark the four attention kernels against single-launch torch SDPA.
Suites (--suite): decode, prefill, paged_decode, paged_prefill, all. The
torch side times one SDPA call per step over dense tensors: GQA expansion,
page-table gathers, padding, and masks are built once outside the timed
region, and masked calls prefer the cuDNN backend (the default masked path
is the slow math backend). The kernel's timed work still includes its fused
paged reads and current-token K/V append. Agreement is checked against the
same reference.
"""
from __future__ import annotations
import json
import math
import statistics
from dataclasses import asdict, dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable, Optional
import click
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
from astrai.extension import is_available
from astrai.extension.ops import (
attn_decode,
attn_paged_decode,
attn_paged_prefill,
attn_prefill,
)
from astrai.inference.workspace import MAX_SPLITS, Q_TILE_ROWS
@dataclass(frozen=True)
class GqaConfig:
"""One model family's attention geometry (llama-style GQA)."""
name: str
hq: int
hkv: int
head_dim: int
DEFAULT_CONFIGS = (
GqaConfig("llama2_7b", 32, 8, 128),
GqaConfig("llama3_70b", 64, 8, 128),
GqaConfig("qwen2_7b", 28, 4, 128),
GqaConfig("llama3_8b_d64", 32, 8, 64),
)
# (batch, per-request context length); context includes the token being
# decoded (kv_len = context, the last slot written in-kernel).
DECODE_CASES = ((1, 4096), (8, 4096), (32, 2048), (64, 1024))
# (batch, q_len); prefill from scratch so kv_len == q_len.
PREFILL_CASES = ((1, 2048), (1, 4096), (4, 1024), (8, 512))
def parse_config(value: str) -> GqaConfig:
parts = value.split(":")
if len(parts) != 4 or not parts[0]:
raise click.BadParameter("config must use NAME:HQ:HKV:HEAD_DIM")
try:
hq, hkv, head_dim = (int(item) for item in parts[1:])
except ValueError as exc:
raise click.BadParameter("HQ/HKV/HEAD_DIM must be integers") from exc
if hq <= 0 or hkv <= 0 or head_dim <= 0 or hq % hkv or head_dim % 32:
raise click.BadParameter(
"HQ/HKV positive with HQ % HKV == 0; HEAD_DIM % 32 == 0"
)
return GqaConfig(parts[0], hq, hkv, head_dim)
def time_operation(operation: Callable[[], torch.Tensor], iterations: int) -> float:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(iterations):
operation()
end.record()
end.synchronize()
return start.elapsed_time(end) / iterations
def summarize(values: list[float]) -> dict[str, float]:
ordered = sorted(values)
return {
"median_ms": statistics.median(ordered),
"p90_ms": ordered[max(0, math.ceil(0.9 * len(ordered)) - 1)],
}
def measure_operations(
operations: dict[str, Callable[[], torch.Tensor]],
*,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, list[float]]:
for operation in operations.values():
for _ in range(warmup):
operation()
torch.cuda.synchronize()
samples: dict[str, list[float]] = {name: [] for name in operations}
order = tuple(operations)
# A-B-B-A order balances cache, clock, and temperature drift.
for _ in range(trials):
for name in (*order, *reversed(order)):
samples[name].append(time_operation(operations[name], iterations))
return samples
def repeat_kv_heads(x: torch.Tensor, n_rep: int) -> torch.Tensor:
"""Expand [*, n_kv_heads, head_dim] to [*, n_kv_heads * n_rep, head_dim]
with the backend's grouping (kv head = q head // n_rep)."""
if n_rep == 1:
return x
n_heads, head_dim = x.shape[-2:]
return (
x.unsqueeze(-2)
.expand(*x.shape[:-2], n_heads, n_rep, head_dim)
.reshape(*x.shape[:-2], n_heads * n_rep, head_dim)
)
def sdpa(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, **kwargs) -> torch.Tensor:
"""SDPA over blhd tensors: [batch, seq, heads, head_dim] -> blhd."""
out = F.scaled_dot_product_attention(
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), **kwargs
)
return out.transpose(1, 2)
def prefer_cudnn_sdpa(
call: Callable[[], torch.Tensor],
) -> Callable[[], torch.Tensor]:
"""Return a closure running the masked SDPA ``call`` on cuDNN attention
when the backend accepts the bool mask, else torch's default (the
default masked path falls back to the much slower math backend)."""
def with_cudnn() -> torch.Tensor:
with sdpa_kernel([SDPBackend.CUDNN_ATTENTION]):
return call()
try:
with_cudnn()
except Exception:
return call
return with_cudnn
def ragged_lens(batch: int, span: int) -> list[int]:
"""Deterministic mixed lengths spanning [span // 2, span]."""
if batch == 1:
return [span]
step = max(span // 2 // (batch - 1), 1)
return [span - (batch - 1 - i) * step for i in range(batch)]
def cumsum_indptr(lens: list[int]) -> torch.Tensor:
return torch.tensor(
[0, *torch.tensor(lens).cumsum(0).tolist()], dtype=torch.int32, device="cuda"
)
@dataclass(frozen=True)
class PagedInputs:
"""Standalone replicas of the PagePool / InferenceWorkspace tensors."""
q: torch.Tensor
k_cache: torch.Tensor
v_cache: torch.Tensor
req_to_token: torch.Tensor
req_pool_indices: torch.Tensor
kv_indptr: torch.Tensor
def build_paged_inputs(
config: GqaConfig, kv_lens: list[int], q_lens: Optional[list[int]]
) -> PagedInputs:
"""Flat pool + page table: request ``i`` owns the contiguous slot range
``[offset_i, offset_i + kv_len_i)``. Q is packed across requests when
``q_lens`` is given (ragged prefill), else [B, Hq, D] (decode)."""
batch = len(kv_lens)
pool = torch.randn(
sum(kv_lens),
config.hkv,
config.head_dim,
device="cuda",
dtype=torch.bfloat16,
)
req_to_token = torch.zeros(batch, max(kv_lens), dtype=torch.int32, device="cuda")
offset = 0
for i, length in enumerate(kv_lens):
req_to_token[i, :length] = torch.arange(
offset, offset + length, dtype=torch.int32, device="cuda"
)
offset += length
return PagedInputs(
q=torch.randn(
sum(q_lens) if q_lens is not None else batch,
config.hq,
config.head_dim,
device="cuda",
dtype=torch.bfloat16,
),
k_cache=pool,
v_cache=torch.randn_like(pool),
req_to_token=req_to_token,
req_pool_indices=torch.arange(batch, dtype=torch.int32, device="cuda"),
kv_indptr=cumsum_indptr(kv_lens),
)
def report_result(
suite: str,
config: GqaConfig,
case: dict[str, int],
operations: dict[str, Callable[[], torch.Tensor]],
samples: dict[str, list[float]],
io_bytes: int,
reference: torch.Tensor,
actual: torch.Tensor,
) -> dict[str, object]:
difference = actual.float() - reference.float()
result: dict[str, object] = {
"suite": suite,
"config": config.name,
"agreement": {
"max_abs_error": float(difference.abs().max()),
"cosine_similarity": float(
F.cosine_similarity(
actual.float().flatten(), reference.float().flatten(), dim=0
)
),
},
"estimated_io_bytes": io_bytes,
**case,
}
for name, operation in operations.items():
latency = summarize(samples[name])
result[name] = {
"effective_bandwidth_gbps": io_bytes / (latency["median_ms"] / 1000) / 1e9,
**latency,
}
speedup = (result["torch"]["median_ms"] / result["cuda"]["median_ms"] - 1.0) * 100.0
label = f"B={case.get('batch')}" + (
f" ctx={case['context']}" if "context" in case else f" q={case['q_len']}"
)
print(
f"{suite},{config.name},{label},{result['torch']['median_ms']:.4f},"
f"{result['cuda']['median_ms']:.4f},{speedup:+.1f}%,"
f"{result['agreement']['max_abs_error']:.4f}"
)
return result
# ---------------------------------------------------------------------------
# Suites
# ---------------------------------------------------------------------------
def benchmark_decode(
config: GqaConfig,
batch: int,
context: int,
*,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, object]:
q = torch.randn(
batch, 1, config.hq, config.head_dim, device="cuda", dtype=torch.bfloat16
)
k = torch.randn(
batch,
context,
config.hkv,
config.head_dim,
device="cuda",
dtype=torch.bfloat16,
)
v = torch.randn_like(k)
# GQA expansion is data preparation, not attention compute — build it
# once so the timed torch side is a single SDPA launch.
k_expanded = repeat_kv_heads(k, config.hq // config.hkv)
v_expanded = repeat_kv_heads(v, config.hq // config.hkv)
def torch_op() -> torch.Tensor:
return sdpa(q, k_expanded, v_expanded)
def cuda_op() -> torch.Tensor:
return attn_decode(q, k, v, is_causal=True)
operations = {"torch": torch_op, "cuda": cuda_op}
samples = measure_operations(
operations, warmup=warmup, iterations=iterations, trials=trials
)
io_bytes = (2 * q.numel() + 2 * k.numel()) * q.element_size()
return report_result(
"decode",
config,
{"batch": batch, "context": context},
operations,
samples,
io_bytes,
torch_op(),
cuda_op(),
)
def benchmark_prefill(
config: GqaConfig,
batch: int,
q_len: int,
*,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, object]:
q = torch.randn(
batch, q_len, config.hq, config.head_dim, device="cuda", dtype=torch.bfloat16
)
k = torch.randn(
batch, q_len, config.hkv, config.head_dim, device="cuda", dtype=torch.bfloat16
)
v = torch.randn_like(k)
k_expanded = repeat_kv_heads(k, config.hq // config.hkv)
v_expanded = repeat_kv_heads(v, config.hq // config.hkv)
def torch_op() -> torch.Tensor:
return sdpa(q, k_expanded, v_expanded, is_causal=True)
def cuda_op() -> torch.Tensor:
return attn_prefill(q, k, v, is_causal=True)
operations = {"torch": torch_op, "cuda": cuda_op}
samples = measure_operations(
operations, warmup=warmup, iterations=iterations, trials=trials
)
io_bytes = (2 * q.numel() + 2 * k.numel()) * q.element_size()
return report_result(
"prefill",
config,
{"batch": batch, "q_len": q_len},
operations,
samples,
io_bytes,
torch_op(),
cuda_op(),
)
def benchmark_paged_decode(
config: GqaConfig,
batch: int,
context: int,
*,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, object]:
n_rep = config.hq // config.hkv
kv_lens = [length + 1 for length in ragged_lens(batch, context)]
inputs = build_paged_inputs(config, kv_lens, None)
new_k = torch.randn(
batch, config.hkv, config.head_dim, device="cuda", dtype=torch.bfloat16
)
new_v = torch.randn_like(new_k)
o_part = torch.empty(
batch,
config.hq,
MAX_SPLITS,
config.head_dim,
dtype=torch.float32,
device="cuda",
)
ml_part = torch.empty(
batch, config.hq, MAX_SPLITS, 2, dtype=torch.float32, device="cuda"
)
out_buf = torch.empty(
batch, config.hq, config.head_dim, dtype=torch.bfloat16, device="cuda"
)
def cuda_op() -> torch.Tensor:
return attn_paged_decode(
inputs.q,
inputs.k_cache,
inputs.v_cache,
inputs.req_to_token,
inputs.req_pool_indices,
inputs.kv_indptr,
new_k=new_k,
new_v=new_v,
is_causal=True,
o_part_buf=o_part,
ml_part_buf=ml_part,
out_buf=out_buf,
)
# Reference-side data preparation happens once, outside the timed
# closure: append the current-token K/V into the pool (the kernel does
# this fused inside its launch), gather padded K/V, expand GQA heads.
max_len = max(kv_lens)
slots = inputs.req_to_token[:, :max_len].long()
last_slots = inputs.req_to_token[
torch.arange(batch, device="cuda"), torch.tensor(kv_lens) - 1
].long()
inputs.k_cache[last_slots] = new_k
inputs.v_cache[last_slots] = new_v
k_expanded = repeat_kv_heads(inputs.k_cache[slots], n_rep)
v_expanded = repeat_kv_heads(inputs.v_cache[slots], n_rep)
position = torch.arange(max_len, device="cuda")
lengths = torch.tensor(kv_lens, device="cuda", dtype=torch.long)
keep_mask = (position[None, :] < lengths[:, None])[:, None, None, :]
q_batched = inputs.q.unsqueeze(1)
sdpa_call = prefer_cudnn_sdpa(
lambda: sdpa(q_batched, k_expanded, v_expanded, attn_mask=keep_mask)
)
def torch_op() -> torch.Tensor:
return sdpa_call().squeeze(1) # [B, Hq, D]
operations = {"torch": torch_op, "cuda": cuda_op}
samples = measure_operations(
operations, warmup=warmup, iterations=iterations, trials=trials
)
io_bytes = (
2 * inputs.q.numel() # q read + out write
+ 2 * sum(kv_lens) * config.hkv * config.head_dim # k/v reads
+ 2 * new_k.numel() # new k/v writes
) * inputs.q.element_size()
return report_result(
"paged_decode",
config,
{"batch": batch, "context": context},
operations,
samples,
io_bytes,
torch_op(),
cuda_op(),
)
def benchmark_paged_prefill(
config: GqaConfig,
batch: int,
q_len: int,
*,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, object]:
n_rep = config.hq // config.hkv
q_lens = ragged_lens(batch, q_len)
inputs = build_paged_inputs(config, q_lens, q_lens)
qo_indptr = cumsum_indptr(q_lens)
tile_batches, tile_indices = [], []
for request, length in enumerate(q_lens):
n_tiles = (length + Q_TILE_ROWS - 1) // Q_TILE_ROWS
tile_batches.extend([request] * n_tiles)
tile_indices.extend(range(n_tiles))
q_tile_to_batch = torch.tensor(tile_batches, dtype=torch.int32, device="cuda")
q_tile_to_index = torch.tensor(tile_indices, dtype=torch.int32, device="cuda")
def cuda_op() -> torch.Tensor:
return attn_paged_prefill(
inputs.q,
inputs.k_cache,
inputs.v_cache,
inputs.req_to_token,
inputs.req_pool_indices,
inputs.kv_indptr,
qo_indptr,
q_tile_to_batch,
q_tile_to_index,
is_causal=True,
)
# Same once-only preparation: gather padded K/V, expand GQA heads, pad Q,
# build the causal + validity mask. The timed reference is one SDPA call
# plus the packed-row unpack.
max_len = max(q_lens)
slots = inputs.req_to_token[:, :max_len].long()
k_expanded = repeat_kv_heads(inputs.k_cache[slots], n_rep)
v_expanded = repeat_kv_heads(inputs.v_cache[slots], n_rep)
position = torch.arange(max_len, device="cuda")
lengths = torch.tensor(q_lens, device="cuda", dtype=torch.long)
causal = position[None, :, None] >= position[None, None, :]
keep = position[None, None, :] < lengths[:, None, None]
attn_mask = (causal & keep).unsqueeze(1)
q_padded = torch.zeros(
batch,
max_len,
config.hq,
config.head_dim,
device="cuda",
dtype=inputs.q.dtype,
)
for i, length in enumerate(q_lens):
q_padded[i, :length] = inputs.q[int(qo_indptr[i]) : int(qo_indptr[i + 1])]
sdpa_call = prefer_cudnn_sdpa(
lambda: sdpa(q_padded, k_expanded, v_expanded, attn_mask=attn_mask)
)
def torch_op() -> torch.Tensor:
out = sdpa_call()
return torch.cat([out[i, :length] for i, length in enumerate(q_lens)])
operations = {"torch": torch_op, "cuda": cuda_op}
samples = measure_operations(
operations, warmup=warmup, iterations=iterations, trials=trials
)
io_bytes = (
2 * inputs.q.numel() + 2 * sum(q_lens) * config.hkv * config.head_dim
) * inputs.q.element_size()
return report_result(
"paged_prefill",
config,
{"batch": batch, "q_len": q_len},
operations,
samples,
io_bytes,
torch_op(),
cuda_op(),
)
@click.command(help=__doc__)
@click.option("--output", type=click.Path(path_type=Path), help="Optional JSON output.")
@click.option(
"--suite",
"suites",
type=click.Choice(("decode", "prefill", "paged_decode", "paged_prefill", "all")),
multiple=True,
default=("all",),
show_default=True,
)
@click.option(
"--config",
"config_values",
multiple=True,
help="Filter defaults by bare name, or add/override with NAME:HQ:HKV:HEAD_DIM.",
)
@click.option("--warmup", type=click.IntRange(min=1), default=10, show_default=True)
@click.option("--iterations", type=click.IntRange(min=1), default=50, show_default=True)
@click.option("--trials", type=click.IntRange(min=1), default=10, show_default=True)
@click.option("--seed", type=int, default=0, show_default=True)
def benchmark_command(
output: Path | None,
suites: tuple[str, ...],
config_values: tuple[str, ...],
warmup: int,
iterations: int,
trials: int,
seed: int,
) -> None:
if not torch.cuda.is_available():
raise click.ClickException("CUDA is required")
kernel_for_suite = {
"decode": "attn_decode",
"prefill": "attn_prefill",
"paged_decode": "attn_paged_decode",
"paged_prefill": "attn_paged_prefill",
}
selected = (
tuple(kernel_for_suite) if "all" in suites else tuple(dict.fromkeys(suites))
)
missing = [
kernel_for_suite[suite]
for suite in selected
if not is_available(kernel_for_suite[suite])
]
if missing:
raise click.ClickException(f"built kernels required: {', '.join(missing)}")
# A bare name filters the matching default; a full spec overrides or appends.
chosen: dict[str, GqaConfig] = {}
for value in config_values:
config = (
next((c for c in DEFAULT_CONFIGS if c.name == value), None)
if ":" not in value
else parse_config(value)
)
if config is None:
raise click.BadParameter(f"unknown default config {value!r}")
chosen[config.name] = config
configs = tuple(chosen.values()) or DEFAULT_CONFIGS
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
runners = {
"decode": (benchmark_decode, DECODE_CASES),
"prefill": (benchmark_prefill, PREFILL_CASES),
"paged_decode": (benchmark_paged_decode, DECODE_CASES),
"paged_prefill": (benchmark_paged_prefill, PREFILL_CASES),
}
print("suite,config,case,torch_ms,cuda_ms,speedup,max_abs")
results = []
with torch.inference_mode():
for suite in selected:
runner, cases = runners[suite]
for config in configs:
for case in cases:
results.append(
runner(
config,
*case,
warmup=warmup,
iterations=iterations,
trials=trials,
)
)
torch.cuda.empty_cache()
if output is not None:
props = torch.cuda.get_device_properties(0)
payload = {
"metadata": {
"gpu_name": props.name,
"compute_capability": f"{props.major}.{props.minor}",
"torch_version": torch.__version__,
"cuda_version": torch.version.cuda,
"timestamp_utc": datetime.now(timezone.utc).isoformat(),
},
"settings": {
"warmup": warmup,
"iterations": iterations,
"trials": trials,
"seed": seed,
"order": "A-B-B-A",
"suites": list(selected),
"configs": [asdict(config) for config in configs],
},
"results": results,
}
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
if __name__ == "__main__":
benchmark_command()