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 — BASE case.2#3# This is the generic fused MoE forward pass: router (softmax + top-k) -> token4# permutation -> all_to_all dispatch -> per-expert SiLU MLP -> all_to_all combine ->5# weighted unpermute. Here the expert count is fixed (num_experts = 8) regardless of6# world size, so the EP load pattern depends on the launch configuration.7from typing import List, Optional, Tuple, Union8 9import torch10import torch.distributed as dist11 12 13class _AllToAll(torch.autograd.Function):14 @staticmethod15 def forward(ctx, group, input, output_split_sizes, input_split_sizes):16 ctx.group = group17 ctx.output_split_sizes = output_split_sizes18 ctx.input_split_sizes = input_split_sizes19 if dist.get_world_size(group=group) == 1:20 return input.contiguous()21 input = input.contiguous()22 if output_split_sizes is None:23 output = torch.empty_like(input)24 else:25 output = torch.empty(26 size=(sum(output_split_sizes), input.size(1)),27 dtype=input.dtype,28 device=input.device,29 )30 dist.all_to_all_single(31 output,32 input,33 output_split_sizes=output_split_sizes,34 input_split_sizes=input_split_sizes,35 group=group,36 )37 return output38 39 @staticmethod40 def backward(ctx, grad_output):41 return (42 None,43 _AllToAll.apply(44 ctx.group, grad_output, ctx.input_split_sizes, ctx.output_split_sizes45 ),46 None,47 None,48 )49 50 51def _all_to_all(52 group: dist.ProcessGroup,53 input: torch.Tensor,54 output_split_sizes: Optional[List[int]],55 input_split_sizes: Optional[List[int]],56) -> torch.Tensor:57 return _AllToAll.apply(group, input, output_split_sizes, input_split_sizes)58 59 60def _preprocess(61 expert_mask: torch.Tensor,62 num_experts: int,63 ep_group: dist.ProcessGroup,64) -> Tuple[List[int], List[int], torch.Tensor, torch.Tensor]:65 ep_size = ep_group.size()66 num_local_experts = num_experts // ep_size67 rank = dist.get_rank(ep_group)68 num_local_tokens_per_expert = expert_mask.sum(dim=(1, 2))69 input_splits = (70 num_local_tokens_per_expert.reshape(ep_size, num_local_experts).sum(dim=1).tolist()71 )72 num_local_tokens_per_expert_flat = num_local_tokens_per_expert.contiguous().view(-1)73 output_size = ep_size * num_local_tokens_per_expert_flat.numel()74 num_global_tokens_per_expert_flat = torch.empty(75 output_size,76 dtype=num_local_tokens_per_expert.dtype,77 device=num_local_tokens_per_expert.device,78 )79 dist.all_gather_into_tensor(80 num_global_tokens_per_expert_flat, num_local_tokens_per_expert_flat, group=ep_group81 )82 num_global_tokens_per_expert = num_global_tokens_per_expert_flat.view(83 ep_size, num_local_tokens_per_expert.size(0)84 )85 start_idx, end_idx = rank * num_local_experts, (rank + 1) * num_local_experts86 num_global_tokens_per_local_expert = num_global_tokens_per_expert[87 :, start_idx:end_idx88 ].contiguous()89 output_splits = num_global_tokens_per_local_expert.sum(dim=1).tolist()90 num_global_sum_tokens_per_local_expert = num_global_tokens_per_local_expert.sum(91 dim=092 ).to(torch.device("cpu"), non_blocking=True)93 num_global_tokens_per_local_expert = num_global_tokens_per_local_expert.view(94 -1, num_local_experts95 ).to(torch.device("cpu"), non_blocking=True)96 return (97 input_splits,98 output_splits,99 num_global_tokens_per_local_expert,100 num_global_sum_tokens_per_local_expert,101 )102 103 104def _permute(105 tokens: torch.Tensor, routing_map: torch.Tensor106) -> Tuple[torch.Tensor, torch.Tensor]:107 num_tokens, _ = tokens.shape108 num_experts = routing_map.shape[0]109 routing_map = routing_map.bool()110 token_indices = (111 torch.arange(num_tokens, device=routing_map.device)112 .unsqueeze(0)113 .expand(num_experts, -1)114 )115 sorted_indices = token_indices.masked_select(routing_map)116 permuted_input = tokens.index_select(0, sorted_indices)117 return permuted_input, sorted_indices118 119 120def _sort_chunks_by_idxs(121 input: torch.Tensor,122 split_sizes: Union[torch.Tensor, List[int]],123 sorted_idxs: List[int],124) -> torch.Tensor:125 if isinstance(split_sizes, torch.Tensor):126 split_sizes = split_sizes.tolist()127 chunks = torch.split(input, split_sizes, dim=0)128 return torch.cat([chunks[i] for i in sorted_idxs], dim=0)129 130 131def _generate_weights_idx(132 routing_weights: torch.Tensor,133 selected_experts: torch.Tensor,134 num_experts: int,135) -> torch.Tensor:136 num_tokens, topk = routing_weights.shape137 weights_idx = torch.zeros(138 (num_tokens, num_experts),139 dtype=routing_weights.dtype,140 device=routing_weights.device,141 )142 weights_idx.scatter_add_(1, selected_experts, routing_weights)143 return weights_idx144 145 146def _unpermute(147 tokens: torch.Tensor,148 routing_weights: torch.Tensor,149 hidden_states_shape: torch.Size,150 permutation_mapping: torch.Tensor,151 routing_map: torch.Tensor,152) -> torch.Tensor:153 tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())154 tokens = tokens * tokens_weight.unsqueeze(-1)155 hidden_dim = hidden_states_shape[-1]156 unpermuted_tokens = torch.zeros(157 hidden_states_shape, device=tokens.device, dtype=tokens.dtype158 )159 expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)160 unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)161 return unpermuted_tokens162 163 164def token_pre_all2all(165 hidden_states: torch.Tensor,166 expert_mask: torch.Tensor,167 num_experts: int,168 input_splits: List[int],169 output_splits: List[int],170 num_global_tokens_per_local_expert: torch.Tensor,171 group: Optional[dist.ProcessGroup] = None,172) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Size]:173 group = group or dist.group.WORLD174 hidden_dim = hidden_states.size(-1)175 hidden_states = hidden_states.reshape(-1, hidden_dim)176 org_hidden_states_shape = hidden_states.shape177 routing_map = expert_mask.sum(dim=1)178 179 local_permuted_hidden_states, local_input_permutation_mapping = _permute(180 hidden_states, routing_map181 )182 expected_tokens = sum(input_splits)183 actual_tokens = local_permuted_hidden_states.shape[0]184 if expected_tokens != actual_tokens:185 raise RuntimeError(186 f"EP split mismatch: input_splits sum ({expected_tokens}) != "187 f"permuted tokens ({actual_tokens})"188 )189 190 global_permuted_hidden_states = _all_to_all(191 group, local_permuted_hidden_states, output_splits, input_splits192 )193 num_local_experts = num_experts // dist.get_world_size(group)194 permute_order = (195 torch.arange(num_experts).reshape(-1, num_local_experts).T.ravel().tolist()196 )197 split_sizes = num_global_tokens_per_local_expert.ravel().tolist()198 global_permuted_hidden_states = _sort_chunks_by_idxs(199 global_permuted_hidden_states, split_sizes, permute_order200 )201 return (202 global_permuted_hidden_states,203 routing_map,204 local_input_permutation_mapping,205 org_hidden_states_shape,206 )207 208 209def tokens_post_all2all(210 expert_outputs: torch.Tensor,211 routing_weights: torch.Tensor,212 selected_experts: torch.Tensor,213 num_experts: int,214 input_splits: List[int],215 output_splits: List[int],216 num_global_tokens_per_local_expert: torch.Tensor,217 routing_map: torch.Tensor,218 local_input_permutation_mapping: torch.Tensor,219 org_hidden_states_shape: torch.Size,220 group: Optional[dist.ProcessGroup] = None,221) -> torch.Tensor:222 group = group or dist.group.WORLD223 num_local_experts = num_experts // dist.get_world_size(group)224 unpermute_order = (225 torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()226 )227 split_sizes = num_global_tokens_per_local_expert.T.ravel().tolist()228 expert_outputs = _sort_chunks_by_idxs(229 expert_outputs, split_sizes, unpermute_order230 )231 unpermute_outputs = _all_to_all(group, expert_outputs, input_splits, output_splits)232 weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)233 unpermute_outputs = _unpermute(234 unpermute_outputs,235 weights_idx,236 org_hidden_states_shape,237 local_input_permutation_mapping,238 routing_map,239 )240 return unpermute_outputs241 242 243def expert_forward(244 x: torch.Tensor,245 gate_proj: torch.nn.Linear,246 up_proj: torch.nn.Linear,247 down_proj: torch.nn.Linear,248) -> torch.Tensor:249 gate = torch.nn.functional.silu(gate_proj(x))250 up = up_proj(x)251 return down_proj(gate * up)252 253 254def solution(255 hidden_states: torch.Tensor,256 gate_weight: torch.Tensor,257 gate_bias: Optional[torch.Tensor],258 gate_proj: torch.nn.Linear,259 up_proj: torch.nn.Linear,260 down_proj: torch.nn.Linear,261 num_experts: int,262 top_k: int,263 group: Optional[dist.ProcessGroup] = None,264) -> torch.Tensor:265 group = group or dist.group.WORLD266 hidden_dim = hidden_states.size(-1)267 num_tokens = hidden_states.reshape(-1, hidden_dim).size(0)268 269 router_logits = torch.nn.functional.linear(270 hidden_states.reshape(-1, hidden_dim), gate_weight, gate_bias271 )272 routing_weights, selected_experts = torch.topk(273 torch.softmax(router_logits, dim=-1), top_k, dim=-1274 )275 expert_mask = torch.nn.functional.one_hot(276 selected_experts, num_classes=num_experts277 ).permute(2, 1, 0)278 279 input_splits, output_splits, num_global_tokens_per_local_expert, _ = _preprocess(280 expert_mask, num_experts, group281 )282 283 (284 global_permuted_hidden_states,285 routing_map,286 local_input_permutation_mapping,287 org_hidden_states_shape,288 ) = token_pre_all2all(289 hidden_states,290 expert_mask,291 num_experts,292 input_splits,293 output_splits,294 num_global_tokens_per_local_expert,295 group,296 )297 298 expert_outputs = expert_forward(299 global_permuted_hidden_states, gate_proj, up_proj, down_proj300 )301 302 out = tokens_post_all2all(303 expert_outputs,304 routing_weights,305 selected_experts,306 num_experts,307 input_splits,308 output_splits,309 num_global_tokens_per_local_expert,310 routing_map,311 local_input_permutation_mapping,312 org_hidden_states_shape,313 group,314 )315 return out316 