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:
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user