def estimate_memory_usage(shapes, operation="gemm"): """Estimate GPU memory usage for DeepGEMM operations""" m, n, k = shapes # Input tensors lhs_fp8 = m * k * 1 # FP8 lhs_scales = m * ((k + 127) // 128) * 4 # FP32 rhs_fp8 = n * k * 1 # FP8 rhs_scales = ((n + 127) // 128) * ((k + 127) // 128) * 4 # FP32 # Output if operation == "gemm": output = m * n * 2 # BF16 else: # weight_grad output = m * n * 4 # FP32 # Temporary workspace (estimated) workspace = max(m, n) * 1024 * 4 # Conservative estimate total = lhs_fp8 + lhs_scales + rhs_fp8 + rhs_scales + output + workspace print(f"Memory usage for {shapes}:") print(f" Inputs: {(lhs_fp8 + lhs_scales + rhs_fp8 + rhs_scales) / 1024**2:.1f} MB") print(f" Output: {output / 1024**2:.1f} MB") print(f" Workspace: {workspace / 1024**2:.1f} MB") print(f" Total: {total / 1024**2:.1f} MB") return total # Estimate for different problem sizes estimate_memory_usage((1024, 2048, 4096), "gemm") estimate_memory_usage((2048, 4096, 8192), "weight_grad")