flashrt-fp8-ffn / benchmarks /benchmark_bf16_entry.py
liangsu9988's picture
Promote latest kernel artifacts to main
f17d8e9 verified
Raw
History Blame Contribute Delete
13.6 kB
#!/usr/bin/env python3
"""Benchmark the BF16-to-FP8 GELU FFN region boundary."""
from __future__ import annotations
import argparse
import json
import statistics
from dataclasses import asdict, dataclass
from pathlib import Path
import torch
import torch.nn.functional as F
import benchmark as base
SHAPES = {
"siglip_m8": (8, 1152, 4304, 1152),
"siglip_m51": (51, 1152, 4304, 1152),
"siglip_m64": (64, 1152, 4304, 1152),
"siglip_m105": (105, 1152, 4304, 1152),
"siglip_m128": (128, 1152, 4304, 1152),
"dit_m8": (8, 1536, 6144, 1536),
"dit_m51": (51, 1536, 6144, 1536),
"dit_m64": (64, 1536, 6144, 1536),
"dit_m105": (105, 1536, 6144, 1536),
"dit_m128": (128, 1536, 6144, 1536),
}
@dataclass
class Result:
shape: str
M: int
K: int
H: int
N: int
flashrt_bf16_entry_us: float
flashrt_cuda_graph_us: float
separate_quant_us: float
fp8_kernel_only_us: float
torch_bf16_eager_us: float
torch_bf16_compile_us: float | None
speedup_vs_separate_quant: float
speedup_vs_eager: float
speedup_vs_compile: float | None
compile_status: str
flashrt_compile_status: str
input_quant_exact: bool
output_dtype: str
staged_max_abs: float
staged_mean_abs: float
staged_p99_abs: float
staged_cosine: float
bf16_max_abs: float
bf16_mean_abs: float
bf16_p99_abs: float
bf16_cosine: float
performance_status: str
status: str
def percentile(x: torch.Tensor, q: float) -> float:
flat = x.flatten()
k = max(1, min(flat.numel(), int(q * flat.numel() + 0.999999)))
return float(flat.kthvalue(k).values.item())
def metrics(got: torch.Tensor, expected: torch.Tensor) -> dict[str, float]:
diff = (got.float() - expected.float()).abs().flatten()
cosine = F.cosine_similarity(
got.float().flatten(), expected.float().flatten(), dim=0
)
return {
"max_abs": float(diff.max().item()),
"mean_abs": float(diff.mean().item()),
"p99_abs": percentile(diff, 0.99),
"cosine": float(cosine.item()),
}
def quantize_input(x: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
inv_scale = 1.0 / scale.float()
return torch.clamp(
x.float() * inv_scale, -base.fp8_max(), base.fp8_max()
).to(base.fp8_dtype())
def midm_padded_rows(rows: int) -> int:
if (
torch.version.hip is None
and torch.cuda.get_device_capability(0) == (11, 0)
and 9 <= rows <= 128
):
return ((rows + 63) // 64) * 64
return rows
def median_time_us(fn, args) -> float:
return statistics.median(
base.time_us(fn, warmup=args.warmup, iters=args.iters)
for _ in range(args.rounds)
)
def abba_time_us(a, b, args) -> tuple[float, float]:
a_samples = []
b_samples = []
for _ in range(args.rounds):
a_samples.append(base.time_us(a, warmup=args.warmup, iters=args.iters))
b_samples.append(base.time_us(b, warmup=args.warmup, iters=args.iters))
b_samples.append(base.time_us(b, warmup=args.warmup, iters=args.iters))
a_samples.append(base.time_us(a, warmup=args.warmup, iters=args.iters))
return statistics.median(a_samples), statistics.median(b_samples)
def make_case(M: int, K: int, H: int, N: int):
x = torch.randn((M, K), device="cuda", dtype=torch.bfloat16) * 0.25
up = torch.randn((H, K), device="cuda", dtype=torch.bfloat16) * (K**-0.5)
down = torch.randn((N, H), device="cuda", dtype=torch.bfloat16) * (H**-0.5)
up_bias = torch.randn((H,), device="cuda", dtype=torch.bfloat16) * 0.01
down_bias = torch.randn((N,), device="cuda", dtype=torch.bfloat16) * 0.01
def scale_for(tensor: torch.Tensor) -> torch.Tensor:
return (
tensor.float().abs().max() / (0.9 * base.fp8_max())
).clamp_min(1e-6).reshape(1)
x_scale = scale_for(x)
up_scale = scale_for(up)
down_scale = scale_for(down)
x_fp8 = quantize_input(x, x_scale)
up_fp8 = base.quantize_fp8(up, up_scale)
down_fp8 = base.quantize_fp8(down, down_scale)
calibrated_hidden = F.gelu(
(x_fp8.float() * x_scale) @ (up_fp8.float() * up_scale).T
+ up_bias.float(),
approximate="tanh",
)
hidden_scale = scale_for(calibrated_hidden)
return (
x,
up,
up_bias,
down,
down_bias,
x_fp8,
up_fp8,
down_fp8,
x_scale,
up_scale,
hidden_scale,
down_scale,
)
def run_shape(ops, name: str, shape, args) -> Result:
M, K, H, N = shape
(
x,
up,
up_bias,
down,
down_bias,
x_fp8,
up_fp8,
down_fp8,
x_scale,
up_scale,
hidden_scale,
down_scale,
) = make_case(M, K, H, N)
padded_m = midm_padded_rows(M)
input_fp8 = torch.empty((padded_m, K), device="cuda", dtype=base.fp8_dtype())
hidden_bf16 = torch.empty((padded_m, H), device="cuda", dtype=torch.bfloat16)
hidden_fp8 = torch.empty_like(hidden_bf16, dtype=base.fp8_dtype())
out = torch.empty((padded_m, N), device="cuda", dtype=torch.bfloat16)
exact_hidden = torch.empty_like(hidden_bf16)
exact_hidden_fp8 = torch.empty_like(hidden_fp8)
exact_out = torch.empty_like(out)
staged_hidden = torch.empty((M, H), device="cuda", dtype=torch.bfloat16)
staged_hidden_fp8 = torch.empty_like(staged_hidden, dtype=base.fp8_dtype())
staged_out = torch.empty((M, N), device="cuda", dtype=torch.bfloat16)
def flashrt_call():
return ops.bf16_fp8_gelu_mlp_bf16(
x,
up_fp8,
up_bias,
down_fp8,
down_bias,
x_scale,
up_scale,
hidden_scale,
down_scale,
input_fp8=input_fp8,
hidden_bf16=hidden_bf16,
hidden_fp8=hidden_fp8,
out=out,
pad_to=padded_m,
)
def staged_call(input_arg=x_fp8):
return ops.fp8_gelu_mlp_bf16(
input_arg,
up_fp8,
up_bias,
down_fp8,
down_bias,
x_scale,
up_scale,
hidden_scale,
down_scale,
hidden_bf16=staged_hidden,
hidden_fp8=staged_hidden_fp8,
out=staged_out,
)
def separate_quant_call():
return staged_call(quantize_input(x, x_scale))
def exact_staged_call():
return ops.fp8_gelu_mlp_bf16(
input_fp8,
up_fp8,
up_bias,
down_fp8,
down_bias,
x_scale,
up_scale,
hidden_scale,
down_scale,
hidden_bf16=exact_hidden,
hidden_fp8=exact_hidden_fp8,
out=exact_out,
)[:M]
def torch_bf16_reference():
hidden = F.gelu(F.linear(x, up, up_bias), approximate="tanh")
return F.linear(hidden, down, down_bias)
got = flashrt_call().clone()
staged = exact_staged_call().clone()
torch.cuda.synchronize()
quant_exact = bool(
torch.equal(input_fp8[:M], x_fp8)
and (padded_m == M or torch.count_nonzero(input_fp8[M:]).item() == 0)
)
staged_metrics = metrics(got, staged)
bf16_expected = torch_bf16_reference()
bf16_metrics = metrics(got, bf16_expected)
staged_compatible = quant_exact and staged_metrics["max_abs"] == 0.0
flashrt_us, eager_us = abba_time_us(
flashrt_call, torch_bf16_reference, args
)
graph = torch.cuda.CUDAGraph()
flashrt_call()
torch.cuda.synchronize()
with torch.cuda.graph(graph):
flashrt_call()
graph_us = median_time_us(graph.replay, args)
separate_us = median_time_us(separate_quant_call, args)
kernel_us = median_time_us(staged_call, args)
compile_us = None
compile_status = "not_requested"
flashrt_compile_status = "not_requested"
if args.compile_baseline:
try:
compiled = torch.compile(
torch_bf16_reference, fullgraph=True, mode="reduce-overhead"
)
compiled_out = compiled()
torch.cuda.synchronize()
compiled_metrics = metrics(compiled_out, bf16_expected)
if compiled_metrics["cosine"] < 0.9999:
compile_status = (
"mismatch: cosine="
f"{compiled_metrics['cosine']:.8f}"
)
else:
compile_us = median_time_us(compiled, args)
compile_status = "fullgraph-ok"
except Exception as exc: # noqa: BLE001
compile_status = f"failed: {type(exc).__name__}: {exc}"
try:
compiled_flashrt = torch.compile(
flashrt_call, fullgraph=True, mode="reduce-overhead"
)
compiled_got = compiled_flashrt().clone()
torch.cuda.synchronize()
flashrt_compile_status = (
"fullgraph-ok"
if metrics(compiled_got, got)["max_abs"] == 0.0
else "mismatch"
)
except Exception as exc: # noqa: BLE001
flashrt_compile_status = f"failed: {type(exc).__name__}: {exc}"
speedup_eager = eager_us / flashrt_us
speedup_separate = separate_us / flashrt_us
perf_status = (
("PASS" if speedup_eager >= 1.3 else "FAIL")
if M == 51
else "DIAGNOSTIC"
)
status = "PASS" if staged_compatible and perf_status != "FAIL" else "FAIL"
return Result(
shape=name,
M=M,
K=K,
H=H,
N=N,
flashrt_bf16_entry_us=flashrt_us,
flashrt_cuda_graph_us=graph_us,
separate_quant_us=separate_us,
fp8_kernel_only_us=kernel_us,
torch_bf16_eager_us=eager_us,
torch_bf16_compile_us=compile_us,
speedup_vs_separate_quant=speedup_separate,
speedup_vs_eager=speedup_eager,
speedup_vs_compile=compile_us / flashrt_us if compile_us else None,
compile_status=compile_status,
flashrt_compile_status=flashrt_compile_status,
input_quant_exact=quant_exact,
output_dtype=str(got.dtype),
staged_max_abs=staged_metrics["max_abs"],
staged_mean_abs=staged_metrics["mean_abs"],
staged_p99_abs=staged_metrics["p99_abs"],
staged_cosine=staged_metrics["cosine"],
bf16_max_abs=bf16_metrics["max_abs"],
bf16_mean_abs=bf16_metrics["mean_abs"],
bf16_p99_abs=bf16_metrics["p99_abs"],
bf16_cosine=bf16_metrics["cosine"],
performance_status=perf_status,
status=status,
)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--backend", choices=["source", "installed", "hub"], default="source")
parser.add_argument("--artifact")
parser.add_argument("--repo-id", default="flashrt/flashrt-fp8-ffn")
parser.add_argument("--version", type=int, default=1)
parser.add_argument("--shapes", default="all")
parser.add_argument("--warmup", type=int, default=20)
parser.add_argument("--iters", type=int, default=100)
parser.add_argument("--rounds", type=int, default=5)
parser.add_argument("--compile-baseline", action="store_true")
parser.add_argument("--output", type=Path)
args = parser.parse_args()
if not torch.cuda.is_available():
raise SystemExit("CUDA is required")
torch.manual_seed(19)
if args.backend == "source":
ops = base.load_source_ops()
elif args.backend == "installed":
ops = base.load_installed_ops(args.artifact)
else:
ops = base.load_hub_ops(args.repo_id, args.version)
names = list(SHAPES) if args.shapes == "all" else args.shapes.split(",")
unknown = [name for name in names if name not in SHAPES]
if unknown:
raise SystemExit(f"unknown shapes: {unknown}")
results = []
for name in names:
result = run_shape(ops, name, SHAPES[name], args)
results.append(result)
compile_text = (
f"{result.torch_bf16_compile_us:.3f}us"
if result.torch_bf16_compile_us is not None
else result.compile_status
)
print(
f"{result.status} {name}: flashrt={result.flashrt_bf16_entry_us:.3f}us "
f"graph={result.flashrt_cuda_graph_us:.3f}us "
f"separate={result.separate_quant_us:.3f}us "
f"kernel_only={result.fp8_kernel_only_us:.3f}us "
f"eager={result.torch_bf16_eager_us:.3f}us "
f"compile={compile_text} vs_eager={result.speedup_vs_eager:.2f}x "
f"vs_separate={result.speedup_vs_separate_quant:.2f}x "
f"staged_max={result.staged_max_abs:.6f} "
f"bf16_cos={result.bf16_cosine:.8f} "
f"perf={result.performance_status} "
f"op_compile={result.flashrt_compile_status}"
)
torch.cuda.empty_cache()
payload = {
"backend": args.backend,
"device": torch.cuda.get_device_name(0),
"torch": torch.__version__,
"warmup": args.warmup,
"iters": args.iters,
"rounds": args.rounds,
"primary_order": "A-B-B-A median",
"results": [asdict(result) for result in results],
}
if args.output:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(payload, indent=2), encoding="utf-8")
if any(result.status != "PASS" for result in results):
raise SystemExit(1)
if __name__ == "__main__":
main()