import time def benchmark_gemm(func, *args, num_runs=10): # Warmup for _ in range(3): func(*args) torch.cuda.synchronize() # Timing start = time.time() for _ in range(num_runs): func(*args) torch.cuda.synchronize() return (time.time() - start) / num_runs # Benchmark both versions fp8_time = benchmark_gemm(deep_gemm.gemm_fp8_fp8_bf16_nt, lhs_input, rhs_input, output) bf16_time = benchmark_gemm(lambda x, y, out: out.copy_(x @ y.t()), lhs, rhs, reference) # Calculate throughput (TFLOPS) ops = 2 * m * n * k # Multiply-accumulate operations fp8_tflops = ops / fp8_time / 1e12 bf16_tflops = ops / bf16_time / 1e12 print(f"FP8 GEMM: {fp8_time*1000:.2f}ms ({fp8_tflops:.1f} TFLOPS)") print(f"BF16 GEMM: {bf16_time*1000:.2f}ms ({bf16_tflops:.1f} TFLOPS)") print(f"Speedup: {bf16_time/fp8_time:.1f}x")