# Perform the FP8 GEMM deep_gemm.gemm_fp8_fp8_bf16_nt(lhs_input, rhs_input, output) # Verify correctness reference = lhs @ rhs.t() error = torch.abs(output - reference).max().item() relative_error = (error / torch.abs(reference).max().item()) * 100 print(f"Max absolute error: {error:.6f}") print(f"Relative error: {relative_error:.3f}%") print("✓ FP8 GEMM completed successfully!")