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
0likes259downloads
38_ulysses_gather_seq_scatter_heads_qkv.py189 linesDownload Raw Back to reference
1from typing import Any, Optional, Tuple2 3import torch4import torch.distributed as dist5from torch import Tensor6from torch.distributed import ProcessGroup7 8 9def _pad_tensor(x: Tensor, dim: int, padding_size: int, padding_value: int = 0) -> Tensor:10    shape = list(x.shape)11    shape[dim] = padding_size12    pad = torch.full(shape, padding_value, dtype=x.dtype, device=x.device)13    return torch.cat([x, pad], dim=dim)14 15 16def _unpad_tensor(x: Tensor, dim: int, padding_size: int) -> Tensor:17    slc = [slice(None)] * len(x.shape)18    slc[dim] = slice(0, -padding_size)19    return x[tuple(slc)]20 21 22def _all_to_all_single(23    x: Tensor,24    scatter_dim: int,25    gather_dim: int,26    group: Optional[dist.ProcessGroup] = None,27    async_op: bool = False,28):29    group = group or dist.group.WORLD30    sp_world_size = dist.get_world_size(group)31    assert scatter_dim <= 1, "scatter_dim must be 0 or 1 when using all_to_all_single!"32    assert gather_dim <= 1, "gather_dim must be 0 or 1 when using all_to_all_single!"33    if scatter_dim != 0:34        gather_dim_bef = x.shape[gather_dim]35        scatter_dim_bef = x.shape[scatter_dim]36        x = (37            x.reshape(38                [gather_dim_bef, sp_world_size, scatter_dim_bef // sp_world_size]39                + list(x.shape[2:])40            )41            .transpose(0, 1)42            .reshape(43                [gather_dim_bef * sp_world_size, scatter_dim_bef // sp_world_size]44                + list(x.shape[2:])45            )46            .contiguous()47        )48 49    output = torch.empty_like(x)50    comm = dist.all_to_all_single(output, x.contiguous(), group=group, async_op=async_op)51 52    if async_op:53 54        def wait():55            comm.wait()56            if scatter_dim == 0:57                return torch.cat(output.split(x.size(0) // sp_world_size), dim=gather_dim)58            else:59                return output60 61        return wait62 63    if scatter_dim == 0:64        output = torch.cat(output.split(x.size(0) // sp_world_size), dim=gather_dim)65    return output66 67 68def _all_to_all(69    local_input: Tensor,70    scatter_dim: int,71    gather_dim: int,72    group: Optional[dist.ProcessGroup] = None,73    async_op: bool = False,74):75    group = group or dist.group.WORLD76    seq_world_size = dist.get_world_size(group)77    input_list = [78        t.contiguous()79        for t in torch.tensor_split(local_input, seq_world_size, scatter_dim)80    ]81    output_list = [torch.empty_like(input_list[0]) for _ in range(seq_world_size)]82    comm = dist.all_to_all(output_list, input_list, group=group, async_op=async_op)83    if async_op:84 85        def wait():86            comm.wait()87            return torch.cat(output_list, dim=gather_dim).contiguous()88 89        return wait90    return torch.cat(output_list, dim=gather_dim).contiguous()91 92 93def _all_to_all_tensor(94    x: Tensor,95    scatter_dim: int,96    gather_dim: int,97    group: dist.ProcessGroup,98    async_op: bool = False,99):100    if scatter_dim <= 1 and gather_dim <= 1:101        return _all_to_all_single(x, scatter_dim, gather_dim, group, async_op)102    return _all_to_all(x, scatter_dim, gather_dim, group, async_op)103 104 105class _SeqAllToAll(torch.autograd.Function):106    @staticmethod107    def forward(108        ctx: Any,109        group: dist.ProcessGroup,110        local_input: Tensor,111        scatter_dim: int,112        gather_dim: int,113        async_op: bool,114    ) -> Tensor:115        ctx.group = group116        ctx.scatter_dim = scatter_dim117        ctx.gather_dim = gather_dim118        ctx.async_op = async_op119        return _all_to_all_tensor(local_input, scatter_dim, gather_dim, group, async_op)120 121    @staticmethod122    def backward(ctx: Any, *grad_output: Tensor) -> Tuple[None, Tensor, None, None, None]:123        if ctx.async_op:124            input_t = torch.cat(grad_output[1:], dim=ctx.gather_dim).contiguous()125        else:126            input_t = grad_output[0]127        return (128            None,129            _all_to_all_tensor(130                input_t, ctx.gather_dim, ctx.scatter_dim, ctx.group, False131            ),132            None,133            None,134            None,135        )136 137 138def gather_seq_scatter_heads_qkv(139    qkv_tensor: Tensor,140    seq_dim: int,141    unpadded_dim_size: Optional[int] = None,142    restore_shape: bool = True,143    async_op: bool = False,144    group: Optional[ProcessGroup] = None,145) -> Tensor:146    group = group or dist.group.WORLD147    if not group:148        return qkv_tensor149    sp_world = dist.get_world_size(group)150    orig_shape = qkv_tensor.shape151    scatter_dim = qkv_tensor.dim()152    bef_all2all_shape = list(orig_shape)153    qkv_proj_dim = bef_all2all_shape[-1]154    bef_all2all_shape = bef_all2all_shape[:-1] + [3, qkv_proj_dim // 3]155    qkv_tensor = qkv_tensor.view(bef_all2all_shape)156    if async_op:157        return _SeqAllToAll.apply(group, qkv_tensor, scatter_dim, seq_dim, async_op)158    qkv_tensor = _SeqAllToAll.apply(group, qkv_tensor, scatter_dim, seq_dim, async_op)159 160    if restore_shape:161        out_shape = list(orig_shape)162        out_shape[seq_dim] *= sp_world163        out_shape[-1] = qkv_proj_dim // sp_world164        qkv_tensor = qkv_tensor.view(out_shape)165 166    if unpadded_dim_size and unpadded_dim_size % sp_world != 0:167        padding_size = qkv_tensor.size(seq_dim) - unpadded_dim_size168        qkv_tensor = _unpad_tensor(qkv_tensor, seq_dim, padding_size)169 170    return qkv_tensor171 172 173def solution(174    qkv_tensor: torch.Tensor,175    seq_dim: int,176    group: Optional[ProcessGroup] = None,177    unpadded_dim_size: Optional[int] = None,178    restore_shape: bool = True,179) -> torch.Tensor:180    group = group or dist.group.WORLD181    return gather_seq_scatter_heads_qkv(182        qkv_tensor,183        seq_dim=seq_dim,184        unpadded_dim_size=unpadded_dim_size or 0,185        restore_shape=restore_shape,186        async_op=False,187        group=group,188    )189