# Create inputs for grouped GEMM lhs_grouped = torch.randn((total_aligned, k), device='cuda', dtype=torch.bfloat16) rhs_grouped = torch.randn((num_experts, n, k), device='cuda', dtype=torch.bfloat16) output_grouped = torch.empty((total_aligned, n), device='cuda', dtype=torch.bfloat16) # Create mapping tensor m_indices = torch.empty(total_aligned, device='cuda', dtype=torch.int32) start = 0 for expert_id, (orig_tokens, aligned_tokens) in enumerate(zip(tokens_per_expert, aligned_tokens)): # Real tokens get expert ID m_indices[start:start + orig_tokens] = expert_id # Padding tokens get -1 (ignored) m_indices[start + orig_tokens:start + aligned_tokens] = -1 start += aligned_tokens print(f"Mapping tensor shape: {m_indices.shape}") print(f"Expert assignments: {m_indices[:20]}") # First 20 tokens