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
This commit is contained in:
2026-09-05 01:38:10 +08:00
parent a77e35dd51
commit 6709534d64
24 changed files with 1306 additions and 3394 deletions
+649
View File
@@ -0,0 +1,649 @@
"""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()
+394
View File
@@ -0,0 +1,394 @@
"""Benchmark the FP8 quantize and GEMM kernels against torch baselines.
Suites (--suite): quantize (plain / delayed-scaling ring / dual-orientation
entries vs the aten float8 cast) and gemm (pre-quantized ``mm_fp8`` in the
NT orientation the fp8 linear path uses, vs bf16 ``F.linear``). GEMM
agreement reports both kernel error (vs the fp32 dequantized fp8 product)
and format error (that product vs the bf16 matmul). FP8 MMA requires
compute capability 89+.
"""
from __future__ import annotations
import json
import math
import statistics
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable
import click
import torch
import torch.nn.functional as F
from astrai.extension import is_available
from astrai.extension.ops.fp8 import mm_fp8, quantize, quantize_dual
FP8_MAX = {"e4m3": 448.0, "e5m2": 57344.0}
@dataclass(frozen=True)
class MatrixShape:
name: str
rows: int
cols: int
QUANTIZE_SHAPES = (
MatrixShape("astrai_1b_act", 2048, 1536),
MatrixShape("llama2_7b_act", 2048, 4096),
MatrixShape("llama2_7b_down_w", 4096, 11008),
MatrixShape("llama3_70b_act", 2048, 8192),
)
# GEMM shapes as (N, K) weight mats; M comes from --m-values.
GEMM_SHAPES = (
MatrixShape("llama2_7b_qkv", 4096, 4096),
MatrixShape("llama2_7b_up_gate", 11008, 4096),
MatrixShape("llama2_7b_down", 4096, 11008),
MatrixShape("llama3_70b_up_gate", 28672, 8192),
)
def parse_positive_ints(value: str) -> tuple[int, ...]:
try:
values = tuple(dict.fromkeys(int(item.strip()) for item in value.split(",")))
except ValueError as exc:
raise click.BadParameter("expected comma-separated integers") from exc
if not values or any(item <= 0 for item in values):
raise click.BadParameter("values must be positive integers")
return values
def parse_shape(value: str) -> MatrixShape:
parts = value.split(":")
if len(parts) != 3 or not parts[0]:
raise click.BadParameter("shape must use NAME:ROWS:COLS")
try:
rows, cols = (int(item) for item in parts[1:])
except ValueError as exc:
raise click.BadParameter("ROWS and COLS must be integers") from exc
if rows <= 0 or cols <= 0:
raise click.BadParameter("ROWS and COLS must be positive")
return MatrixShape(parts[0], rows, cols)
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-C-C-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 quant_step(x: torch.Tensor, fmt: str) -> torch.Tensor:
"""Quantization step (dequant scale) from the current amax; ``quantize``
takes its reciprocal as the multiplier."""
amax = x.abs().amax().to(torch.float32).clamp_min(1e-12)
return amax / FP8_MAX[fmt]
def benchmark_quantize(
shape: MatrixShape,
fmt: str,
*,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, object]:
x = torch.randn(shape.rows, shape.cols, device="cuda", dtype=torch.bfloat16) * 0.1
multiplier = quant_step(x, fmt).reciprocal()
fp8_dtype = torch.float8_e4m3fn if fmt == "e4m3" else torch.float8_e5m2
ring_state = torch.zeros(16 + 4, dtype=torch.float32, device="cuda")
operations: dict[str, Callable[[], torch.Tensor]] = {
"torch_cast": lambda: x.to(fp8_dtype),
"plain": lambda: quantize(x, multiplier, fmt)[0],
"ring": lambda: quantize(
x,
multiplier,
fmt,
ring_state=ring_state,
hist_idx=0,
fp8_max=FP8_MAX[fmt],
)[0],
"dual": lambda: quantize_dual(x, multiplier, fmt)[0],
}
samples = measure_operations(
operations, warmup=warmup, iterations=iterations, trials=trials
)
with torch.no_grad():
step = quant_step(x, fmt)
x8, _ = quantize(x, multiplier, fmt)
dequant_error = float((x8.to(torch.float32) * step - x.float()).abs().max())
dual_t = quantize_dual(x, multiplier, fmt)[1]
dual_matches = bool(torch.equal(dual_t.t().contiguous(), x8.contiguous()))
io_bytes = 3 * x.numel() + 4 # bf16 read + fp8 write + f32 amax
result: dict[str, object] = {
"suite": "quantize",
"shape": shape.name,
"rows": shape.rows,
"cols": shape.cols,
"fmt": fmt,
"estimated_io_bytes": io_bytes,
"dequant_max_abs_error": dequant_error,
"dual_transpose_matches": dual_matches,
}
for name, samples_ms in samples.items():
latency = summarize(samples_ms)
bytes_per_call = io_bytes + x.numel() if name == "dual" else io_bytes
result[name] = {
"effective_bandwidth_gbps": bytes_per_call
/ (latency["median_ms"] / 1000)
/ 1e9,
**latency,
}
speedup = (
result["torch_cast"]["median_ms"] / result["plain"]["median_ms"] - 1.0
) * 100.0
print(
f"quantize,{shape.name},{shape.rows}x{shape.cols},"
f"{result['torch_cast']['median_ms']:.4f},{result['plain']['median_ms']:.4f},"
f"{result['ring']['median_ms']:.4f},{result['dual']['median_ms']:.4f},"
f"{speedup:+.1f}%,{dequant_error:.5f},{dual_matches}"
)
return result
def benchmark_gemm(
shape: MatrixShape,
m: int,
fmt: str,
*,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, object]:
x = torch.randn(m, shape.cols, device="cuda", dtype=torch.bfloat16) * 0.1
w = (
torch.randn(shape.rows, shape.cols, device="cuda", dtype=torch.bfloat16)
* shape.cols**-0.5
)
sx, sw = quant_step(x, fmt), quant_step(w, fmt)
with torch.no_grad():
x8, _ = quantize(x, sx.reciprocal(), fmt)
w8, _ = quantize(w, sw.reciprocal(), fmt)
dequant_scale = sx * sw
def torch_op() -> torch.Tensor:
return F.linear(x, w)
def fp8_op() -> torch.Tensor:
return mm_fp8(x8, w8, dequant_scale, trans_b=True)
operations = {"torch_bf16": torch_op, "fp8": fp8_op}
samples = measure_operations(
operations, warmup=warmup, iterations=iterations, trials=trials
)
with torch.no_grad():
actual = fp8_op().float()
dequant_reference = (
x8.to(torch.float32) @ w8.to(torch.float32).t() * dequant_scale
)
bf16_reference = torch_op().float()
kernel_difference = actual - dequant_reference
format_difference = dequant_reference - bf16_reference
io_bytes = m * shape.cols + shape.rows * shape.cols + 2 * m * shape.rows
result: dict[str, object] = {
"suite": "gemm",
"shape": shape.name,
"m": m,
"n": shape.rows,
"k": shape.cols,
"fmt": fmt,
"estimated_io_bytes": io_bytes,
"kernel_max_abs": float(kernel_difference.abs().max()),
"kernel_rel_l2": float(
kernel_difference.norm() / dequant_reference.norm().clamp_min(1e-12)
),
"format_rel_l2": float(
format_difference.norm() / bf16_reference.norm().clamp_min(1e-12)
),
}
for name, samples_ms in samples.items():
latency = summarize(samples_ms)
result[name] = {
"effective_bandwidth_gbps": io_bytes / (latency["median_ms"] / 1000) / 1e9,
**latency,
}
speedup = (
result["torch_bf16"]["median_ms"] / result["fp8"]["median_ms"] - 1.0
) * 100.0
print(
f"gemm,{shape.name},{m}x{shape.rows}x{shape.cols},"
f"{result['torch_bf16']['median_ms']:.4f},{result['fp8']['median_ms']:.4f},"
f"{speedup:+.1f}%,{result['kernel_max_abs']:.4f},"
f"{result['kernel_rel_l2']:.6f},{result['format_rel_l2']:.6f}"
)
return result
@click.command(help=__doc__)
@click.option("--output", type=click.Path(path_type=Path), help="Optional JSON output.")
@click.option(
"--suite",
"suites",
type=click.Choice(("quantize", "gemm", "all")),
multiple=True,
default=("all",),
show_default=True,
)
@click.option("--fmt", type=click.Choice(("e4m3", "e5m2")), default="e4m3")
@click.option("--m-values", default="512,2048,4096", show_default=True)
@click.option(
"--shape",
"shape_values",
multiple=True,
help="Filter defaults by bare name (either suite), or add/override with "
"NAME:ROWS:COLS.",
)
@click.option("--warmup", type=click.IntRange(min=1), default=10, show_default=True)
@click.option(
"--iterations", type=click.IntRange(min=1), default=100, 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, ...],
fmt: str,
m_values: str,
shape_values: tuple[str, ...],
warmup: int,
iterations: int,
trials: int,
seed: int,
) -> None:
if not torch.cuda.is_available():
raise click.ClickException("CUDA is required")
if not is_available("fp8_ops"):
raise click.ClickException(
"the built fp8_ops kernel is required (compute capability 89+)"
)
selected = ("quantize", "gemm") if "all" in suites else tuple(dict.fromkeys(suites))
m_values_parsed = parse_positive_ints(m_values)
# Any --shape selection replaces the defaults for both suites: a bare
# name keeps that suite's matching default, a NAME:ROWS:COLS spec
# overrides the same-name default or adds a new one.
bare_names = {value for value in shape_values if ":" not in value}
known = {shape.name for shape in QUANTIZE_SHAPES + GEMM_SHAPES}
unknown = sorted(bare_names - known)
if unknown:
raise click.BadParameter(f"unknown default shape names: {', '.join(unknown)}")
specs = [parse_shape(value) for value in shape_values if ":" in value]
def resolve_shapes(defaults: tuple[MatrixShape, ...]) -> list[MatrixShape]:
if not shape_values:
return list(defaults)
by_name = {shape.name: shape for shape in defaults if shape.name in bare_names}
for spec in specs:
by_name[spec.name] = spec
return list(by_name.values())
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
results = []
with torch.inference_mode():
if "quantize" in selected:
print(
"suite,shape,size,cast_ms,plain_ms,ring_ms,dual_ms,vs_cast,"
"dequant_max,dual_ok"
)
for shape in resolve_shapes(QUANTIZE_SHAPES):
results.append(
benchmark_quantize(
shape, fmt, warmup=warmup, iterations=iterations, trials=trials
)
)
torch.cuda.empty_cache()
if "gemm" in selected:
print(
"suite,shape,mxn_xk,bf16_ms,fp8_ms,speedup,kernel_max,"
"kernel_rel_l2,format_rel_l2"
)
for shape in resolve_shapes(GEMM_SHAPES):
for m in m_values_parsed:
results.append(
benchmark_gemm(
shape,
m,
fmt,
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(),
"fmt": fmt,
},
"settings": {
"warmup": warmup,
"iterations": iterations,
"trials": trials,
"seed": seed,
"order": "A-B-C-C-B-A",
"suites": list(selected),
"m_values": list(m_values_parsed),
},
"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()
-333
View File
@@ -1,333 +0,0 @@
"""Benchmark decode-time linear shapes before enabling custom GEMM dispatch.
The benchmark deliberately calls ``torch.nn.functional.linear`` directly. It
establishes the per-architecture cuBLAS baseline that later GEMM primitives and
dispatch decisions must beat.
"""
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, Iterable
import click
import torch
import torch.nn.functional as F
@dataclass(frozen=True)
class LinearShape:
name: str
n: int
k: int
DEFAULT_SHAPES = (
LinearShape("q_proj", 1536, 1536),
LinearShape("k_proj", 256, 1536),
LinearShape("v_proj", 256, 1536),
LinearShape("attn_out", 1536, 1536),
LinearShape("mlp_up", 6912, 1536),
LinearShape("mlp_gate", 6912, 1536),
LinearShape("mlp_down", 1536, 6912),
LinearShape("lm_head", 100000, 1536),
)
DTYPES = {"bfloat16": torch.bfloat16, "float16": torch.float16}
def parse_positive_ints(value: str) -> tuple[int, ...]:
"""Parse a comma-separated, duplicate-free list of positive integers."""
try:
values = tuple(dict.fromkeys(int(item.strip()) for item in value.split(",")))
except ValueError as exc:
raise click.BadParameter("expected comma-separated integers") from exc
if not values or any(item <= 0 for item in values):
raise click.BadParameter("values must be positive integers")
return values
def parse_shape(value: str) -> LinearShape:
"""Parse NAME:N:K into a benchmark shape."""
parts = value.split(":")
if len(parts) != 3 or not parts[0]:
raise click.BadParameter("shape must use NAME:N:K")
try:
n, k = (int(item) for item in parts[1:])
except ValueError as exc:
raise click.BadParameter("N and K must be integers") from exc
if n <= 0 or k <= 0:
raise click.BadParameter("N and K must be positive")
return LinearShape(parts[0], n, k)
def estimate_io_bytes(
m: int, n: int, k: int, element_size: int, *, has_bias: bool
) -> int:
"""Estimate bytes touched once by Y[M,N] = X[M,K] @ W[N,K].T."""
elements = m * k + n * k + m * n
if has_bias:
elements += n
return elements * element_size
def percentile(values: Iterable[float], quantile: float) -> float:
ordered = sorted(values)
if not ordered:
raise ValueError("percentile requires at least one sample")
rank = (len(ordered) - 1) * quantile
lower = math.floor(rank)
upper = math.ceil(rank)
if lower == upper:
return ordered[lower]
fraction = rank - lower
return ordered[lower] * (1 - fraction) + ordered[upper] * fraction
def summarize_latency(samples_ms: list[float]) -> dict[str, float]:
return {
"median_ms": statistics.median(samples_ms),
"p90_ms": percentile(samples_ms, 0.90),
"p99_ms": percentile(samples_ms, 0.99),
"min_ms": min(samples_ms),
"max_ms": max(samples_ms),
}
def measure_cuda_ms(
operation: Callable[[], torch.Tensor], *, warmup: int, iterations: int, trials: int
) -> list[float]:
for _ in range(warmup):
operation()
torch.cuda.synchronize()
samples = []
for _ in range(trials):
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()
samples.append(start.elapsed_time(end) / iterations)
return samples
def count_cuda_kernels(
operation: Callable[[], torch.Tensor], repeats: int = 5
) -> float:
"""Profile a few calls and return the average device events per call."""
with torch.profiler.profile(
activities=[
torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA,
],
acc_events=True,
) as profile:
for _ in range(repeats):
operation()
torch.cuda.synchronize()
device_type = torch.autograd.DeviceType.CUDA
events = [event for event in profile.events() if event.device_type == device_type]
return len(events) / repeats
def capture_linear(
x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None
) -> tuple[torch.cuda.CUDAGraph, torch.Tensor]:
for _ in range(3):
F.linear(x, weight, bias)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
output = F.linear(x, weight, bias)
return graph, output
def benchmark_case(
shape: LinearShape,
m: int,
*,
dtype: torch.dtype,
mode: str,
bias_enabled: bool,
warmup: int,
iterations: int,
trials: int,
) -> dict[str, object]:
x = torch.randn((m, shape.k), device="cuda", dtype=dtype)
weight = torch.randn((shape.n, shape.k), device="cuda", dtype=dtype)
bias = torch.randn(shape.n, device="cuda", dtype=dtype) if bias_enabled else None
graph = None
graph_output = None
if mode == "graph":
graph, graph_output = capture_linear(x, weight, bias)
def operation() -> torch.Tensor:
graph.replay()
return graph_output
else:
def operation() -> torch.Tensor:
return F.linear(x, weight, bias)
samples_ms = measure_cuda_ms(
operation, warmup=warmup, iterations=iterations, trials=trials
)
latency = summarize_latency(samples_ms)
io_bytes = estimate_io_bytes(
m, shape.n, shape.k, x.element_size(), has_bias=bias is not None
)
median_seconds = latency["median_ms"] / 1000
result: dict[str, object] = {
"name": shape.name,
"m": m,
"n": shape.n,
"k": shape.k,
"mode": mode,
"bias": bias is not None,
"estimated_io_bytes": io_bytes,
"effective_bandwidth_gbps": io_bytes / median_seconds / 1e9,
"cuda_kernel_launches_per_call": count_cuda_kernels(operation),
**latency,
"samples_ms": samples_ms,
}
return result
def render_markdown(payload: dict[str, object]) -> str:
metadata = payload["metadata"]
assert isinstance(metadata, dict)
results = payload["results"]
assert isinstance(results, list)
lines = [
"# Decode linear baseline",
"",
f"- GPU: {metadata['gpu_name']}",
f"- Compute capability: {metadata['compute_capability']}",
f"- PyTorch / CUDA: {metadata['torch_version']} / {metadata['cuda_version']}",
f"- Dtype: {metadata['dtype']}",
"",
"| Layer | M | N | K | Mode | Median (ms) | p99 (ms) | GB/s | CUDA kernels/call |",
"|---|---:|---:|---:|---|---:|---:|---:|---:|",
]
for item in results:
assert isinstance(item, dict)
lines.append(
"| {name} | {m} | {n} | {k} | {mode} | {median_ms:.4f} | "
"{p99_ms:.4f} | {effective_bandwidth_gbps:.1f} | "
"{cuda_kernel_launches_per_call:.2f} |".format(**item)
)
lines.append("")
return "\n".join(lines)
def device_metadata(dtype_name: str) -> dict[str, object]:
props = torch.cuda.get_device_properties(0)
return {
"timestamp_utc": datetime.now(timezone.utc).isoformat(),
"gpu_name": props.name,
"compute_capability": f"{props.major}.{props.minor}",
"total_memory_bytes": props.total_memory,
"torch_version": torch.__version__,
"cuda_version": torch.version.cuda,
"dtype": dtype_name,
}
@click.command(help=__doc__)
@click.option("--output", type=click.Path(path_type=Path), required=True)
@click.option("--markdown-output", type=click.Path(path_type=Path))
@click.option("--m-values", default="1,2,4,8,16,32", show_default=True)
@click.option(
"--shape",
"shape_values",
multiple=True,
help="Override defaults with repeatable NAME:N:K shapes.",
)
@click.option("--dtype", type=click.Choice(tuple(DTYPES)), default="bfloat16")
@click.option("--mode", type=click.Choice(("eager", "graph", "both")), default="both")
@click.option("--bias/--no-bias", default=False)
@click.option("--warmup", type=click.IntRange(min=1), default=10, show_default=True)
@click.option(
"--iterations", type=click.IntRange(min=1), default=100, show_default=True
)
@click.option("--trials", type=click.IntRange(min=1), default=20, show_default=True)
@click.option("--seed", type=int, default=0, show_default=True)
def benchmark_command(
output: Path,
markdown_output: Path | None,
m_values: str,
shape_values: tuple[str, ...],
dtype: str,
mode: str,
bias: bool,
warmup: int,
iterations: int,
trials: int,
seed: int,
) -> None:
if not torch.cuda.is_available():
raise click.ClickException("CUDA is required")
parsed_m = parse_positive_ints(m_values)
shapes = tuple(parse_shape(item) for item in shape_values) or DEFAULT_SHAPES
modes = ("eager", "graph") if mode == "both" else (mode,)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
results = []
for shape in shapes:
for m in parsed_m:
for current_mode in modes:
click.echo(
f"{shape.name}: M={m} N={shape.n} K={shape.k} {current_mode}"
)
results.append(
benchmark_case(
shape,
m,
dtype=DTYPES[dtype],
mode=current_mode,
bias_enabled=bias,
warmup=warmup,
iterations=iterations,
trials=trials,
)
)
payload: dict[str, object] = {
"schema_version": 1,
"metadata": device_metadata(dtype),
"parameters": {
"m_values": list(parsed_m),
"shapes": [asdict(shape) for shape in shapes],
"modes": list(modes),
"bias": bias,
"warmup": warmup,
"iterations": iterations,
"trials": trials,
"seed": seed,
},
"results": results,
}
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
if markdown_output is not None:
markdown_output.parent.mkdir(parents=True, exist_ok=True)
markdown_output.write_text(render_markdown(payload), encoding="utf-8")
if __name__ == "__main__":
benchmark_command()
-450
View File
@@ -1,450 +0,0 @@
"""Benchmark the BF16 GEMM primitive and guarded linear dispatcher.
The kernel suite covers AstrAI's native projections plus common LLaMA and
GPT-NeoX matrix shapes. The chain suite is a synthetic projection/MLP chain;
it measures dispatcher overhead and dependent MLP work, but is deliberately
not presented as a whole-model throughput benchmark.
"""
import argparse
import gc
import json
import math
import os
import statistics
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
import torch
import torch.nn.functional as F
from astrai.extension import bf16_gemm, is_available, linear
@dataclass(frozen=True)
class Shape:
label: str
n: int
k: int
@dataclass(frozen=True)
class Chain:
label: str
hidden: int
kv: int
intermediate: int
fused_qkv: bool = False
gated_mlp: bool = True
@dataclass(frozen=True)
class Timing:
median_ms: float
p90_ms: float
ASTRAI_SHAPES = (
Shape("astrai_qkv", 256, 1536),
Shape("astrai_square", 1536, 1536),
Shape("astrai_up_gate", 6912, 1536),
Shape("astrai_down", 1536, 6912),
Shape("astrai_lm_head", 100000, 1536),
)
TRADITIONAL_SHAPES = (
Shape("llama2_7b_qo", 4096, 4096),
Shape("llama2_7b_up_gate", 11008, 4096),
Shape("llama2_7b_down", 4096, 11008),
Shape("llama3_8b_kv", 1024, 4096),
Shape("llama3_8b_up_gate", 14336, 4096),
Shape("llama3_8b_down", 4096, 14336),
Shape("llama2_13b_qo", 5120, 5120),
Shape("llama2_13b_up_gate", 13824, 5120),
Shape("llama2_13b_down", 5120, 13824),
Shape("gpt_neox_up", 16384, 4096),
Shape("gpt_neox_down", 4096, 16384),
Shape("qwen2_7b_kv", 512, 3584),
Shape("qwen2_7b_qo", 3584, 3584),
Shape("qwen2_7b_up_gate", 18944, 3584),
Shape("qwen2_7b_down", 3584, 18944),
Shape("llama3_70b_kv", 1024, 8192),
Shape("llama3_70b_qo", 8192, 8192),
Shape("llama3_70b_up_gate", 28672, 8192),
Shape("llama3_70b_down", 8192, 28672),
Shape("opt_1_3b_qkvo", 2048, 2048),
Shape("opt_1_3b_up", 8192, 2048),
Shape("opt_1_3b_down", 2048, 8192),
)
CHAINS = (
Chain("llama2_7b", 4096, 4096, 11008),
Chain("llama3_8b", 4096, 1024, 14336),
Chain("llama2_13b", 5120, 5120, 13824),
Chain("gpt_neox_20b", 4096, 4096, 16384, fused_qkv=True),
Chain("qwen2_7b", 3584, 512, 18944),
Chain("llama3_70b", 8192, 1024, 28672),
Chain("opt_1_3b", 2048, 2048, 8192, gated_mlp=False),
)
def _elapsed_ms(fn: Callable[[], torch.Tensor], inner: int) -> float:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(inner):
fn()
end.record()
end.synchronize()
return start.elapsed_time(end) / inner
def _timing(values: list[float]) -> Timing:
ordered = sorted(values)
p90_index = max(0, math.ceil(0.9 * len(ordered)) - 1)
return Timing(statistics.median(ordered), ordered[p90_index])
def _measure_pair(
baseline: Callable[[], torch.Tensor],
candidate: Callable[[], torch.Tensor],
*,
warmup: int,
samples: int,
inner: int,
prepare_baseline: Callable[[], None] = lambda: None,
prepare_candidate: Callable[[], None] = lambda: None,
) -> tuple[Timing, Timing]:
cases = (
("baseline", prepare_baseline, baseline),
("candidate", prepare_candidate, candidate),
)
for iteration in range(warmup):
_, prepare, fn = cases[iteration % 2]
prepare()
fn()
torch.cuda.synchronize()
values: dict[str, list[float]] = {"baseline": [], "candidate": []}
for sample in range(samples):
order = cases if sample % 2 == 0 else tuple(reversed(cases))
for label, prepare, fn in order:
prepare()
values[label].append(_elapsed_ms(fn, inner))
return _timing(values["baseline"]), _timing(values["candidate"])
def _print_header() -> None:
print(
"suite,label,m,n,k,torch_median_ms,torch_p90_ms,"
"candidate_median_ms,candidate_p90_ms,speedup_pct,"
"max_abs,relative_l2,argmax_equal"
)
def _print_result(
suite: str,
label: str,
m: int,
n: int,
k: int,
baseline: Timing,
candidate: Timing,
reference: torch.Tensor,
actual: torch.Tensor,
) -> dict[str, object]:
difference = actual.float() - reference.float()
max_abs = difference.abs().max().item()
relative_l2 = difference.norm().item() / max(reference.float().norm().item(), 1e-12)
argmax_equal = torch.equal(actual.argmax(dim=-1), reference.argmax(dim=-1))
speedup = (baseline.median_ms / candidate.median_ms - 1.0) * 100.0
result: dict[str, object] = {
"suite": suite,
"label": label,
"m": m,
"n": n,
"k": k,
"torch_median_ms": baseline.median_ms,
"torch_p90_ms": baseline.p90_ms,
"candidate_median_ms": candidate.median_ms,
"candidate_p90_ms": candidate.p90_ms,
"speedup_pct": speedup,
"max_abs": max_abs,
"relative_l2": relative_l2,
"argmax_equal": argmax_equal,
}
print(
f"{suite},{label},{m},{n},{k},"
f"{baseline.median_ms:.6f},{baseline.p90_ms:.6f},"
f"{candidate.median_ms:.6f},{candidate.p90_ms:.6f},"
f"{speedup:+.2f},{max_abs:.6f},{relative_l2:.8f},"
f"{str(argmax_equal).lower()}",
flush=True,
)
return result
def _weight(n: int, k: int, device: torch.device, std: float) -> torch.Tensor:
weight = torch.empty((n, k), device=device, dtype=torch.bfloat16)
weight.normal_(mean=0.0, std=std)
return weight.requires_grad_(True)
def _kernel_functions(
x: torch.Tensor, weight: torch.Tensor
) -> tuple[Callable[[], torch.Tensor], Callable[[], torch.Tensor]]:
def baseline() -> torch.Tensor:
return F.linear(x, weight)
def candidate() -> torch.Tensor:
return bf16_gemm(x, weight.detach())
return baseline, candidate
def benchmark_kernels(
args: argparse.Namespace, device: torch.device
) -> list[dict[str, object]]:
if args.family == "astrai":
shapes = ASTRAI_SHAPES
elif args.family == "traditional":
shapes = TRADITIONAL_SHAPES
else:
shapes = ASTRAI_SHAPES + TRADITIONAL_SHAPES
if args.shape_label:
requested = set(args.shape_label)
shapes = tuple(shape for shape in shapes if shape.label in requested)
missing = requested - {shape.label for shape in shapes}
if missing:
raise ValueError(f"unknown shape labels: {', '.join(sorted(missing))}")
results: list[dict[str, object]] = []
for shape in shapes:
weight = _weight(shape.n, shape.k, device, args.weight_std)
for m in args.m:
x = torch.randn((m, shape.k), device=device, dtype=torch.bfloat16)
baseline_fn, candidate_fn = _kernel_functions(x, weight)
with torch.inference_mode():
reference = baseline_fn()
actual = candidate_fn()
baseline, candidate = _measure_pair(
baseline_fn,
candidate_fn,
warmup=args.warmup,
samples=args.samples,
inner=args.inner,
)
results.append(
_print_result(
"kernel",
shape.label,
m,
shape.n,
shape.k,
baseline,
candidate,
reference,
actual,
)
)
del baseline_fn, candidate_fn, x, reference, actual
del weight
gc.collect()
torch.cuda.empty_cache()
return results
def _set_mode(mode: str) -> None:
os.environ["ASTRAI_GEMM"] = mode
def _chain_weights(
spec: Chain, device: torch.device, std: float
) -> dict[str, torch.Tensor]:
weights = {
"o": _weight(spec.hidden, spec.hidden, device, std),
"up": _weight(spec.intermediate, spec.hidden, device, std),
"down": _weight(spec.hidden, spec.intermediate, device, std),
}
if spec.fused_qkv:
weights["qkv"] = _weight(3 * spec.hidden, spec.hidden, device, std)
else:
weights.update(
{
"q": _weight(spec.hidden, spec.hidden, device, std),
"k": _weight(spec.kv, spec.hidden, device, std),
"v": _weight(spec.kv, spec.hidden, device, std),
}
)
if spec.gated_mlp:
weights["gate"] = _weight(spec.intermediate, spec.hidden, device, std)
return weights
def _chain_fn(
x: torch.Tensor, weights: dict[str, torch.Tensor], spec: Chain
) -> Callable[[], torch.Tensor]:
def run() -> torch.Tensor:
output_projection = linear(x, weights["o"])
up = linear(x, weights["up"])
if spec.fused_qkv:
attention_projection = linear(x, weights["qkv"])[..., : x.shape[-1]]
hidden = F.gelu(up)
else:
attention_projection = linear(x, weights["q"])
linear(x, weights["k"])
linear(x, weights["v"])
if spec.gated_mlp:
gate = linear(x, weights["gate"])
hidden = F.silu(gate) * up
else:
hidden = F.gelu(up)
down = linear(hidden, weights["down"])
return attention_projection + output_projection + down
return run
def benchmark_chains(
args: argparse.Namespace, device: torch.device
) -> list[dict[str, object]]:
results: list[dict[str, object]] = []
chains = CHAINS
if args.chain_label:
requested = set(args.chain_label)
chains = tuple(chain for chain in chains if chain.label in requested)
missing = requested - {chain.label for chain in chains}
if missing:
raise ValueError(f"unknown chain labels: {', '.join(sorted(missing))}")
for spec in chains:
weights = _chain_weights(spec, device, args.weight_std)
for m in args.m:
x = torch.randn((m, spec.hidden), device=device, dtype=torch.bfloat16)
run = _chain_fn(x, weights, spec)
with torch.inference_mode():
_set_mode("0")
reference = run()
_set_mode(args.candidate_mode)
actual = run()
baseline, candidate = _measure_pair(
run,
run,
warmup=args.warmup,
samples=args.samples,
inner=args.chain_inner,
prepare_baseline=lambda: _set_mode("0"),
prepare_candidate=lambda: _set_mode(args.candidate_mode),
)
results.append(
_print_result(
"synthetic_chain",
spec.label,
m,
spec.hidden,
spec.intermediate,
baseline,
candidate,
reference,
actual,
)
)
del x, reference, actual
del weights
gc.collect()
torch.cuda.empty_cache()
return results
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--suite", choices=("kernel", "chain", "all"), default="all")
parser.add_argument(
"--family", choices=("astrai", "traditional", "all"), default="all"
)
parser.add_argument(
"--m", type=int, nargs="+", choices=(1, 2, 4, 8), default=(1, 2, 4, 8)
)
parser.add_argument(
"--shape-label",
action="append",
help="limit the kernel suite to one or more named shape labels",
)
parser.add_argument(
"--chain-label",
action="append",
help="limit the chain suite to one or more named model families",
)
parser.add_argument("--device", type=int, default=0)
parser.add_argument("--warmup", type=int, default=20)
parser.add_argument("--samples", type=int, default=9)
parser.add_argument("--inner", type=int, default=100)
parser.add_argument("--chain-inner", type=int, default=20)
parser.add_argument(
"--candidate-mode",
choices=("auto", "1"),
default="auto",
help="dispatcher mode for the candidate side of the chain suite",
)
parser.add_argument("--weight-std", type=float, default=0.02)
parser.add_argument("--seed", type=int, default=20260902)
parser.add_argument(
"--output",
type=Path,
help="optional JSON output; stdout always retains the compact CSV table",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
if not torch.cuda.is_available() or not is_available("bf16_gemm"):
raise RuntimeError("benchmark requires CUDA and the built bf16_gemm extension")
if args.warmup < 0 or args.samples < 1 or args.inner < 1 or args.chain_inner < 1:
raise ValueError("warmup must be non-negative and sample/inner counts positive")
torch.cuda.set_device(args.device)
device = torch.device("cuda", args.device)
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
properties = torch.cuda.get_device_properties(device)
print(
f"# device={properties.name}, capability={properties.major}.{properties.minor}, "
f"seed={args.seed}, weight_std={args.weight_std}"
)
_print_header()
results: list[dict[str, object]] = []
if args.suite in ("kernel", "all"):
results.extend(benchmark_kernels(args, device))
if args.suite in ("chain", "all"):
results.extend(benchmark_chains(args, device))
if args.output is not None:
payload = {
"environment": {
"device": properties.name,
"capability": f"{properties.major}.{properties.minor}",
"torch": torch.__version__,
"cuda": torch.version.cuda,
},
"parameters": {
"suite": args.suite,
"family": args.family,
"m": args.m,
"shape_labels": args.shape_label,
"chain_labels": args.chain_label,
"candidate_mode": args.candidate_mode,
"seed": args.seed,
"weight_std": args.weight_std,
"warmup": args.warmup,
"samples": args.samples,
"inner": args.inner,
"chain_inner": args.chain_inner,
},
"results": results,
}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(payload, indent=2) + "\n")
if __name__ == "__main__":
main()
+252
View File
@@ -0,0 +1,252 @@
"""Benchmark the fused rotary-embedding kernel against the torch fallback.
The baseline is the complex-multiply fallback from
``astrai.extension.backend.rotary``. Layouts mirror the production call
shapes: packed 3D [tokens, n_heads, head_dim] and dense 4D
[batch, seq_len, n_heads, head_dim]; positions are random integers so every
row exercises a distinct cos/sin gather.
"""
from __future__ import annotations
import json
import math
import statistics
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable
import click
import torch
import torch.nn.functional as F
from astrai.extension import is_available
from astrai.extension.ops import rotary_emb
@dataclass(frozen=True)
class RotaryCase:
name: str
layout: str
batch: int
seq_len: int
heads: int
head_dim: int
DEFAULT_CASES = (
RotaryCase("decode_bs1", "packed", 1, 1, 32, 128),
RotaryCase("decode_bs32", "packed", 32, 1, 32, 128),
RotaryCase("prefill_4k_llama7b", "packed", 1, 4096, 32, 128),
RotaryCase("prefill_4k_llama70b", "packed", 1, 4096, 64, 128),
RotaryCase("train_8x2k_llama7b", "dense", 8, 2048, 32, 128),
RotaryCase("train_4x1k_d64", "dense", 4, 1024, 32, 64),
)
def parse_case(value: str) -> RotaryCase:
parts = value.split(":")
if len(parts) != 6 or not parts[0]:
raise click.BadParameter("case must use NAME:LAYOUT:BATCH:SEQ:HEADS:HEAD_DIM")
name, layout, batch, seq_len, heads, head_dim = parts
try:
fields = (int(batch), int(seq_len), int(heads), int(head_dim))
except ValueError as exc:
raise click.BadParameter("fields must be integers") from exc
if layout not in ("packed", "dense"):
raise click.BadParameter("layout must be 'packed' or 'dense'")
if any(field <= 0 for field in fields) or head_dim % 2:
raise click.BadParameter("fields must be positive; HEAD_DIM even")
return RotaryCase(name, layout, fields[0], fields[1], fields[2], fields[3])
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 torch_apply(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
"""The complex-multiply fallback (mirrors backend.rotary._torch_apply)."""
cos, sin = freqs_cis[..., 0], freqs_cis[..., 1]
dtype = x.dtype
x_ = x.float().reshape(*x.shape[:-1], -1, 2)
x_complex = torch.view_as_complex(x_)
freqs_cis_complex = torch.complex(cos, sin).unsqueeze(-2)
x_rotated = x_complex * freqs_cis_complex
return torch.view_as_real(x_rotated).flatten(-2).to(dtype)
def build_freqs(head_dim: int, positions: torch.Tensor) -> torch.Tensor:
"""[cos, sin] pairs for the given positions, laid out [..., head_dim/2, 2]."""
theta = 10000.0 ** (
-torch.arange(0, head_dim, 2, dtype=torch.float64, device=positions.device)
/ head_dim
)
freqs = positions.double().unsqueeze(-1) * theta
return torch.stack([freqs.cos(), freqs.sin()], dim=-1).float()
def benchmark_case(
case: RotaryCase, *, warmup: int, iterations: int, trials: int
) -> dict[str, object]:
if case.layout == "packed":
tokens = case.batch * case.seq_len
x = torch.randn(
tokens, case.heads, case.head_dim, device="cuda", dtype=torch.bfloat16
)
positions = torch.randint(0, 65536, (tokens,), device="cuda")
else:
x = torch.randn(
case.batch,
case.seq_len,
case.heads,
case.head_dim,
device="cuda",
dtype=torch.bfloat16,
)
positions = torch.randint(0, 65536, (case.batch, case.seq_len), device="cuda")
freqs_cis = build_freqs(case.head_dim, positions)
operations: dict[str, Callable[[], torch.Tensor]] = {
"torch": lambda: torch_apply(x, freqs_cis),
"cuda": lambda: rotary_emb(x, freqs_cis),
}
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))
with torch.no_grad():
expected = operations["torch"]().float()
actual = operations["cuda"]().float()
difference = actual - expected
io_bytes = (
2 * x.numel() * x.element_size() + freqs_cis.numel() * freqs_cis.element_size()
)
result: dict[str, object] = {
"case": case.name,
"layout": case.layout,
"heads": case.heads,
"head_dim": case.head_dim,
"estimated_io_bytes": io_bytes,
"max_abs_error": float(difference.abs().max()),
"cosine_similarity": float(
F.cosine_similarity(actual.flatten(), expected.flatten(), dim=0)
),
}
for name, samples_ms in samples.items():
latency = summarize(samples_ms)
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
print(
f"{case.name},{case.layout},{result['torch']['median_ms']:.5f},"
f"{result['cuda']['median_ms']:.5f},{speedup:+.1f}%,"
f"{result['max_abs_error']:.5f}"
)
return result
@click.command(help=__doc__)
@click.option("--output", type=click.Path(path_type=Path), help="Optional JSON output.")
@click.option(
"--case",
"case_values",
multiple=True,
help="Filter defaults by bare name, or add/override with "
"NAME:LAYOUT:BATCH:SEQ:HEADS:HEAD_DIM.",
)
@click.option("--warmup", type=click.IntRange(min=1), default=10, show_default=True)
@click.option(
"--iterations", type=click.IntRange(min=1), default=100, 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,
case_values: tuple[str, ...],
warmup: int,
iterations: int,
trials: int,
seed: int,
) -> None:
if not torch.cuda.is_available():
raise click.ClickException("CUDA is required")
if not is_available("rotary_emb"):
raise click.ClickException("the built rotary_emb kernel is required")
# A bare name filters the matching default; a full spec overrides or appends.
chosen: dict[str, RotaryCase] = {}
for value in case_values:
case = (
next((c for c in DEFAULT_CASES if c.name == value), None)
if ":" not in value
else parse_case(value)
)
if case is None:
raise click.BadParameter(f"unknown default case {value!r}")
chosen[case.name] = case
cases = tuple(chosen.values()) or DEFAULT_CASES
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
print("case,layout,torch_ms,cuda_ms,speedup,max_abs")
results = []
with torch.inference_mode():
for case in cases:
results.append(
benchmark_case(
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",
},
"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()
-325
View File
@@ -1,325 +0,0 @@
"""Benchmark fused BF16 SwiGLU against torch and unfused GEMM chains."""
from __future__ import annotations
import json
import math
import statistics
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable, Iterable
import click
import torch
import torch.nn.functional as F
from astrai.extension import bf16_gemm, bf16_swiglu, is_available
@dataclass(frozen=True)
class SwiGLUShape:
name: str
n: int
k: int
DEFAULT_SHAPES = (
SwiGLUShape("astrai_1b", 6912, 1536),
SwiGLUShape("llama2_7b", 11008, 4096),
SwiGLUShape("llama3_8b", 14336, 4096),
SwiGLUShape("llama2_13b", 13824, 5120),
SwiGLUShape("gpt_neox_20b", 16384, 6144),
)
def parse_positive_ints(value: str) -> tuple[int, ...]:
try:
values = tuple(dict.fromkeys(int(item.strip()) for item in value.split(",")))
except ValueError as exc:
raise click.BadParameter("expected comma-separated integers") from exc
if not values or any(item <= 0 for item in values):
raise click.BadParameter("values must be positive integers")
return values
def parse_shape(value: str) -> SwiGLUShape:
parts = value.split(":")
if len(parts) != 3 or not parts[0]:
raise click.BadParameter("shape must use NAME:N:K")
try:
n, k = (int(item) for item in parts[1:])
except ValueError as exc:
raise click.BadParameter("N and K must be integers") from exc
if n <= 0 or k <= 0 or k % 8:
raise click.BadParameter("N must be positive and K positive/divisible by 8")
return SwiGLUShape(parts[0], n, k)
def percentile(values: Iterable[float], quantile: float) -> float:
ordered = sorted(values)
rank = (len(ordered) - 1) * quantile
lower = math.floor(rank)
upper = math.ceil(rank)
if lower == upper:
return ordered[lower]
fraction = rank - lower
return ordered[lower] * (1 - fraction) + ordered[upper] * fraction
def summarize(values: list[float]) -> dict[str, float]:
return {
"median_ms": statistics.median(values),
"p90_ms": percentile(values, 0.90),
"p99_ms": percentile(values, 0.99),
"min_ms": min(values),
"max_ms": max(values),
}
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 count_cuda_kernels(
operation: Callable[[], torch.Tensor], repeats: int = 5
) -> float:
with torch.profiler.profile(
activities=[
torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA,
],
acc_events=True,
) as profile:
for _ in range(repeats):
operation()
torch.cuda.synchronize()
device_type = torch.autograd.DeviceType.CUDA
events = [event for event in profile.events() if event.device_type == device_type]
return len(events) / repeats
def capture(operation: Callable[[], torch.Tensor]):
for _ in range(3):
operation()
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
output = operation()
def replay() -> torch.Tensor:
graph.replay()
return output
return replay
def make_operations(x, up_weight, gate_weight, mode: str):
operations: dict[str, Callable[[], torch.Tensor]] = {
"torch": lambda: F.linear(x, up_weight) * F.silu(F.linear(x, gate_weight)),
"gemm_chain": lambda: (
bf16_gemm(x, up_weight) * F.silu(bf16_gemm(x, gate_weight))
),
"fused": lambda: bf16_swiglu(x, up_weight, gate_weight),
}
if mode == "graph":
operations = {name: capture(op) for name, op in operations.items()}
return operations
def benchmark_case(
shape: SwiGLUShape,
m: int,
mode: str,
*,
warmup: int,
iterations: int,
trials: int,
) -> list[dict[str, object]]:
x = torch.randn((m, shape.k), device="cuda", dtype=torch.bfloat16) * 0.1
scale = shape.k**-0.5
up_weight = (
torch.randn((shape.n, shape.k), device="cuda", dtype=torch.bfloat16) * scale
)
gate_weight = (
torch.randn((shape.n, shape.k), device="cuda", dtype=torch.bfloat16) * scale
)
operations = make_operations(x, up_weight, gate_weight, mode)
for operation in operations.values():
for _ in range(warmup):
operation()
torch.cuda.synchronize()
samples = {name: [] for name in operations}
forward_order = tuple(operations)
# A-B-C-C-B-A order balances cache, clock, and temperature drift.
for _ in range(trials):
for name in (*forward_order, *reversed(forward_order)):
samples[name].append(time_operation(operations[name], iterations))
with torch.no_grad():
expected = operations["torch"]().clone()
actual = operations["fused"]().clone()
difference = (actual.float() - expected.float()).abs()
max_abs_error = float(difference.max())
mean_abs_error = float(difference.mean())
cosine_similarity = float(
F.cosine_similarity(actual.float().flatten(), expected.float().flatten(), dim=0)
)
results = []
for name, operation in operations.items():
result: dict[str, object] = {
"shape": shape.name,
"m": m,
"n": shape.n,
"k": shape.k,
"mode": mode,
"implementation": name,
"cuda_kernel_launches_per_call": count_cuda_kernels(operation),
**summarize(samples[name]),
}
if name == "fused":
result.update(
max_abs_error=max_abs_error,
mean_abs_error=mean_abs_error,
cosine_similarity=cosine_similarity,
)
results.append(result)
return results
def device_metadata() -> dict[str, object]:
props = torch.cuda.get_device_properties(0)
return {
"timestamp_utc": datetime.now(timezone.utc).isoformat(),
"gpu_name": props.name,
"compute_capability": f"{props.major}.{props.minor}",
"total_memory_bytes": props.total_memory,
"torch_version": torch.__version__,
"cuda_version": torch.version.cuda,
"dtype": "bfloat16",
}
def render_markdown(payload: dict[str, object]) -> str:
metadata = payload["metadata"]
results = payload["results"]
assert isinstance(metadata, dict)
assert isinstance(results, list)
by_case = {
(item["shape"], item["m"], item["mode"], item["implementation"]): item
for item in results
}
cases = sorted({(item["shape"], item["m"], item["mode"]) for item in results})
lines = [
"# Fused SwiGLU benchmark",
"",
f"- GPU: {metadata['gpu_name']}",
f"- Compute capability: {metadata['compute_capability']}",
f"- PyTorch / CUDA: {metadata['torch_version']} / {metadata['cuda_version']}",
"",
"| Shape | M | Mode | torch ms | GEMM chain ms | fused ms | "
"vs best unfused | fused kernels | max abs | cosine |",
"|---|---:|---|---:|---:|---:|---:|---:|---:|---:|",
]
for shape, m, mode in cases:
torch_item = by_case[(shape, m, mode, "torch")]
gemm_item = by_case[(shape, m, mode, "gemm_chain")]
fused_item = by_case[(shape, m, mode, "fused")]
best = min(torch_item["median_ms"], gemm_item["median_ms"])
improvement = (best / fused_item["median_ms"] - 1) * 100
lines.append(
f"| {shape} | {m} | {mode} | {torch_item['median_ms']:.5f} | "
f"{gemm_item['median_ms']:.5f} | {fused_item['median_ms']:.5f} | "
f"{improvement:+.2f}% | "
f"{fused_item['cuda_kernel_launches_per_call']:.1f} | "
f"{fused_item['max_abs_error']:.5f} | "
f"{fused_item['cosine_similarity']:.8f} |"
)
lines.append("")
return "\n".join(lines)
@click.command(help=__doc__)
@click.option("--output", type=click.Path(path_type=Path), required=True)
@click.option("--markdown-output", type=click.Path(path_type=Path))
@click.option("--m-values", default="1,2,4,8", show_default=True)
@click.option("--shape", "shape_values", multiple=True, help="Repeat NAME:N:K.")
@click.option("--mode", type=click.Choice(("eager", "graph", "both")), default="both")
@click.option("--warmup", type=click.IntRange(min=1), default=10, show_default=True)
@click.option(
"--iterations", type=click.IntRange(min=1), default=100, 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,
markdown_output: Path | None,
m_values: str,
shape_values: tuple[str, ...],
mode: str,
warmup: int,
iterations: int,
trials: int,
seed: int,
) -> None:
if not torch.cuda.is_available():
raise click.ClickException("CUDA is required")
if not is_available("bf16_gemm") or not is_available("bf16_swiglu"):
raise click.ClickException("built bf16_gemm and bf16_swiglu are required")
shapes = tuple(parse_shape(value) for value in shape_values) or DEFAULT_SHAPES
m_values_parsed = parse_positive_ints(m_values)
if any(m > 8 for m in m_values_parsed):
raise click.BadParameter("fused primitive supports M up to 8")
modes = ("eager", "graph") if mode == "both" else (mode,)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
results = []
with torch.inference_mode():
for shape in shapes:
for m in m_values_parsed:
for current_mode in modes:
click.echo(
f"{shape.name}: M={m} N={shape.n} K={shape.k} {current_mode}"
)
results.extend(
benchmark_case(
shape,
m,
current_mode,
warmup=warmup,
iterations=iterations,
trials=trials,
)
)
torch.cuda.empty_cache()
payload: dict[str, object] = {
"metadata": device_metadata(),
"settings": {
"warmup": warmup,
"iterations": iterations,
"trials": trials,
"seed": seed,
"order": "A-B-C-C-B-A",
},
"results": results,
}
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(payload, indent=2) + "\n")
if markdown_output is not None:
markdown_output.parent.mkdir(parents=True, exist_ok=True)
markdown_output.write_text(render_markdown(payload))
if __name__ == "__main__":
benchmark_command()