def cast_to_fp8_per_token(x: torch.Tensor): """Convert tensor to FP8 with per-token (per-row) scaling""" assert x.dim() == 2 m, n = x.shape # Pad to 128-element boundaries (FP8 requirement) pad_size = (128 - (n % 128)) % 128 if pad_size > 0: x = torch.nn.functional.pad(x, (0, pad_size), value=0) # Reshape for scaling calculation x_view = x.view(m, -1, 128) # [m, n/128, 128] # Find max absolute value per 128-element block x_amax = x_view.abs().float().amax(dim=2).clamp(1e-4) # [m, n/128] # Scale to FP8 range (448.0 is max representable value) fp8_data = (x_view * (448.0 / x_amax.unsqueeze(2))).to(torch.float8_e4m3fn) scale_factors = (x_amax / 448.0) return fp8_data.view(m, -1)[:, :n], scale_factors # Convert our matrices lhs_fp8, lhs_scales = cast_to_fp8_per_token(lhs) print(f"Original: {lhs.dtype}, Converted: {lhs_fp8.dtype}") print(f"Scale factors shape: {lhs_scales.shape}")