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.
0266
1from typing import Optional, Tuple2 3import torch4import torch.distributed as dist5 6 7def _generate_permutation_remainder(8 idx: torch.Tensor,9 world_size: int,10) -> Tuple[torch.Tensor, torch.Tensor]:11 owner = (idx % world_size).long()12 send_splits = torch.bincount(owner, minlength=world_size)13 perm = torch.argsort(owner, stable=True).long()14 return perm, send_splits15 16 17@torch.no_grad()18def solution(19 idx: torch.Tensor,20 value: torch.Tensor,21 num_nodes: int,22 group: Optional[dist.ProcessGroup] = None,23) -> Tuple[torch.Tensor, torch.Tensor]:24 group = group or dist.group.WORLD25 world_size = dist.get_world_size(group)26 if world_size == 1:27 return idx, value28 29 perm, send_splits = _generate_permutation_remainder(idx, world_size)30 31 recv_splits = torch.empty_like(send_splits)32 dist.all_to_all_single(recv_splits, send_splits, group=group)33 34 recv_splits = recv_splits.to("cpu", non_blocking=True)35 send_splits = send_splits.to("cpu", non_blocking=True)36 send_idx = idx[perm]37 send_value = value[perm]38 if idx.is_cuda:39 torch.cuda.current_stream().synchronize()40 41 recv_count = int(recv_splits.sum().item())42 recv_splits_list = recv_splits.tolist()43 send_splits_list = send_splits.tolist()44 45 recv_idx = torch.empty((recv_count,), dtype=idx.dtype, device=idx.device)46 dist.all_to_all_single(47 recv_idx,48 send_idx,49 output_split_sizes=recv_splits_list,50 input_split_sizes=send_splits_list,51 group=group,52 )53 54 recv_value = torch.empty(55 (recv_count, *value.shape[1:]),56 dtype=value.dtype,57 device=value.device,58 )59 dist.all_to_all_single(60 recv_value,61 send_value,62 output_split_sizes=recv_splits_list,63 input_split_sizes=send_splits_list,64 group=group,65 )66 67 return recv_idx, recv_value68 