Download benchmarks/benchmark_bf16_entry.py from flashrt/flashrt-fp8-ffn: direct link, hf CLI and curl.
- Browser
- Download file 13.6 kB
-
https://huggingface.co/flashrt/flashrt-fp8-ffn/resolve/main/benchmarks/benchmark_bf16_entry.py
- Command line
-
hf download hf://flashrt/flashrt-fp8-ffn/benchmarks/benchmark_bf16_entry.py
-
curl -L -o benchmark_bf16_entry.py https://huggingface.co/flashrt/flashrt-fp8-ffn/resolve/main/benchmarks/benchmark_bf16_entry.py
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), | |
| } | |
| 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() | |