# Perform grouped GEMM deep_gemm.m_grouped_gemm_fp8_fp8_bf16_nt_contiguous( lhs_grouped_input, rhs_grouped_input, output_grouped, m_indices ) # Verify by computing reference reference_grouped = torch.zeros_like(output_grouped) start = 0 for expert_id, aligned_tokens in enumerate(aligned_tokens): end = start + aligned_tokens reference_grouped[start:end] = lhs_grouped[start:end] @ rhs_grouped[expert_id].t() start = end # Mask out padding tokens for comparison valid_mask = (m_indices != -1).unsqueeze(1) output_masked = torch.where(valid_mask, output_grouped, torch.zeros_like(output_grouped)) reference_masked = torch.where(valid_mask, reference_grouped, torch.zeros_like(reference_grouped)) error = torch.abs(output_masked - reference_masked).max().item() print(f"Grouped GEMM error: {error:.6f}") print("✓ Grouped GEMM completed successfully!")