willychan21/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. Files Path Description data/problems.parquet One row per problem (tabular access) reference/*.py Reference solution() implementations utils/input_output_tensors.py Input/output tensor… See the full description on the dataset page: https://huggingface.co/datasets/willychan21/ParallelKernelBench_Problems.
0263
1from typing import List, Optional, Tuple, Union2 3import torch4import torch.distributed as dist5 6_EP_SUBGROUP_CACHE: dict[tuple[int, int], None | list] = {}7 8 9def _resolve_ep_group_for_narrow_moe(num_experts: int) -> dist.ProcessGroup:10 if not dist.is_initialized():11 raise RuntimeError("torch.distributed must be initialized")12 ws = dist.get_world_size()13 rank = dist.get_rank()14 key = (ws, num_experts)15 if key not in _EP_SUBGROUP_CACHE:16 if num_experts >= ws:17 _EP_SUBGROUP_CACHE[key] = None18 elif ws % num_experts != 0:19 raise ValueError(20 f"narrow EP requires world_size ({ws}) % num_experts ({num_experts}) == 0"21 )22 else:23 groups: list = []24 for r in range(ws // num_experts):25 ranks = list(range(r * num_experts, (r + 1) * num_experts))26 groups.append(dist.new_group(ranks))27 _EP_SUBGROUP_CACHE[key] = groups28 entry = _EP_SUBGROUP_CACHE[key]29 if entry is None:30 return dist.group.WORLD31 return entry[rank // num_experts]32 33 34class _AllToAll(torch.autograd.Function):35 @staticmethod36 def forward(ctx, group, input, output_split_sizes, input_split_sizes):37 ctx.group = group38 ctx.output_split_sizes = output_split_sizes39 ctx.input_split_sizes = input_split_sizes40 if dist.get_world_size(group=group) == 1:41 return input.contiguous()42 input = input.contiguous()43 if output_split_sizes is None:44 output = torch.empty_like(input)45 else:46 output = torch.empty(47 size=(sum(output_split_sizes), input.size(1)),48 dtype=input.dtype,49 device=input.device,50 )51 dist.all_to_all_single(52 output,53 input,54 output_split_sizes=output_split_sizes,55 input_split_sizes=input_split_sizes,56 group=group,57 )58 return output59 60 @staticmethod61 def backward(ctx, grad_output):62 return (63 None,64 _AllToAll.apply(65 ctx.group, grad_output, ctx.input_split_sizes, ctx.output_split_sizes66 ),67 None,68 None,69 )70 71 72def _all_to_all(73 group: dist.ProcessGroup,74 input: torch.Tensor,75 output_split_sizes: Optional[List[int]],76 input_split_sizes: Optional[List[int]],77) -> torch.Tensor:78 return _AllToAll.apply(group, input, output_split_sizes, input_split_sizes)79 80 81def _preprocess(82 expert_mask: torch.Tensor,83 num_experts: int,84 ep_group: dist.ProcessGroup,85) -> Tuple[List[int], List[int], torch.Tensor, torch.Tensor]:86 ep_size = ep_group.size()87 num_local_experts = num_experts // ep_size88 rank = dist.get_rank(ep_group)89 num_local_tokens_per_expert = expert_mask.sum(dim=(1, 2))90 input_splits = (91 num_local_tokens_per_expert.reshape(ep_size, num_local_experts).sum(dim=1).tolist()92 )93 num_local_tokens_per_expert_flat = num_local_tokens_per_expert.contiguous().view(-1)94 output_size = ep_size * num_local_tokens_per_expert_flat.numel()95 num_global_tokens_per_expert_flat = torch.empty(96 output_size,97 dtype=num_local_tokens_per_expert.dtype,98 device=num_local_tokens_per_expert.device,99 )100 dist.all_gather_into_tensor(101 num_global_tokens_per_expert_flat, num_local_tokens_per_expert_flat, group=ep_group102 )103 num_global_tokens_per_expert = num_global_tokens_per_expert_flat.view(104 ep_size, num_local_tokens_per_expert.size(0)105 )106 start_idx, end_idx = rank * num_local_experts, (rank + 1) * num_local_experts107 num_global_tokens_per_local_expert = num_global_tokens_per_expert[108 :, start_idx:end_idx109 ].contiguous()110 output_splits = num_global_tokens_per_local_expert.sum(dim=1).tolist()111 num_global_sum_tokens_per_local_expert = num_global_tokens_per_local_expert.sum(112 dim=0113 ).to(torch.device("cpu"), non_blocking=True)114 num_global_tokens_per_local_expert = num_global_tokens_per_local_expert.view(115 -1, num_local_experts116 ).to(torch.device("cpu"), non_blocking=True)117 return (118 input_splits,119 output_splits,120 num_global_tokens_per_local_expert,121 num_global_sum_tokens_per_local_expert,122 )123 124 125def _permute(126 tokens: torch.Tensor, routing_map: torch.Tensor127) -> Tuple[torch.Tensor, torch.Tensor]:128 num_tokens, _ = tokens.shape129 num_experts = routing_map.shape[0]130 routing_map = routing_map.bool()131 token_indices = (132 torch.arange(num_tokens, device=routing_map.device)133 .unsqueeze(0)134 .expand(num_experts, -1)135 )136 sorted_indices = token_indices.masked_select(routing_map)137 permuted_input = tokens.index_select(0, sorted_indices)138 return permuted_input, sorted_indices139 140 141def _sort_chunks_by_idxs(142 input: torch.Tensor,143 split_sizes: Union[torch.Tensor, List[int]],144 sorted_idxs: List[int],145) -> torch.Tensor:146 if isinstance(split_sizes, torch.Tensor):147 split_sizes = split_sizes.tolist()148 chunks = torch.split(input, split_sizes, dim=0)149 return torch.cat([chunks[i] for i in sorted_idxs], dim=0)150 151 152def _generate_weights_idx(153 routing_weights: torch.Tensor,154 selected_experts: torch.Tensor,155 num_experts: int,156) -> torch.Tensor:157 num_tokens, topk = routing_weights.shape158 weights_idx = torch.zeros(159 (num_tokens, num_experts),160 dtype=routing_weights.dtype,161 device=routing_weights.device,162 )163 weights_idx.scatter_add_(1, selected_experts, routing_weights)164 return weights_idx165 166 167def _unpermute(168 tokens: torch.Tensor,169 routing_weights: torch.Tensor,170 hidden_states_shape: torch.Size,171 permutation_mapping: torch.Tensor,172 routing_map: torch.Tensor,173) -> torch.Tensor:174 tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())175 tokens = tokens * tokens_weight.unsqueeze(-1)176 hidden_dim = hidden_states_shape[-1]177 unpermuted_tokens = torch.zeros(178 hidden_states_shape, device=tokens.device, dtype=tokens.dtype179 )180 expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)181 unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)182 return unpermuted_tokens183 184 185def token_pre_all2all(186 hidden_states: torch.Tensor,187 expert_mask: torch.Tensor,188 num_experts: int,189 input_splits: List[int],190 output_splits: List[int],191 num_global_tokens_per_local_expert: torch.Tensor,192 group: Optional[dist.ProcessGroup] = None,193) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Size]:194 group = group or dist.group.WORLD195 hidden_dim = hidden_states.size(-1)196 hidden_states = hidden_states.reshape(-1, hidden_dim)197 org_hidden_states_shape = hidden_states.shape198 routing_map = expert_mask.sum(dim=1)199 200 local_permuted_hidden_states, local_input_permutation_mapping = _permute(201 hidden_states, routing_map202 )203 expected_tokens = sum(input_splits)204 actual_tokens = local_permuted_hidden_states.shape[0]205 if expected_tokens != actual_tokens:206 raise RuntimeError(207 f"EP split mismatch: input_splits sum ({expected_tokens}) != "208 f"permuted tokens ({actual_tokens})"209 )210 211 global_permuted_hidden_states = _all_to_all(212 group, local_permuted_hidden_states, output_splits, input_splits213 )214 num_local_experts = num_experts // dist.get_world_size(group)215 permute_order = (216 torch.arange(num_experts).reshape(-1, num_local_experts).T.ravel().tolist()217 )218 split_sizes = num_global_tokens_per_local_expert.ravel().tolist()219 global_permuted_hidden_states = _sort_chunks_by_idxs(220 global_permuted_hidden_states, split_sizes, permute_order221 )222 return (223 global_permuted_hidden_states,224 routing_map,225 local_input_permutation_mapping,226 org_hidden_states_shape,227 )228 229 230def tokens_post_all2all(231 expert_outputs: torch.Tensor,232 routing_weights: torch.Tensor,233 selected_experts: torch.Tensor,234 num_experts: int,235 input_splits: List[int],236 output_splits: List[int],237 num_global_tokens_per_local_expert: torch.Tensor,238 routing_map: torch.Tensor,239 local_input_permutation_mapping: torch.Tensor,240 org_hidden_states_shape: torch.Size,241 group: Optional[dist.ProcessGroup] = None,242) -> torch.Tensor:243 group = group or dist.group.WORLD244 num_local_experts = num_experts // dist.get_world_size(group)245 unpermute_order = (246 torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()247 )248 split_sizes = num_global_tokens_per_local_expert.T.ravel().tolist()249 expert_outputs = _sort_chunks_by_idxs(250 expert_outputs, split_sizes, unpermute_order251 )252 unpermute_outputs = _all_to_all(group, expert_outputs, input_splits, output_splits)253 weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)254 unpermute_outputs = _unpermute(255 unpermute_outputs,256 weights_idx,257 org_hidden_states_shape,258 local_input_permutation_mapping,259 routing_map,260 )261 return unpermute_outputs262 263 264def expert_forward(265 x: torch.Tensor,266 gate_proj: torch.nn.Linear,267 up_proj: torch.nn.Linear,268 down_proj: torch.nn.Linear,269) -> torch.Tensor:270 gate = torch.nn.functional.silu(gate_proj(x))271 up = up_proj(x)272 return down_proj(gate * up)273 274 275def solution(276 hidden_states: torch.Tensor,277 gate_weight: torch.Tensor,278 gate_bias: Optional[torch.Tensor],279 gate_proj: torch.nn.Linear,280 up_proj: torch.nn.Linear,281 down_proj: torch.nn.Linear,282 num_experts: int,283 top_k: int,284 group: Optional[dist.ProcessGroup] = None,285) -> torch.Tensor:286 if group is None:287 group = _resolve_ep_group_for_narrow_moe(num_experts)288 hidden_dim = hidden_states.size(-1)289 num_tokens = hidden_states.reshape(-1, hidden_dim).size(0)290 291 router_logits = torch.nn.functional.linear(292 hidden_states.reshape(-1, hidden_dim), gate_weight, gate_bias293 )294 routing_weights, selected_experts = torch.topk(295 torch.softmax(router_logits, dim=-1), top_k, dim=-1296 )297 expert_mask = torch.nn.functional.one_hot(298 selected_experts, num_classes=num_experts299 ).permute(2, 1, 0)300 301 input_splits, output_splits, num_global_tokens_per_local_expert, _ = _preprocess(302 expert_mask, num_experts, group303 )304 305 (306 global_permuted_hidden_states,307 routing_map,308 local_input_permutation_mapping,309 org_hidden_states_shape,310 ) = token_pre_all2all(311 hidden_states,312 expert_mask,313 num_experts,314 input_splits,315 output_splits,316 num_global_tokens_per_local_expert,317 group,318 )319 320 expert_outputs = expert_forward(321 global_permuted_hidden_states, gate_proj, up_proj, down_proj322 )323 324 out = tokens_post_all2all(325 expert_outputs,326 routing_weights,327 selected_experts,328 num_experts,329 input_splits,330 output_splits,331 num_global_tokens_per_local_expert,332 routing_map,333 local_input_permutation_mapping,334 org_hidden_states_shape,335 group,336 )337 return out338 