def analyze_gemm_performance(m, n, k, operation="forward"): # Theoretical peak performance # H100 has ~1600 TFLOPS FP8 peak ops = 2 * m * n * k peak_time = ops / 1600e12 # Theoretical minimum time # Memory bandwidth fp8_bytes = (m * k + n * k) * 1 # FP8 inputs bf16_bytes = m * n * 2 # BF16 output scale_bytes = ((m * k) // 128 + (n * k) // 128) * 4 # FP32 scales total_bytes = fp8_bytes + bf16_bytes + scale_bytes # H100 has ~3TB/s memory bandwidth bandwidth_time = total_bytes / 3e12 print(f"\n{operation.upper()} GEMM Analysis (M={m}, N={n}, K={k}):") print(f"Operations: {ops/1e9:.1f} GigaOps") print(f"Compute bound time: {peak_time*1000:.2f}ms") print(f"Memory bound time: {bandwidth_time*1000:.2f}ms") print(f"Bottleneck: {'Compute' if peak_time > bandwidth_time else 'Memory'}") # Analyze our configurations analyze_gemm_performance(128, 256, 512, "forward") analyze_gemm_performance(256, 512, 1024, "weight_grad")