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
0likes263downloads
36_ulysses_gather_seq_scatter_heads.py76 linesDownload Raw Back to reference
1from typing import Optional2 3import torch4import torch.distributed as dist5from torch.distributed import ProcessGroup6 7 8def _all_to_all(9    local_input: torch.Tensor,10    scatter_dim: int,11    gather_dim: int,12    group: dist.ProcessGroup,13) -> torch.Tensor:14    seq_world_size = dist.get_world_size(group)15    input_list = [t.contiguous() for t in torch.tensor_split(local_input, seq_world_size, scatter_dim)]16    output_list = [torch.empty_like(input_list[0]) for _ in range(seq_world_size)]17    dist.all_to_all(output_list, input_list, group=group)18    return torch.cat(output_list, dim=gather_dim).contiguous()19 20 21def _all_to_all_single(22    x: torch.Tensor,23    scatter_dim: int,24    gather_dim: int,25    group: dist.ProcessGroup,26) -> torch.Tensor:27    sp_world_size = dist.get_world_size(group)28    assert scatter_dim <= 1 and gather_dim <= 129    if scatter_dim != 0:30        gather_dim_bef = x.shape[gather_dim]31        scatter_dim_bef = x.shape[scatter_dim]32        x = (33            x.reshape(34                [gather_dim_bef, sp_world_size, scatter_dim_bef // sp_world_size] + list(x.shape[2:])35            )36            .transpose(0, 1)37            .reshape(38                [gather_dim_bef * sp_world_size, scatter_dim_bef // sp_world_size] + list(x.shape[2:])39            )40            .contiguous()41        )42    output = torch.empty_like(x)43    dist.all_to_all_single(output, x.contiguous(), group=group)44    if scatter_dim == 0:45        output = torch.cat(output.split(x.size(0) // sp_world_size), dim=gather_dim)46    return output47 48 49def _all_to_all_tensor(50    x: torch.Tensor,51    scatter_dim: int,52    gather_dim: int,53    group: dist.ProcessGroup,54) -> torch.Tensor:55    if scatter_dim <= 1 and gather_dim <= 1:56        return _all_to_all_single(x, scatter_dim, gather_dim, group)57    return _all_to_all(x, scatter_dim, gather_dim, group)58 59 60def solution(61    x: torch.Tensor,62    seq_dim: int,63    head_dim: int,64    group: Optional[ProcessGroup] = None,65    unpadded_dim_size: int = 0,66) -> torch.Tensor:67    group = group or dist.group.WORLD68    sp_world = dist.get_world_size(group)69    x = _all_to_all_tensor(x, scatter_dim=head_dim, gather_dim=seq_dim, group=group)70    if unpadded_dim_size and unpadded_dim_size % sp_world != 0:71        padding_size = x.size(seq_dim) - unpadded_dim_size72        slc = [slice(None)] * x.dim()73        slc[seq_dim] = slice(0, -padding_size)74        x = x[tuple(slc)]75    return x76