CoolFace
Datasetpublic

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.

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
0likes266downloads
67_gnn_sparse_embedding_all2all.py68 linesDownload Raw Back to reference
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