togethercomputer/ParallelKernelBench_Problems
ParallelKernelBench (benchmark) Reference problems for ParallelKernelBench: a benchmark for LLM-generated multi-GPU CUDA kernels. This dataset contains 87 reference implementations in reference/ and the input tensor specification in utils/input_output_tensors.py. Inputs are deterministic — reproduce them with create_input_tensor(rank, world_size, problem_id, base_shape, dtype, trial) from that file; you do not need stored .pt files. Files Path Description… See the full description on the dataset page: https://huggingface.co/datasets/togethercomputer/ParallelKernelBench_Problems.
0177
1# Expert-parallel (EP) fused MoE forward — BALANCED EP.2#3# Same kernel as problem 31 (router -> permute -> all_to_all dispatch -> per-expert4# SiLU MLP -> all_to_all combine -> unpermute), but the harness sets5# num_experts == world_size, i.e. exactly one expert per rank. With uniform routing6# this gives the balanced all_to_all dispatch pattern (each rank sends/receives a7# roughly equal token count).8from typing import List, Optional, Tuple, Union9 10import torch11import torch.distributed as dist12 13 14class _AllToAll(torch.autograd.Function):15 @staticmethod16 def forward(ctx, group, input, output_split_sizes, input_split_sizes):17 ctx.group = group18 ctx.output_split_sizes = output_split_sizes19 ctx.input_split_sizes = input_split_sizes20 if dist.get_world_size(group=group) == 1:21 return input.contiguous()22 input = input.contiguous()23 if output_split_sizes is None:24 output = torch.empty_like(input)25 else:26 output = torch.empty(27 size=(sum(output_split_sizes), input.size(1)),28 dtype=input.dtype,29 device=input.device,30 )31 dist.all_to_all_single(32 output,33 input,34 output_split_sizes=output_split_sizes,35 input_split_sizes=input_split_sizes,36 group=group,37 )38 return output39 40 @staticmethod41 def backward(ctx, grad_output):42 return (43 None,44 _AllToAll.apply(45 ctx.group, grad_output, ctx.input_split_sizes, ctx.output_split_sizes46 ),47 None,48 None,49 )50 51 52def _all_to_all(53 group: dist.ProcessGroup,54 input: torch.Tensor,55 output_split_sizes: Optional[List[int]],56 input_split_sizes: Optional[List[int]],57) -> torch.Tensor:58 return _AllToAll.apply(group, input, output_split_sizes, input_split_sizes)59 60 61def _preprocess(62 expert_mask: torch.Tensor,63 num_experts: int,64 ep_group: dist.ProcessGroup,65) -> Tuple[List[int], List[int], torch.Tensor, torch.Tensor]:66 ep_size = ep_group.size()67 num_local_experts = num_experts // ep_size68 rank = dist.get_rank(ep_group)69 num_local_tokens_per_expert = expert_mask.sum(dim=(1, 2))70 input_splits = (71 num_local_tokens_per_expert.reshape(ep_size, num_local_experts).sum(dim=1).tolist()72 )73 num_local_tokens_per_expert_flat = num_local_tokens_per_expert.contiguous().view(-1)74 output_size = ep_size * num_local_tokens_per_expert_flat.numel()75 num_global_tokens_per_expert_flat = torch.empty(76 output_size,77 dtype=num_local_tokens_per_expert.dtype,78 device=num_local_tokens_per_expert.device,79 )80 dist.all_gather_into_tensor(81 num_global_tokens_per_expert_flat, num_local_tokens_per_expert_flat, group=ep_group82 )83 num_global_tokens_per_expert = num_global_tokens_per_expert_flat.view(84 ep_size, num_local_tokens_per_expert.size(0)85 )86 start_idx, end_idx = rank * num_local_experts, (rank + 1) * num_local_experts87 num_global_tokens_per_local_expert = num_global_tokens_per_expert[88 :, start_idx:end_idx89 ].contiguous()90 output_splits = num_global_tokens_per_local_expert.sum(dim=1).tolist()91 num_global_sum_tokens_per_local_expert = num_global_tokens_per_local_expert.sum(92 dim=093 ).to(torch.device("cpu"), non_blocking=True)94 num_global_tokens_per_local_expert = num_global_tokens_per_local_expert.view(95 -1, num_local_experts96 ).to(torch.device("cpu"), non_blocking=True)97 return (98 input_splits,99 output_splits,100 num_global_tokens_per_local_expert,101 num_global_sum_tokens_per_local_expert,102 )103 104 105def _permute(106 tokens: torch.Tensor, routing_map: torch.Tensor107) -> Tuple[torch.Tensor, torch.Tensor]:108 num_tokens, _ = tokens.shape109 num_experts = routing_map.shape[0]110 routing_map = routing_map.bool()111 token_indices = (112 torch.arange(num_tokens, device=routing_map.device)113 .unsqueeze(0)114 .expand(num_experts, -1)115 )116 sorted_indices = token_indices.masked_select(routing_map)117 permuted_input = tokens.index_select(0, sorted_indices)118 return permuted_input, sorted_indices119 120 121def _sort_chunks_by_idxs(122 input: torch.Tensor,123 split_sizes: Union[torch.Tensor, List[int]],124 sorted_idxs: List[int],125) -> torch.Tensor:126 if isinstance(split_sizes, torch.Tensor):127 split_sizes = split_sizes.tolist()128 chunks = torch.split(input, split_sizes, dim=0)129 return torch.cat([chunks[i] for i in sorted_idxs], dim=0)130 131 132def _generate_weights_idx(133 routing_weights: torch.Tensor,134 selected_experts: torch.Tensor,135 num_experts: int,136) -> torch.Tensor:137 num_tokens, topk = routing_weights.shape138 weights_idx = torch.zeros(139 (num_tokens, num_experts),140 dtype=routing_weights.dtype,141 device=routing_weights.device,142 )143 weights_idx.scatter_add_(1, selected_experts, routing_weights)144 return weights_idx145 146 147def _unpermute(148 tokens: torch.Tensor,149 routing_weights: torch.Tensor,150 hidden_states_shape: torch.Size,151 permutation_mapping: torch.Tensor,152 routing_map: torch.Tensor,153) -> torch.Tensor:154 tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())155 tokens = tokens * tokens_weight.unsqueeze(-1)156 hidden_dim = hidden_states_shape[-1]157 unpermuted_tokens = torch.zeros(158 hidden_states_shape, device=tokens.device, dtype=tokens.dtype159 )160 expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)161 unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)162 return unpermuted_tokens163 164 165def token_pre_all2all(166 hidden_states: torch.Tensor,167 expert_mask: torch.Tensor,168 num_experts: int,169 input_splits: List[int],170 output_splits: List[int],171 num_global_tokens_per_local_expert: torch.Tensor,172 group: Optional[dist.ProcessGroup] = None,173) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Size]:174 group = group or dist.group.WORLD175 hidden_dim = hidden_states.size(-1)176 hidden_states = hidden_states.reshape(-1, hidden_dim)177 org_hidden_states_shape = hidden_states.shape178 routing_map = expert_mask.sum(dim=1)179 180 local_permuted_hidden_states, local_input_permutation_mapping = _permute(181 hidden_states, routing_map182 )183 expected_tokens = sum(input_splits)184 actual_tokens = local_permuted_hidden_states.shape[0]185 if expected_tokens != actual_tokens:186 raise RuntimeError(187 f"EP split mismatch: input_splits sum ({expected_tokens}) != "188 f"permuted tokens ({actual_tokens})"189 )190 191 global_permuted_hidden_states = _all_to_all(192 group, local_permuted_hidden_states, output_splits, input_splits193 )194 num_local_experts = num_experts // dist.get_world_size(group)195 permute_order = (196 torch.arange(num_experts).reshape(-1, num_local_experts).T.ravel().tolist()197 )198 split_sizes = num_global_tokens_per_local_expert.ravel().tolist()199 global_permuted_hidden_states = _sort_chunks_by_idxs(200 global_permuted_hidden_states, split_sizes, permute_order201 )202 return (203 global_permuted_hidden_states,204 routing_map,205 local_input_permutation_mapping,206 org_hidden_states_shape,207 )208 209 210def tokens_post_all2all(211 expert_outputs: torch.Tensor,212 routing_weights: torch.Tensor,213 selected_experts: torch.Tensor,214 num_experts: int,215 input_splits: List[int],216 output_splits: List[int],217 num_global_tokens_per_local_expert: torch.Tensor,218 routing_map: torch.Tensor,219 local_input_permutation_mapping: torch.Tensor,220 org_hidden_states_shape: torch.Size,221 group: Optional[dist.ProcessGroup] = None,222) -> torch.Tensor:223 group = group or dist.group.WORLD224 num_local_experts = num_experts // dist.get_world_size(group)225 unpermute_order = (226 torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()227 )228 split_sizes = num_global_tokens_per_local_expert.T.ravel().tolist()229 expert_outputs = _sort_chunks_by_idxs(230 expert_outputs, split_sizes, unpermute_order231 )232 unpermute_outputs = _all_to_all(group, expert_outputs, input_splits, output_splits)233 weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)234 unpermute_outputs = _unpermute(235 unpermute_outputs,236 weights_idx,237 org_hidden_states_shape,238 local_input_permutation_mapping,239 routing_map,240 )241 return unpermute_outputs242 243 244def expert_forward(245 x: torch.Tensor,246 gate_proj: torch.nn.Linear,247 up_proj: torch.nn.Linear,248 down_proj: torch.nn.Linear,249) -> torch.Tensor:250 gate = torch.nn.functional.silu(gate_proj(x))251 up = up_proj(x)252 return down_proj(gate * up)253 254 255def solution(256 hidden_states: torch.Tensor,257 gate_weight: torch.Tensor,258 gate_bias: Optional[torch.Tensor],259 gate_proj: torch.nn.Linear,260 up_proj: torch.nn.Linear,261 down_proj: torch.nn.Linear,262 num_experts: int,263 top_k: int,264 group: Optional[dist.ProcessGroup] = None,265) -> torch.Tensor:266 group = group or dist.group.WORLD267 hidden_dim = hidden_states.size(-1)268 num_tokens = hidden_states.reshape(-1, hidden_dim).size(0)269 270 router_logits = torch.nn.functional.linear(271 hidden_states.reshape(-1, hidden_dim), gate_weight, gate_bias272 )273 routing_weights, selected_experts = torch.topk(274 torch.softmax(router_logits, dim=-1), top_k, dim=-1275 )276 expert_mask = torch.nn.functional.one_hot(277 selected_experts, num_classes=num_experts278 ).permute(2, 1, 0)279 280 input_splits, output_splits, num_global_tokens_per_local_expert, _ = _preprocess(281 expert_mask, num_experts, group282 )283 284 (285 global_permuted_hidden_states,286 routing_map,287 local_input_permutation_mapping,288 org_hidden_states_shape,289 ) = token_pre_all2all(290 hidden_states,291 expert_mask,292 num_experts,293 input_splits,294 output_splits,295 num_global_tokens_per_local_expert,296 group,297 )298 299 expert_outputs = expert_forward(300 global_permuted_hidden_states, gate_proj, up_proj, down_proj301 )302 303 out = tokens_post_all2all(304 expert_outputs,305 routing_weights,306 selected_experts,307 num_experts,308 input_splits,309 output_splits,310 num_global_tokens_per_local_expert,311 routing_map,312 local_input_permutation_mapping,313 org_hidden_states_shape,314 group,315 )316 return out317 