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.
0259
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 