Files
AstrAI/scripts/tools/benchmark_gemv_common.py
T
0z5a d4a292b36b perf: tune bf16 gemv and add opt-in fused swiglu
- deepen common-shape BF16 GEMV tuning with warp-row tiling for LLaMA/Qwen2/GPT-NeoX/OPT decode projections
- add fused BF16 up/gate SwiGLU CUDA primitive with ASTRAI_SWIGLU=0/1/auto dispatch
- keep the unfused linear backend as the default path; auto enables no shape until per-architecture checkpoint gates pass
- fall back to the linear/torch chain when kernels are absent, on CPU, in training, or outside supported M/K/dtype shapes
- add gemv/swiglu benchmark scripts, dispatch and parity tests, and kernel documentation

Benchmark: NVIDIA L20 (sm_89), CUDA 12.8, PyTorch 2.11.0+cu128, idle GPU. AstrAI 1B config (24 layers, hidden 1536, vocab 100000), BF16, prompt 128, 32 greedy decode tokens, CUDA graphs enabled, A/B in separate interleaved processes (3 rounds, 8 trials each, medians). Default vs ASTRAI_SWIGLU=1 per generate call: batch 1 134.8->129.1 ms (+4.44%), batch 2 136.2->130.9 ms (+4.06%), batch 4 145.5->140.3 ms (+3.66%). Greedy output identical at batch 1, differs at batch 2/4, so auto stays unfused by default; kernelless fallback verified bit-identical greedy.
2026-09-03 04:26:53 +08:00

451 lines
15 KiB
Python

"""Benchmark the BF16 GEMV 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_gemv, 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_gemv(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_GEMV"] = 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_gemv"):
raise RuntimeError("benchmark requires CUDA and the built bf16_gemv 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()