def cast_to_fp8_per_block(x: torch.Tensor): """Convert tensor to FP8 with per-block scaling (128x128 blocks)""" m, n = x.shape # Pad to 128x128 blocks padded_m = ((m + 127) // 128) * 128 padded_n = ((n + 127) // 128) * 128 x_padded = torch.zeros((padded_m, padded_n), dtype=x.dtype, device=x.device) x_padded[:m, :n] = x # Reshape into 128x128 blocks x_view = x_padded.view(-1, 128, x_padded.size(1) // 128, 128) # Find max per block x_amax = x_view.abs().float().amax(dim=(1, 3), keepdim=True).clamp(1e-4) # Scale to FP8 x_scaled = (x_view * (448.0 / x_amax)).to(torch.float8_e4m3fn) scale_factors = (x_amax / 448.0).view(x_view.size(0), x_view.size(2)) return x_scaled.view_as(x_padded)[:m, :n], scale_factors # Convert RHS with block scaling rhs_fp8, rhs_scales = cast_to_fp8_per_block(rhs) print(f"RHS FP8 shape: {rhs_fp8.shape}") print(f"RHS scales shape: {rhs_scales.shape}")