# Convert LHS (same as before) lhs_grouped_fp8, lhs_grouped_scales = cast_to_fp8_per_token(lhs_grouped) lhs_grouped_scales = get_col_major_tma_aligned_tensor(lhs_grouped_scales) # Convert each expert's RHS separately rhs_grouped_fp8 = torch.empty_like(rhs_grouped, dtype=torch.float8_e4m3fn) rhs_grouped_scales = torch.empty((num_experts, (n + 127) // 128, (k + 127) // 128), device='cuda', dtype=torch.float32) for expert_id in range(num_experts): rhs_grouped_fp8[expert_id], rhs_grouped_scales[expert_id] = cast_to_fp8_per_block(rhs_grouped[expert_id]) # Package inputs lhs_grouped_input = (lhs_grouped_fp8, lhs_grouped_scales) rhs_grouped_input = (rhs_grouped_fp8, rhs_grouped_scales) print("✓ Grouped data converted to FP8")