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.
0123
1from typing import List, Optional, Union2 3import torch4import torch.distributed as dist5 6 7def _sort_chunks_by_idxs(8 input: torch.Tensor,9 split_sizes: Union[torch.Tensor, List[int]],10 sorted_idxs: List[int],11) -> torch.Tensor:12 if isinstance(split_sizes, torch.Tensor):13 split_sizes = split_sizes.tolist()14 chunks = torch.split(input, split_sizes, dim=0)15 return torch.cat([chunks[i] for i in sorted_idxs], dim=0)16 17 18def _all_to_all_forward(19 group: dist.ProcessGroup,20 input: torch.Tensor,21 output_split_sizes: Optional[List[int]],22 input_split_sizes: Optional[List[int]],23) -> torch.Tensor:24 if dist.get_world_size(group) == 1:25 return input.contiguous()26 input = input.contiguous()27 out_size = sum(output_split_sizes) if output_split_sizes else input.size(0)28 output = torch.empty((out_size, input.size(1)), dtype=input.dtype, device=input.device)29 dist.all_to_all_single(30 output, input,31 output_split_sizes=output_split_sizes,32 input_split_sizes=input_split_sizes,33 group=group,34 )35 return output36 37 38def _generate_weights_idx(39 routing_weights: torch.Tensor,40 selected_experts: torch.Tensor,41 num_experts: int,42) -> torch.Tensor:43 num_tokens, topk = routing_weights.shape44 weights_idx = torch.zeros(45 (num_tokens, num_experts), dtype=routing_weights.dtype, device=routing_weights.device46 )47 weights_idx.scatter_add_(1, selected_experts, routing_weights)48 return weights_idx49 50 51def _unpermute(52 tokens: torch.Tensor,53 routing_weights: torch.Tensor,54 hidden_states_shape: torch.Size,55 permutation_mapping: torch.Tensor,56 routing_map: torch.Tensor,57) -> torch.Tensor:58 tokens_weight = routing_weights.T.contiguous().masked_select(routing_map.bool())59 tokens = tokens * tokens_weight.unsqueeze(-1)60 hidden_dim = hidden_states_shape[-1]61 unpermuted_tokens = torch.zeros(hidden_states_shape, device=tokens.device, dtype=tokens.dtype)62 expanded_mapping = permutation_mapping.unsqueeze(1).expand(-1, hidden_dim)63 unpermuted_tokens.scatter_add_(0, expanded_mapping, tokens)64 return unpermuted_tokens65 66 67def solution(68 expert_outputs: torch.Tensor,69 routing_weights: torch.Tensor,70 selected_experts: torch.Tensor,71 num_experts: int,72 input_splits: Union[List[int], torch.Tensor],73 output_splits: Union[List[int], torch.Tensor],74 num_global_tokens_per_local_expert: torch.Tensor,75 routing_map: torch.Tensor,76 local_input_permutation_mapping: torch.Tensor,77 org_hidden_states_shape: torch.Size,78 group: Optional[dist.ProcessGroup] = None,79) -> torch.Tensor:80 group = group or dist.group.WORLD81 num_local_experts = num_experts // dist.get_world_size(group)82 unpermute_order = torch.arange(num_experts).reshape(num_local_experts, -1).T.ravel().tolist()83 84 expert_outputs = _sort_chunks_by_idxs(85 expert_outputs,86 num_global_tokens_per_local_expert.T.ravel(),87 unpermute_order,88 )89 90 unpermute_outputs = _all_to_all_forward(group, expert_outputs, input_splits, output_splits)91 92 weights_idx = _generate_weights_idx(routing_weights, selected_experts, num_experts)93 unpermute_outputs = _unpermute(94 unpermute_outputs,95 weights_idx,96 org_hidden_states_shape,97 local_input_permutation_mapping,98 routing_map,99 )100 101 return unpermute_outputs102 