def validate_fp8_conversion(original, fp8_data, scales): """Check if FP8 conversion preserves data accurately""" # Reconstruct original from FP8 if fp8_data.dim() == 2: # Per-token scaling m, n = fp8_data.shape fp8_view = fp8_data.view(m, -1, 128) scales_expanded = scales.unsqueeze(2) reconstructed = fp8_view.float() * scales_expanded reconstructed = reconstructed.view(m, -1)[:, :original.shape[1]] # Compare abs_error = torch.abs(original.float() - reconstructed).max().item() rel_error = abs_error / torch.abs(original.float()).max().item() print(f"FP8 conversion error: {abs_error:.6f} ({rel_error*100:.3f}%)") return abs_error < 1e-2 # Reasonable threshold for FP8 # Validate our conversions lhs_valid = validate_fp8_conversion(lhs, lhs_fp8, lhs_scales) rhs_valid = validate_fp8_conversion(rhs, rhs_fp8[0], rhs_scales[0]) print(f"LHS conversion valid: {lhs_valid}") print(f"RHS conversion valid: {rhs_valid}")